waxal2026-backup / code /bench_lin.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
6.57 kB
#!/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")