File size: 3,569 Bytes
9818d29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""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