BSG-BAT / export_onnx.py
tphakala's picture
Add BSG-BAT v0.21 ONNX ensemble (6 checkpoints), labels, original preprocessing code, model card
a4d9f29 verified
Raw History Blame Contribute Delete
3.21 kB
"""
Export BSG-BAT PyTorch checkpoints to ONNX.
The model (code/supervised.py `Net`) is a plain 6-conv + 3-fc CNN.
Input : float32 [batch, 1, 512, 128] (1 channel, ntime=512, nfreq=128)
a log10 mel spectrogram of 384 kHz mono audio
(n_fft=1024, hop=768, n_mels=128, fmin=9000, fmax=150000), per-segment normalized.
Output : float32 [batch, 22] logits (21 European bat species + Background).
probability = sigmoid(logit) (multi-label, BCEWithLogits training).
Usage: python export_onnx.py <model_in.pt> <model_out.onnx>
"""
import sys
import torch
import torch.nn as nn
import torch.nn.functional as F
NTIME = 512
NFREQ = 128
NCLASSES = 22
class Net(nn.Module):
# Copied verbatim from bsgbat/code/supervised.py so we do not depend on its sys.path.
def __init__(self, ntime=NTIME, nfreq=NFREQ, nclasses=NCLASSES):
super().__init__()
self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.conv3 = nn.Conv2d(64, 128, 3, padding=1)
self.conv4 = nn.Conv2d(128, 256, 3, padding=1)
self.conv5 = nn.Conv2d(256, 512, 3, padding=1)
self.conv6 = nn.Conv2d(512, 512, 3, padding=1)
n_maxpool = 6
nt, nr = ntime, nfreq
for _ in range(n_maxpool):
nt //= 2
nr //= 2
nr *= 4
self.fc1 = nn.Linear(512 * nt * nr, 512)
self.fc2 = nn.Linear(512, 128)
self.fc3 = nn.Linear(128, nclasses)
def forward(self, x):
x = F.max_pool2d(F.relu(self.conv1(x)), (2, 2))
x = F.max_pool2d(F.relu(self.conv2(x)), 2)
x = F.max_pool2d(F.relu(self.conv3(x)), 2)
x = F.max_pool2d(F.relu(self.conv4(x)), (2, 2))
x = F.max_pool2d(F.relu(self.conv5(x)), (2, 1))
x = F.max_pool2d(F.relu(self.conv6(x)), (2, 1))
x = torch.flatten(x, 1)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
x = self.fc3(x)
return x
def main(argv):
model_in, model_out = argv[0], argv[1]
device = torch.device("cpu")
model = Net(nclasses=NCLASSES)
state = torch.load(model_in, map_location=device)
model.load_state_dict(state)
model.eval()
dummy = torch.randn(1, 1, NTIME, NFREQ, dtype=torch.float32)
torch.onnx.export(
model,
dummy,
model_out,
input_names=["spectrogram"],
output_names=["logits"],
dynamic_axes={"spectrogram": {0: "batch"}, "logits": {0: "batch"}},
opset_version=17,
do_constant_folding=True,
)
print(f"exported {model_in} -> {model_out}")
# Numerical parity check PyTorch vs onnxruntime.
import numpy as np
import onnxruntime as ort
with torch.no_grad():
ref = model(dummy).numpy()
sess = ort.InferenceSession(model_out, providers=["CPUExecutionProvider"])
got = sess.run(["logits"], {"spectrogram": dummy.numpy()})[0]
max_abs = float(np.max(np.abs(ref - got)))
print(f"output shape {got.shape}, max abs diff PyTorch vs ORT = {max_abs:.3e}")
assert got.shape == (1, NCLASSES), got.shape
assert max_abs < 1e-4, max_abs
print("PARITY_OK")
if __name__ == "__main__":
main(sys.argv[1:])