File size: 3,208 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 | #!/usr/bin/env python3
"""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]))
# ou est la casse ? initiales de phrase vs mots internes
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")
|