File size: 5,814 Bytes
6eed659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
#!/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()