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")