Download export_onnx.py from tphakala/BSG-BAT: direct link, hf CLI and curl.
- Browser
- Download file 3.21 kB
-
https://huggingface.co/tphakala/BSG-BAT/resolve/main/export_onnx.py
- Command line
-
hf download hf://tphakala/BSG-BAT/export_onnx.py
-
curl -L -o export_onnx.py https://huggingface.co/tphakala/BSG-BAT/resolve/main/export_onnx.py
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:]) | |