"""Compare cached ONNX inference against the official PyTorch trajectories.""" import json from pathlib import Path import time import numpy as np import onnxruntime as ort out = Path('validation') fixtures = json.loads((out / 'fixtures.json').read_text()) bf16 = np.load(out / 'reference-bf16.npz') fp16 = np.load(out / 'reference-fp16.npz') def metrics(reference, actual): r = reference.astype(np.float64); a = actual.astype(np.float64) p = np.exp(r - r.max(-1, keepdims=True)); p /= p.sum(-1, keepdims=True) q = np.exp(a - a.max(-1, keepdims=True)); q /= q.sum(-1, keepdims=True) return {'rmse': float(np.sqrt(np.mean((a-r)**2))), 'max_abs': float(np.max(np.abs(a-r))), 'cosine': float(np.mean(np.sum(a*r,-1)/(np.linalg.norm(a,axis=-1)*np.linalg.norm(r,axis=-1)))), 'argmax_agreement': float(np.mean(np.argmax(a,-1)==np.argmax(r,-1))), 'mean_kl_reference_to_actual': float(np.mean(np.sum(p*np.log((p+1e-30)/(q+1e-30)),axis=-1)))} reports = {} for name, path in [('fp16', 'models/own-fp16/model.onnx'), ('int4', 'models/own-int4/model.onnx'), ('packaged', 'models/minicpm5-webgpu/onnx/model_q4f16.onnx')]: options = ort.SessionOptions(); options.intra_op_num_threads = 8; options.inter_op_num_threads = 1 start = time.time() session = ort.InferenceSession(path, sess_options=options, providers=['CPUExecutionProvider']) print('Loaded', name, time.time()-start, flush=True) output_names = [o.name for o in session.get_outputs()] samples = {} for fixture in fixtures: cache = {i.name: np.zeros((1,2,0,128), dtype=np.float16) for i in session.get_inputs() if i.name.startswith('past_key_values')} ids = np.array([fixture['input_ids']], dtype=np.int64) total = ids.shape[1] logits = [] for next_id in fixture['continuation']: feeds = {'input_ids': ids, 'attention_mask': np.ones((1,total),dtype=np.int64), **cache} results = dict(zip(output_names, session.run(None, feeds))) assert results['logits'].shape == (1,1,130560) assert np.isfinite(results['logits']).all() logits.append(results['logits'][0,-1].astype(np.float32)) cache = {key.replace('present', 'past_key_values'): value for key,value in results.items() if key.startswith('present')} ids = np.array([[next_id]],dtype=np.int64); total += 1 samples[fixture['id']] = np.stack(logits) print(name, fixture['id'], 'checked', len(logits), 'steps', flush=True) np.savez(out / f'onnx-{name}.npz', **samples) actual = np.concatenate(list(samples.values())) reports[name] = {'vs_bf16': metrics(np.concatenate([bf16[f['id']] for f in fixtures]),actual), 'vs_fp16': metrics(np.concatenate([fp16[f['id']] for f in fixtures]),actual), 'positions': len(actual), 'elapsed_s': time.time()-start} if name == 'packaged': raw = np.load(out/'onnx-int4.npz') reports[name]['vs_unmodified_export'] = metrics(np.concatenate([raw[f['id']] for f in fixtures]), actual) assert all(np.array_equal(samples[f['id']],raw[f['id']]) for f in fixtures), 'Packaging changed logits' print(json.dumps(reports[name]), flush=True) (out/'onnx-validation.json').write_text(json.dumps(reports,indent=2)) del session, cache, results assert reports['fp16']['vs_fp16']['cosine'] > 0.999, 'Unquantized export deviates from PyTorch' assert reports['fp16']['vs_fp16']['mean_kl_reference_to_actual'] < 0.01, 'Unquantized export KL too high' print('Export parity and lossless packaging checks passed', flush=True)