File size: 5,411 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
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
#!/usr/bin/env python3
"""Taxonomie des erreurs de nos modeles sur la validation (CPU, alignements jiwer).
Complete l'analyse profonde : substitutions par categorie + frontieres de mots + clips pathologiques."""
import csv
import json
import unicodedata

import jiwer
import numpy as np

LANGS = ["lin", "lug", "sna"]


def lower_np(t):
    return unicodedata.normalize("NFC", t).lower()


def strip_punct(w):
    return "".join(c for c in w if c.isalnum())


def strip_diac(w):
    w = w.lower().replace("ɛ", "e").replace("ɔ", "o")
    return "".join(c for c in unicodedata.normalize("NFD", w) if not unicodedata.combining(c))


def classify_sub(rw, hw):
    if rw == hw:
        return "identique"
    if rw.lower() == hw.lower():
        return "casse_seule"
    if strip_punct(rw) == strip_punct(hw) and strip_punct(rw):
        return "ponct_seule"
    if strip_punct(rw.lower()) == strip_punct(hw.lower()) and strip_punct(rw):
        return "casse+ponct"
    if strip_diac(strip_punct(rw)) == strip_diac(strip_punct(hw)) and strip_punct(rw):
        return "diacritique"
    return "lexicale"


def main():
    mapping = json.load(open("/scratch/rehearsal/mapping.json"))
    hyps = {r["ID"]: r["Target"] for r in csv.DictReader(open("/scratch/rehearsal/sub_rehearsal.csv", encoding="utf-8"))}
    report = {}

    for lang in LANGS:
        clips = [(a, i["text"], hyps.get(a, "")) for a, i in mapping.items()
                 if i["true_lang"] == lang and i["text"].strip()]
        cats = {"casse_seule": 0, "ponct_seule": 0, "casse+ponct": 0, "diacritique": 0, "lexicale": 0}
        n_sub = n_ins = n_del = 0
        clip_wer = []
        # frontieres de mots : CER avec vs sans espaces
        ref_ns, hyp_ns, ref_sp, hyp_sp = [], [], [], []
        patho = []
        for a, ref, hyp in clips:
            out = jiwer.process_words(ref, hyp)
            for chunk in out.alignments[0]:
                if chunk.type == "substitute":
                    for k in range(chunk.ref_end_idx - chunk.ref_start_idx):
                        rw = out.references[0][chunk.ref_start_idx + k]
                        hw = out.hypotheses[0][chunk.hyp_start_idx + k] if chunk.hyp_start_idx + k < chunk.hyp_end_idx else ""
                        cats[classify_sub(rw, hw)] += 1
                        n_sub += 1
                elif chunk.type == "insert":
                    n_ins += chunk.hyp_end_idx - chunk.hyp_start_idx
                elif chunk.type == "delete":
                    n_del += chunk.ref_end_idx - chunk.ref_start_idx
            w = jiwer.wer(ref, hyp)
            clip_wer.append((w, a, ref, hyp))
            ref_sp.append(ref.lower()); hyp_sp.append(hyp.lower())
            ref_ns.append(ref.lower().replace(" ", "")); hyp_ns.append(hyp.lower().replace(" ", ""))
            if w >= 1.0:
                patho.append((a, ref, hyp, len(ref.split())))

        tot_err = n_sub + n_ins + n_del
        conv = cats["casse_seule"] + cats["ponct_seule"] + cats["casse+ponct"] + cats["diacritique"]
        cer_sp = jiwer.cer(ref_sp, hyp_sp)
        cer_ns = jiwer.cer(ref_ns, hyp_ns)
        wers = np.array([w for w, *_ in clip_wer])
        clip_wer.sort(reverse=True)
        top5_err = sum(min(w, 3.0) * len(r.split()) for w, a, r, h in clip_wer[:max(1, len(clips)//20)])
        all_err = sum(min(w, 3.0) * len(r.split()) for w, a, r, h in clip_wer)

        report[lang] = {
            "n_clips": len(clips), "total_erreurs_mots": tot_err,
            "substitutions": n_sub, "insertions": n_ins, "deletions": n_del,
            "categories_subs": cats,
            "pct_convention_only_sur_total": round(100 * conv / max(tot_err, 1), 1),
            "pct_lexicale_sur_subs": round(100 * cats["lexicale"] / max(n_sub, 1), 1),
            "cer_avec_espaces": round(cer_sp, 4),
            "cer_sans_espaces": round(cer_ns, 4),
            "gain_potentiel_frontieres_cer": round(cer_sp - cer_ns, 4),
            "wer_par_clip": {
                "=0%": int((wers == 0).sum()), "0-20%": int(((wers > 0) & (wers <= 0.2)).sum()),
                "20-50%": int(((wers > 0.2) & (wers <= 0.5)).sum()),
                "50-100%": int(((wers > 0.5) & (wers < 1.0)).sum()), ">=100%": int((wers >= 1.0).sum()),
            },
            "pct_erreurs_dans_5pct_pires_clips": round(100 * top5_err / max(all_err, 1), 1),
            "n_patho_wer100": len(patho),
        }
        report[lang]["exemples_patho"] = [
            {"ref": r[:90], "hyp": h[:90], "n_mots": nm} for _, r, h, nm in
            sorted(patho, key=lambda x: -x[3])[:6]
        ]
        print(f"=== {lang} ===", flush=True)
        print(f"  erreurs mots: {tot_err} (sub {n_sub}/ins {n_ins}/del {n_del})", flush=True)
        print(f"  categories subs: {cats}", flush=True)
        print(f"  convention-only sur total erreurs: {report[lang]['pct_convention_only_sur_total']}%", flush=True)
        print(f"  CER avec espaces {cer_sp:.4f} vs sans espaces {cer_ns:.4f} -> frontieres = {cer_sp-cer_ns:.4f}", flush=True)
        print(f"  WER/clip: {report[lang]['wer_par_clip']}", flush=True)
        print(f"  5% pires clips portent {report[lang]['pct_erreurs_dans_5pct_pires_clips']}% des erreurs; {len(patho)} clips WER>=100%", flush=True)

    json.dump(report, open("/root/models/error_taxonomy.json", "w"), ensure_ascii=False, indent=1)
    print("TAXONOMY_DONE", flush=True)


if __name__ == "__main__":
    main()