File size: 1,777 Bytes
3b7b3fb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 | """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")
|