"""Standalone loader for stratus-labs/nocturne-v1-mini (EfficientNet-B1 student, distilled from the AST teacher). from model import load_nocturne_mini, predict_file model, vocab, thresholds = load_nocturne_mini(path) # strict load print(predict_file(model, "clip.wav", vocab, top_k=5)) Requires: torch, torchaudio, timm, safetensors, soundfile. ~9.4 M params; runs on CPU. """ from __future__ import annotations import json from pathlib import Path import timm, torch, torch.nn as nn, torch.nn.functional as F import torchaudio.transforms as T SR, N_MELS, HOP, WIN, N_FFT, MAX_LEN = 16_000, 128, 160, 400, 400, 1024 AUDIOSET_MEAN, AUDIOSET_STD = -4.2677393, 4.5689974 class LogMel(nn.Module): def __init__(self): super().__init__() self.mel = T.MelSpectrogram(sample_rate=SR, n_fft=N_FFT, hop_length=HOP, win_length=WIN, n_mels=N_MELS, f_min=0, f_max=SR // 2, power=2.0) def forward(self, wav: torch.Tensor) -> torch.Tensor: x = (self.mel(wav) + 1e-6).log().transpose(1, 2) t = x.shape[1] x = F.pad(x, (0, 0, 0, MAX_LEN - t)) if t < MAX_LEN else x[:, :MAX_LEN, :] return (x - AUDIOSET_MEAN) / (AUDIOSET_STD * 2) class EffNetStudent(nn.Module): def __init__(self, num_classes: int, name: str = "efficientnet_b1"): super().__init__() self.logmel = LogMel() self.backbone = timm.create_model(name, pretrained=False, in_chans=1, num_classes=0) d = self.backbone.num_features self.head = nn.Sequential(nn.LayerNorm(d), nn.Dropout(0.2), nn.Linear(d, num_classes)) def forward(self, wav: torch.Tensor) -> torch.Tensor: return self.head(self.backbone(self.logmel(wav).unsqueeze(1))) def load_nocturne_mini(path: str | Path, device: str = "cpu"): from safetensors.torch import load_file path = Path(path) cfg = json.loads((path / "config.json").read_text()) vocab_map = json.loads((path / "vocab.json").read_text()) vocab = [None] * len(vocab_map) for name, idx in vocab_map.items(): vocab[idx] = name model = EffNetStudent(num_classes=cfg["num_labels"]) missing, unexpected = model.load_state_dict(load_file(str(path / "model.safetensors")), strict=False) if missing or unexpected: raise RuntimeError(f"checkpoint/model mismatch: missing={len(missing)} unexpected={len(unexpected)} " f"(e.g. {sorted(missing)[:2]} / {sorted(unexpected)[:2]}); refusing a partial load") thresholds = {} tp = path / "thresholds.json" if tp.exists(): tj = json.loads(tp.read_text()); thresholds = tj.get("per_class", tj) return model.to(device).eval(), vocab, thresholds def load_wav_16k(fp: str | Path, seconds: float = 10.24) -> torch.Tensor: import soundfile as sf, torchaudio.functional as AF wav, sr = sf.read(str(fp), dtype="float32", always_2d=True) wav = torch.from_numpy(wav.mean(axis=1)) if sr != SR: wav = AF.resample(wav, sr, SR) n = int(seconds * SR) return (wav[:n] if wav.numel() >= n else F.pad(wav, (0, n - wav.numel()))).unsqueeze(0) @torch.no_grad() def predict_file(model, fp, vocab, top_k: int = 5, threshold: float | None = None, device: str = "cpu"): probs = torch.sigmoid(model(load_wav_16k(fp).to(device)))[0].cpu() order = probs.argsort(descending=True) out = [(vocab[i], float(probs[i])) for i in order[: (len(order) if threshold is not None else top_k)]] return [(s, p) for s, p in out if p >= threshold] if threshold is not None else out