nocturne-v1-mini / model.py
stratusvale's picture
model.py: standalone strict loader (EffNet-B1 student)
9818d29 verified
Raw History Blame Contribute Delete
3.57 kB
"""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