#!/usr/bin/env python3 """BANC D'ESSAI LINGALA sur la validation officielle (1844 clips, 0% de fuite de clips verifiee ; KenLM verifie propre a 0.1%). devhard est INUTILISABLE (100% de ses clips sont dans le train). But n1 : CALIBRER L'INSTRUMENT. Trois modeles ont un verdict LB connu : joint_cont = CHAMPION lin_s4 = LB -0.0072 (doit finir DERRIERE joint_cont) joint_linps = LB -0.0062 (doit finir DERRIERE joint_cont) Si la validation ne reproduit pas cet ordre, elle ne sert pas non plus a selectionner un modele, et il faudra reentrainer une baseline B0 sur train_lin_min pour retrouver un gate honnete. But n2 : donner leur premiere chance equitable aux modeles jamais retenus (mms1b_lin, mmsjoint, mmsjointpl, joint_cont2, joint_best) -- ils avaient ete ecartes sur devhard, terrain ou nos propres modeles etaient gonfles par la memorisation et eux non. """ import json, os, sys, time import numpy as np, soundfile as sf, torch from multiprocessing import Pool from pyctcdecode import build_ctcdecoder from transformers import AutoModelForCTC, AutoProcessor VAL = "/scratch/prep/manifests/waxal_lin_validation.jsonl" ARPA = "/scratch/lm/lin_5g.arpa" ALPHA, BETA, LSB, BW = 0.5, 1.0, True, 64 MODELS = [ ("joint_cont [CHAMPION]", "/root/models/joint_cont_best"), ("lin_s4 [LB -0.0072]", "/root/models/lin_s4_best"), ("joint_linps [LB -0.0062]", "/scratch/runs/joint_linps/best"), ("joint_cont2", "/root/models/joint_cont2_best"), ("joint_best", "/root/models/joint_best"), ("mms1b_lin", "/root/models/mms1b_lin_best"), ("mmsjoint", "/root/models/mmsjoint_best"), ("mmsjointpl", "/root/models/mmsjointpl_best"), ] def lev(a, b): if a == b: return 0 if not a: return len(b) if not b: return len(a) prev = list(range(len(b) + 1)) for i, ca in enumerate(a, 1): cur = [i] for j, cb in enumerate(b, 1): cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb))) prev = cur return prev[-1] def score(hyps, refs): we = wl = ce = cl = 0 per = [] for h, r in zip(hyps, refs): hw, rw = h.split(), r.split() e1, e2 = lev(hw, rw), lev(h, r) we += e1; wl += len(rw); ce += e2; cl += len(r) per.append(0.5 * e1 / max(len(rw), 1) + 0.5 * e2 / max(len(r), 1)) wer, cer = we / max(wl, 1), ce / max(cl, 1) return wer, cer, 0.5 * wer + 0.5 * cer, np.array(per) def norm(s): return " ".join(str(s).replace("|", " ").split()) def logits_for(model_dir, files): proc = AutoProcessor.from_pretrained(model_dir) m = AutoModelForCTC.from_pretrained(model_dir, dtype=torch.float32).cuda().eval() out, bs = [], 4 with torch.inference_mode(): i = 0 while i < len(files): b = files[i:i + bs] try: au = [sf.read(f, dtype="float32")[0] for f in b] x = proc(au, sampling_rate=16000, return_tensors="pt", padding=True) x = {k: v.cuda() for k, v in x.items() if k in ("input_values", "input_features", "attention_mask")} lg = m(**x).logits.log_softmax(-1).float().cpu().numpy() for j in range(len(b)): out.append(lg[j]) i += bs except torch.cuda.OutOfMemoryError: torch.cuda.empty_cache() if bs == 1: raise bs = max(1, bs // 2) print(" OOM -> batch %d" % bs, flush=True) del m; torch.cuda.empty_cache() return proc, out rows = [json.loads(l) for l in open(VAL, encoding="utf-8")] rows = [r for r in rows if os.path.exists(r["audio"])] files = [r["audio"] for r in rows] refs = [r["text"] for r in rows] print("validation lingala : %d clips (0%% de fuite de clips, verifie)" % len(files), flush=True) print("decodage KenLM : alpha=%.1f beta=%.1f lsb=%s beam=%d (config exacte du champion)\n" % (ALPHA, BETA, LSB, BW), flush=True) RES = {} for nom, path in MODELS: if not os.path.isdir(path): print("%-26s ABSENT (%s)" % (nom, path), flush=True); continue t0 = time.time() try: proc, L = logits_for(path, files) tok = proc.tokenizer v = tok.get_vocab() lab = [None] * len(v) for t, i in v.items(): lab[i] = t lab[tok.word_delimiter_token_id] = " " lab[tok.unk_token_id] = "⁇" lab[tok.pad_token_id] = "" g = [norm(tok.decode(l.argmax(-1))) for l in L] wg, cg, kg, perg = score(g, refs) dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=ALPHA, beta=BETA, lm_score_boundary=LSB) with Pool(8) as p: hl = [norm(x) for x in dec.decode_batch(p, L, beam_width=BW)] wl_, cl_, kl, perl = score(hl, refs) RES[nom] = dict(greedy=(wg, cg, kg), kenlm=(wl_, cl_, kl), perg=perg, perl=perl) print("%-26s greedy %.4f | KenLM %.4f (%.0f s)" % (nom, kg, kl, time.time() - t0), flush=True) del L except Exception as e: print("%-26s ECHEC : %s" % (nom, type(e).__name__ + " " + str(e)[:120]), flush=True) torch.cuda.empty_cache() print("\n============== CLASSEMENT (validation lingala, combine, plus bas = mieux) ==============") print("%-26s %8s %8s %9s %8s %8s %9s" % ("", "WERg", "CERg", "GREEDY", "WERlm", "CERlm", "KenLM")) for nom in sorted(RES, key=lambda k: RES[k]["kenlm"][2]): g, l = RES[nom]["greedy"], RES[nom]["kenlm"] print("%-26s %8.4f %8.4f %9.4f %8.4f %8.4f %9.4f" % (nom, g[0], g[1], g[2], l[0], l[1], l[2])) ref_key = [k for k in RES if k.startswith("joint_cont ")] if ref_key: rk = ref_key[0] base = RES[rk]["perl"] rng = np.random.default_rng(0) print("\n--- comparaison appariee vs CHAMPION, decodage KenLM (negatif = bat le champion) ---") for nom in sorted(RES, key=lambda k: RES[k]["kenlm"][2]): if nom == rk: continue d = RES[nom]["perl"] - base bs = np.array([d[rng.integers(0, len(d), len(d))].mean() for _ in range(3000)]) lo, hi = np.percentile(bs, 2.5), np.percentile(bs, 97.5) verdict = "BAT LE CHAMPION" if hi < 0 else ("perd" if lo > 0 else "dans le bruit") print("%-26s delta %+.5f IC95 [%+.5f,%+.5f] %s" % (nom, d.mean(), lo, hi, verdict)) print("\n>>> CALIBRATION : lin_s4 et joint_linps DOIVENT finir derriere joint_cont.") print(">>> Si ce n'est pas le cas, la validation ne sert pas a selectionner un modele.") print("BENCH_LIN_DONE")