| """Self-check: ONNX matches PyTorch, and padding does not change results. |
| |
| python test_ambernet.py [model_dir] [ambernet.onnx] |
| """ |
|
|
| import sys |
|
|
| import numpy as np |
| import onnxruntime as ort |
| import torch |
|
|
| from modeling_ambernet import AmberNet |
|
|
| model_dir = sys.argv[1] if len(sys.argv) > 1 else "." |
| onnx_path = sys.argv[2] if len(sys.argv) > 2 else "ambernet.onnx" |
|
|
| model = AmberNet.from_pretrained(model_dir) |
| session = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"]) |
| rng = np.random.default_rng(0) |
|
|
| |
| for seconds in (1.5, 7.0): |
| audio = rng.standard_normal((1, int(16000 * seconds))).astype(np.float32) * 0.1 |
| lens = np.array([audio.shape[1]], dtype=np.int64) |
| torch_logits = model(torch.from_numpy(audio), torch.from_numpy(lens))[0].detach().numpy() |
| onnx_logits = session.run(None, {"audio": audio, "audio_len": lens})[0] |
| diff = np.abs(torch_logits - onnx_logits).max() |
| print(f"{seconds:>4}s onnx-vs-torch max diff {diff:.2e}") |
| assert diff < 1e-3, f"ONNX diverged from PyTorch: {diff}" |
|
|
| |
| short = rng.standard_normal(16000 * 2).astype(np.float32) * 0.1 |
| long = rng.standard_normal(16000 * 5).astype(np.float32) * 0.1 |
| batch = np.zeros((2, len(long)), dtype=np.float32) |
| batch[0, : len(short)] = short |
| batch[1] = long |
| lens = np.array([len(short), len(long)], dtype=np.int64) |
| batched = session.run(None, {"audio": batch, "audio_len": lens})[0] |
| solo = session.run(None, {"audio": short[None], "audio_len": lens[:1]})[0] |
| diff = np.abs(batched[0] - solo[0]).max() |
| print(f"padding invariance max diff {diff:.2e}") |
| assert diff < 1e-3, f"padding changed the result: {diff}" |
|
|
| print("OK") |
|
|