#!/usr/bin/env python3 """SUTA — Single-Utterance Test-time Adaptation (Lin et al., Interspeech 2022). Attaque la CAUSE RACINE diagnostiquée : sur-apprentissage locuteur (Phase 2 = locuteurs inédits). Pour CHAQUE clip : on adapte uniquement les paramètres affines des LayerNorm de l'encodeur en minimisant l'incertitude du modèle sur ce clip, puis on décode, puis on RESET. Perte = 0.3*entropie + 0.7*MCC (minimum class confusion). Garde-fous indispensables : - SGEM : exclure les trames où le BLANK est argmax du calcul d'entropie (sinon la minimisation dégénère vers le blank et vide la sortie). - reset des poids après chaque clip (aucune accumulation). - si la longueur de sortie chute de >15% vs baseline sur un clip -> on garde la baseline. Gate : devhard (locuteurs disjoints = le bon banc). Règle : si gain < 0.005, NE PAS soumettre. Conforme : aucune donnée externe, aucun label, inférence seule. """ import copy, json, os import jiwer, numpy as np, soundfile as sf, torch from transformers import AutoModelForCTC, AutoProcessor LANG = os.environ.get("LANG_ASR", "sna") MODEL = os.environ.get("MODEL", "/root/models/sna_ps_best") STEPS = int(os.environ.get("STEPS", "4")) LR = float(os.environ.get("LR", "2e-4")) W_ENT = float(os.environ.get("W_ENT", "0.3")) W_MCC = float(os.environ.get("W_MCC", "0.7")) NCLIP = int(os.environ.get("NCLIP", "0")) # 0 = tous AUD = "/root/devhard_audio" SR = 16000 def comb(refs, hyps): pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()] a = [x for x, _ in pr]; b = [y for _, y in pr] w = jiwer.wer(a, b); c = jiwer.cer(a, b) return w, c, 0.5 * w + 0.5 * c def losses(logits, blank): """logits: (1,T,V). Entropie SGEM (hors trames blank) + MCC.""" lp = logits.log_softmax(-1) p = lp.exp() keep = logits.argmax(-1).squeeze(0) != blank # SGEM : ignorer les trames blank if keep.sum() < 2: keep = torch.ones_like(keep, dtype=torch.bool) pk = p.squeeze(0)[keep] # (T', V) ent = -(pk * pk.clamp_min(1e-9).log()).sum(-1).mean() # MCC : décourager la confusion entre classes sur les trames retenues corr = pk.t() @ pk # (V,V) off = corr - torch.diag(torch.diag(corr)) mcc = off.sum() / max(pk.shape[0], 1) return ent, mcc def main(): rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] sub = [r for r in rows if r["lang"] == LANG] for r in sub: r["audio"] = os.path.join(AUD, os.path.basename(r["audio"])) sub = [r for r in sub if os.path.exists(r["audio"])] if NCLIP: sub = sub[:NCLIP] refs = [r["text"] for r in sub] print("SUTA %s : %d clips | steps=%d lr=%g (ent %.1f / mcc %.1f)" % (LANG, len(sub), STEPS, LR, W_ENT, W_MCC), flush=True) proc = AutoProcessor.from_pretrained(MODEL) model = AutoModelForCTC.from_pretrained(MODEL, dtype=torch.float32).cuda() tok = proc.tokenizer blank = tok.pad_token_id # ne rendre entraînables QUE les affines des LayerNorm for p in model.parameters(): p.requires_grad_(False) ln_params = [] for mod in model.modules(): if isinstance(mod, torch.nn.LayerNorm): for p in (mod.weight, mod.bias): if p is not None: p.requires_grad_(True); ln_params.append(p) print("parametres adaptes : %d tenseurs LayerNorm (%d valeurs)" % (len(ln_params), sum(p.numel() for p in ln_params)), flush=True) init = copy.deepcopy([p.detach().clone() for p in ln_params]) def decode(lg): return " ".join(tok.decode(lg.argmax(-1)).replace("|", " ").split()) base_h, suta_h, reverted = [], [], 0 for i, r in enumerate(sub): au = sf.read(r["audio"], dtype="float32")[0] if au.ndim > 1: au = au.mean(1) x = proc(au, sampling_rate=SR, return_tensors="pt") x = {k: v.cuda() for k, v in x.items()} # --- baseline (avant adaptation) --- model.eval() with torch.inference_mode(): b = decode(model(**x).logits[0].float().cpu().numpy()) base_h.append(b) # --- adaptation --- for p, p0 in zip(ln_params, init): p.data.copy_(p0) opt = torch.optim.AdamW(ln_params, lr=LR) for _ in range(STEPS): opt.zero_grad() lg = model(**x).logits ent, mcc = losses(lg, blank) (W_ENT * ent + W_MCC * mcc).backward() torch.nn.utils.clip_grad_norm_(ln_params, 1.0) opt.step() with torch.inference_mode(): s = decode(model(**x).logits[0].float().cpu().numpy()) # garde-fou longueur if len(b.split()) and len(s.split()) < 0.85 * len(b.split()): s = b; reverted += 1 suta_h.append(s) if (i + 1) % 50 == 0: print(" %d/%d (reverts %d)" % (i + 1, len(sub), reverted), flush=True) bb = comb(refs, base_h); ss = comb(refs, suta_h) print("\nbaseline greedy : WER %.4f CER %.4f COMBINE %.4f" % bb, flush=True) print("SUTA : WER %.4f CER %.4f COMBINE %.4f" % ss, flush=True) d = ss[2] - bb[2] print("=> gain %+.4f | reverts %d/%d" % (d, reverted, len(sub)), flush=True) print("VERDICT : %s" % ("✅ GO (>=0.005)" if d <= -0.005 else ("~ trop faible, NE PAS soumettre" if d < 0 else "❌ degrade")), flush=True) nch = sum(1 for a, c in zip(base_h, suta_h) if a != c) print("clips modifies : %d/%d" % (nch, len(sub)), flush=True) json.dump({"base": bb[2], "suta": ss[2], "gain": d, "reverts": reverted}, open("/root/suta_%s.json" % LANG, "w")) print("SUTA_DONE", flush=True) if __name__ == "__main__": main()