| |
| """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") |
|
|