| |
| """Compare sna_ps (tour 1) et sna_ps2 (tour 2) sur devhard-sna : 444 clips a |
| LOCUTEURS DISJOINTS du train. C'est le seul jeu ou l'apport des 295 h de |
| locuteurs NOUVEAUX peut se manifester -- la validation officielle partage ses |
| locuteurs avec le train et est restee plate (0.1743 -> 0.1736). |
| |
| Pipeline identique a la soumission : beam CTC pur 256, N-best 50, rescorage |
| par sna_r2_best avec w=4.0. On mesure aussi le greedy, pour separer l'effet |
| du modele acoustique de l'effet du decodage. |
| """ |
| import json, os, sys |
| import numpy as np, torch |
| from multiprocessing import Pool |
| from pyctcdecode import build_ctcdecoder |
|
|
| sys.path.insert(0, "/root") |
| from gen_sna_rescore import compute_logits, ctc_score, encode_for, norm |
|
|
| DEV = "/root/devhard/devhard_sna.jsonl" |
| M1 = "/root/models/sna_ps_best" |
| M2 = "/scratch/runs/sna_ps2/best" |
| RESC = "/scratch/restore/sna_r2_best" |
| W, NBEST, BW = 4.0, 50, 256 |
|
|
|
|
| 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, lower=False): |
| """WER/CER agreges (somme des erreurs / somme des longueurs), + erreurs par clip.""" |
| we = wl = ce = cl = 0 |
| per = [] |
| for h, r in zip(hyps, refs): |
| if lower: |
| h, r = h.lower(), r.lower() |
| hw, rw = h.split(), r.split() |
| e1 = lev(hw, rw) |
| e2 = lev(h, r) |
| we += e1; wl += len(rw); ce += e2; cl += len(r) |
| per.append(e1 / max(len(rw), 1) * 0.5 + e2 / max(len(r), 1) * 0.5) |
| wer, cer = we / max(wl, 1), ce / max(cl, 1) |
| return wer, cer, 0.5 * wer + 0.5 * cer, np.array(per) |
|
|
|
|
| def decode(model_dir, files, LG, t2, tag): |
| proc, L1 = compute_logits(model_dir, 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] = "" |
| greedy = [norm(tok.decode(l.argmax(-1))) for l in L1] |
| print(" [%s] logits + greedy OK" % tag, flush=True) |
|
|
| dec = build_ctcdecoder(lab) |
| with Pool(8) as p: |
| allbeams = dec.decode_beams_batch(p, L1, beam_width=BW) |
| with Pool(8) as p: |
| db = [norm(x) for x in dec.decode_batch(p, L1, beam_width=BW)] |
| print(" [%s] beam %d OK" % (tag, BW), flush=True) |
|
|
| picked = [] |
| for i, bs in enumerate(allbeams): |
| c = [norm(b[0]) for b in bs[:NBEST]] |
| a = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]] |
| for extra in (db[i], greedy[i]): |
| if extra and extra not in c: |
| c.append(extra) |
| a.append(ctc_score(L1[i], encode_for(tok, extra), tok.pad_token_id)) |
| sc = np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in c]) |
| tot = np.array(a) + W * sc |
| picked.append(c[int(np.argmax(tot))] or greedy[i] or "a") |
| print(" [%s] rescorage OK" % tag, flush=True) |
| return greedy, picked |
|
|
|
|
| rows = [json.loads(l) for l in open(DEV, encoding="utf-8")] |
| files, refs = [], [] |
| for r in rows: |
| p = r["audio"] |
| if not os.path.exists(p): |
| p = "/root/devhard_audio/%s.flac" % r["id"] |
| if os.path.exists(p): |
| files.append(p); refs.append(r["text"]) |
| print("devhard-sna : %d clips utilisables sur %d" % (len(files), len(rows)), flush=True) |
| assert len(files) > 400, "trop d'audio manquant" |
|
|
| proc2, LG = compute_logits(RESC, files) |
| t2 = proc2.tokenizer |
| print("rescoreur %s OK" % RESC, flush=True) |
|
|
| g1, p1 = decode(M1, files, LG, t2, "tour1") |
| g2, p2 = decode(M2, files, LG, t2, "tour2") |
|
|
| print("\n================ devhard-sna, %d clips, LOCUTEURS DISJOINTS ================" % len(files)) |
| print("%-34s %8s %8s %9s" % ("", "WER", "CER", "combine")) |
| res = {} |
| for nom, h in (("tour1 greedy", g1), ("tour1 pipeline complet", p1), |
| ("tour2 greedy", g2), ("tour2 pipeline complet", p2)): |
| w, c, k, per = score(h, refs) |
| res[nom] = per |
| print("%-34s %8.4f %8.4f %9.4f" % (nom, w, c, k)) |
|
|
| print("\n--- insensible a la casse (la casse compte-t-elle dans l'ecart ?) ---") |
| for nom, h in (("tour1 pipeline complet", p1), ("tour2 pipeline complet", p2)): |
| w, c, k, _ = score(h, refs, lower=True) |
| print("%-34s %8.4f %8.4f %9.4f" % (nom, w, c, k)) |
|
|
| d = res["tour2 pipeline complet"] - res["tour1 pipeline complet"] |
| mieux = int((d < -1e-9).sum()); pire = int((d > 1e-9).sum()); egal = int((abs(d) <= 1e-9).sum()) |
| print("\n--- par clip (pipeline complet) ---") |
| print("tour2 MEILLEUR sur %d clips | PIRE sur %d | identique sur %d" % (mieux, pire, egal)) |
| print("delta moyen (tour2 - tour1) : %+.5f (negatif = tour2 gagne)" % d.mean()) |
| rng = np.random.default_rng(0) |
| bs = np.array([d[rng.integers(0, len(d), len(d))].mean() for _ in range(4000)]) |
| print("bootstrap 95%% : [%+.5f , %+.5f] | P(tour2 gagne) = %.1f%%" |
| % (np.percentile(bs, 2.5), np.percentile(bs, 97.5), (bs < 0).mean() * 100)) |
| print("DEVHARD_PS2_DONE") |
|
|