Download scripts/validate_onnx.py from TechnoBaptist/MiniCPM5-2B-WebGPU-Pi: direct link, hf CLI and curl.
- Browser
- Download file 3.73 kB
-
https://huggingface.co/spaces/TechnoBaptist/MiniCPM5-2B-WebGPU-Pi/resolve/main/scripts/validate_onnx.py
- Command line
-
hf download hf://spaces/TechnoBaptist/MiniCPM5-2B-WebGPU-Pi/scripts/validate_onnx.py
-
curl -L -o validate_onnx.py https://huggingface.co/spaces/TechnoBaptist/MiniCPM5-2B-WebGPU-Pi/resolve/main/scripts/validate_onnx.py
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) | |