MiniCPM5-2B-WebGPU-Pi / scripts /validate_onnx.py
Mike0021's picture
Load verified weights from pinned companion model repository
17f3094 verified
Raw History Blame Contribute Delete
3.73 kB
"""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)