#!/usr/bin/env python3 """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")