File size: 4,850 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
"""Fouille d'erreurs sur les hypotheses champion (banc 4235 clips, refs connues).
Chaque piste n'est retenue que si son gain NET sur la validation complete est positif :
 A. normalisation apostrophes/typographie   B. paires de substitution frequentes
 C. majuscules systematiques (noms propres) D. chiffres vs lettres  E. hygiene CSV."""
import csv
import json
import unicodedata
from collections import Counter

import jiwer

MAP = json.load(open("/scratch/rehearsal/mapping.json", encoding="utf-8"))
HYPS = {r["ID"]: r["Target"] for r in csv.DictReader(open("/root/subs/sub_rehearsal_final.csv", encoding="utf-8"))}
LANGS = ["lin", "lug", "sna"]


def norm(t):
    return " ".join(str(t).replace("|", " ").split())


def pairs_for(lang):
    rs, hs = [], []
    for anon, info in MAP.items():
        if info["true_lang"] == lang and info["text"].strip():
            rs.append(norm(info["text"]))
            hs.append(norm(HYPS.get(anon, "")))
    return rs, hs


def sc(rs, hs):
    return 0.5 * jiwer.wer(rs, hs) + 0.5 * jiwer.cer(rs, hs)


def apply_map(hs, table):
    out = []
    for h in hs:
        out.append(" ".join(table.get(w, w) for w in h.split()))
    return out


print("=== E. hygiene CSV ===", flush=True)
for name in ("sub_public.csv", "sub_phase2.csv", "sub_rehearsal_final.csv"):
    rows = list(csv.DictReader(open(f"/root/subs/{name}", encoding="utf-8")))
    bad = sum(1 for r in rows if "�" in r["Target"] or "Ã" in r["Target"])
    empt = sum(1 for r in rows if not r["Target"].strip())
    lens = sorted(len(r["Target"].split()) for r in rows)
    print(f"{name}: {len(rows)} lignes | mojibake {bad} | vides {empt} | mots p50/p95 {lens[len(lens)//2]}/{lens[int(len(lens)*0.95)]}", flush=True)

CHAR_VARIANTS = [("’", "'"), ("‘", "'"), ("`", "'"), ("“", '"'), ("”", '"')]
for lang in LANGS:
    rs, hs = pairs_for(lang)
    base = sc(rs, hs)
    print(f"\n===== {lang} | baseline banc {base:.4f} ({len(rs)} clips) =====", flush=True)

    print("--- A. typographie ---", flush=True)
    rjoin, hjoin = " ".join(rs), " ".join(hs)
    for ch, repl in CHAR_VARIANTS:
        cr, chh = rjoin.count(ch), hjoin.count(ch)
        if cr or chh:
            print(f"  car {ch!r}: refs {cr} vs hyps {chh}", flush=True)
    hs2 = [h.translate(str.maketrans({"’": "'", "‘": "'", "`": "'"})) for h in hs]
    s2 = sc(rs, hs2)
    if s2 < base:
        print(f"  normalisation apostrophes: {base:.4f} -> {s2:.4f}  <== MIEUX", flush=True)
    else:
        print(f"  normalisation apostrophes: {s2 - base:+.5f} (neutre/negatif)", flush=True)

    print("--- B. paires de substitution ---", flush=True)
    sub = Counter()
    for r, h in zip(rs, hs):
        rw, hw = r.split(), h.split()
        ops = jiwer.process_words([r], [h])
        for al in ops.alignments[0]:
            if al.type == "substitute":
                for i, j in zip(range(al.ref_start_idx, al.ref_end_idx), range(al.hyp_start_idx, al.hyp_end_idx)):
                    sub[(hw[j], rw[i])] += 1
    cands = [(h, r, n) for (h, r), n in sub.most_common(60) if n >= 8 and h.lower() != r.lower()]
    print(f"  {len(cands)} paires candidates (>=8 occ.)", flush=True)
    kept = {}
    cur = list(hs)
    curs = base
    for h, r, n in cands:
        trial = apply_map(cur, {h: r})
        st = sc(rs, trial)
        if st < curs - 1e-6:
            kept[h] = r
            cur, curs = trial, st
    print(f"  paires retenues: {len(kept)} | banc {base:.4f} -> {curs:.4f} ({curs - base:+.4f})", flush=True)
    if kept:
        print("  ex:", dict(list(kept.items())[:8]), flush=True)

    print("--- C. majuscules systematiques ---", flush=True)
    wc, wcap = Counter(), Counter()
    for l in open(f"/scratch/prep/manifests/waxal_{lang}_train.jsonl", encoding="utf-8"):
        for w in norm(json.loads(l)["text"]).split():
            core = w.strip(".,;:!?\"'")
            if len(core) >= 3:
                wc[core.lower()] += 1
                if core[:1].isupper():
                    wcap[core.lower()] += 1
    always_cap = {w for w in wc if wc[w] >= 5 and wcap[w] / wc[w] >= 0.95}
    table = {}
    for h in hs:
        for w in h.split():
            core = w.strip(".,;:!?\"'")
            if core.lower() in always_cap and core[:1].islower():
                table[w] = w.replace(core, core[:1].upper() + core[1:], 1)
    hs3 = apply_map(hs, table)
    s3 = sc(rs, hs3)
    print(f"  {len(always_cap)} mots toujours-capitalises train | {len(table)} formes corrigees | {base:.4f} -> {s3:.4f} ({s3 - base:+.4f})", flush=True)

    print("--- D. chiffres ---", flush=True)
    rdig = sum(any(c.isdigit() for c in r) for r in rs)
    hdig = sum(any(c.isdigit() for c in h) for h in hs)
    print(f"  refs avec chiffres: {rdig}/{len(rs)} | hyps avec chiffres: {hdig}", flush=True)

print("\nMINE_DONE", flush=True)