File size: 6,565 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 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | #!/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")
|