| |
| """DÉCOMPOSITION DU BUDGET D'ERREUR pour les DEUX langues (le §5 ne l'avait fait que sur le lin). |
| Combien du WER/CER vient de : la CASSE seule ? la PONCTUATION seule ? le vrai CONTENU ? |
| Le scoreur WAXAL est BRUT (casse et ponctuation comptent), donc chaque composante est un |
| levier potentiel — mais seulement si elle est réparable. |
| Mesure : score tel quel, puis en neutralisant chaque composante des DEUX côtés (ref et hyp). |
| """ |
| import json, pickle, re |
| import jiwer |
| from transformers import AutoProcessor |
|
|
| PUNCT = re.compile(r"[^\w\s'ɛɔ]") |
|
|
|
|
| def sc(refs, hyps): |
| pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()] |
| a = [x for x, _ in pr]; b = [y for _, y in pr] |
| w = jiwer.wer(a, b); c = jiwer.cer(a, b) |
| return w, c, 0.5 * w + 0.5 * c |
|
|
|
|
| def variants(t, lower=False, nopunct=False): |
| if lower: |
| t = t.lower() |
| if nopunct: |
| t = PUNCT.sub("", t) |
| return " ".join(t.split()) |
|
|
|
|
| rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] |
|
|
| CFG = { |
| "lin": ("/root/models/joint_cont_best", "/scratch/lm/logits_lin.pkl"), |
| "sna": ("/root/models/sna_ps_best", "/scratch/lm/logits_sna.pkl"), |
| } |
| for lg, (mdl, lgt) in CFG.items(): |
| sub = [r for r in rows if r["lang"] == lg] |
| refs = [r["text"] for r in sub] |
| L1 = pickle.load(open(lgt, "rb")) |
| tok = AutoProcessor.from_pretrained(mdl).tokenizer |
| hyps = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in L1] |
| base = sc(refs, hyps) |
| print("\n" + "=" * 70) |
| print("### %s — %d clips (greedy) : WER %.4f CER %.4f COMBINE %.4f" |
| % (lg.upper(), len(sub), base[0], base[1], base[2])) |
| print("=" * 70) |
| rows_out = [] |
| for label, lower, nop in (("casse neutralisee", True, False), |
| ("ponctuation neutralisee", False, True), |
| ("casse + ponctuation", True, True)): |
| m = sc([variants(t, lower, nop) for t in refs], |
| [variants(t, lower, nop) for t in hyps]) |
| rows_out.append((label, m, base[2] - m[2])) |
| print(" %-26s WER %.4f CER %.4f combine %.4f part du budget : %.1f%%" |
| % (label, m[0], m[1], m[2], 100.0 * (base[2] - m[2]) / base[2])) |
| resid = rows_out[-1][1][2] |
| print(" %-26s combine %.4f = le VRAI contenu (%.1f%%)" |
| % ("-> residuel", resid, 100.0 * resid / base[2])) |
| |
| maj_ref = sum(1 for t in refs if t[:1].isupper()) |
| maj_hyp = sum(1 for t in hyps if t[:1].isupper()) |
| print(" majuscule initiale : refs %d/%d | nous %d/%d" % (maj_ref, len(refs), maj_hyp, len(hyps))) |
| inner_ref = sum(sum(1 for w in t.split()[1:] if w[:1].isupper()) for t in refs) |
| inner_hyp = sum(sum(1 for w in t.split()[1:] if w[:1].isupper()) for t in hyps) |
| print(" majuscules internes : refs %d | nous %d" % (inner_ref, inner_hyp)) |
| fin_ref = sum(1 for t in refs if t.rstrip()[-1:] in ".!?") |
| fin_hyp = sum(1 for t in hyps if t.rstrip()[-1:] in ".!?") |
| print(" ponctuation finale : refs %d/%d | nous %d/%d" % (fin_ref, len(refs), fin_hyp, len(hyps))) |
| print("\nERR_DECOMP_DONE") |
|
|