"""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) # Lengths deliberately differ from the export sample, to exercise the dynamic axes. 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}" # A padded batch must give each item the same answer it gets alone. 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")