"""Standalone inference for the LAION VocalBurst classifier (VoiceCLAP embedding -> MLP head). from inference import VocalBurstClassifier clf = VocalBurstClassifier(".") # a local checkout, or a HF repo id print(clf.predict("clip.wav")) """ import os, json, torch, torch.nn as nn, torchaudio from transformers import AutoModel class _MLP(nn.Module): def __init__(self, d, h, depth, out): super().__init__() L = [nn.Linear(d, h), nn.LayerNorm(h), nn.GELU(), nn.Dropout(0.0)] for _ in range(depth - 1): L += [nn.Linear(h, h), nn.LayerNorm(h), nn.GELU(), nn.Dropout(0.0)] L += [nn.Linear(h, out)] self.net = nn.Sequential(*L) def forward(self, x): return self.net(x) class VocalBurstClassifier: def __init__(self, path_or_repo=".", device=None): self.device = device or ("cuda" if torch.cuda.is_available() else "cpu") d = path_or_repo if not os.path.isdir(d): # a HF repo id -> download from huggingface_hub import snapshot_download d = snapshot_download(path_or_repo) self.cfg = json.load(open(f"{d}/config.json")) self.classes = json.load(open(f"{d}/classes.json")) self.groups = json.load(open(f"{d}/class_to_group.json")) self.nb = self.classes.index("no_burst") # frozen VoiceCLAP-commercial embedder (auto-downloaded from the public repo) self.vc = AutoModel.from_pretrained(self.cfg["embedder"], trust_remote_code=True).to(self.device).eval() self.mlp = _MLP(self.cfg["input_dim"], self.cfg["hidden"], self.cfg["depth"], self.cfg["num_classes"]).to(self.device).eval() self.mlp.load_state_dict(torch.load(f"{d}/model.pt", map_location=self.device)) @torch.no_grad() def probs(self, audio_path): w, sr = torchaudio.load(audio_path); w = w.mean(0) if w.dim() == 2 else w w16 = torchaudio.functional.resample(w, sr, 16000) if sr != 16000 else w if len(w16) > 16000 * 30: w16 = w16[:16000 * 30] # VoiceCLAP/Whisper sees first 30 s e = self.vc.encode_waveform(w16.to(self.device)) logit = self.mlp(e).squeeze(0) p = torch.sigmoid(logit) if self.cfg["activation"] == "sigmoid" else torch.softmax(logit, -1) return p.cpu().numpy() def predict(self, audio_path, topk=5, no_burst_gate=0.5): p = self.probs(audio_path) if float(p[self.nb]) >= no_burst_gate: # no-burst gate return {"no_burst": True, "p_no_burst": float(p[self.nb]), "predictions": []} order = p.argsort()[::-1] preds = [(self.classes[i], float(p[i])) for i in order if self.classes[i] != "no_burst"][:topk] return {"no_burst": False, "p_no_burst": float(p[self.nb]), "top1": preds[0][0], "predictions": preds, "group": self.groups[preds[0][0]]}