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