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")