|
|
| """RÈGLES DE CASSE & PONCTUATION FINALE, apprises sur le corpus TRAIN (la spec).
|
| Écarts systématiques mesurés :
|
| LIN : majuscules internes refs 129 vs nous 9 -> on en met 14x trop PEU (noms propres)
|
| SNA : majuscules internes refs 574 vs nous 635 -> on en met trop
|
| SNA : ponctuation finale refs 411/433 vs nous 431/433 -> on en met trop
|
| Casse+ponctuation = 8.5% du budget lin et 24.2% du budget sna.
|
| Règle : pour chaque mot, le train dit-il majuscule ou minuscule en position INTERNE ?
|
| On ne corrige que les cas tranchés (ratio >= R, effectif >= N).
|
| """
|
| import json, os, pickle, re
|
| from collections import Counter
|
| import jiwer
|
| from transformers import AutoProcessor
|
|
|
| R = float(os.environ.get("R", "4"))
|
| N = int(os.environ.get("N", "5"))
|
|
|
|
|
| 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 key(w):
|
| return re.sub(r"[^\w'ɛɔ]", "", w).lower()
|
|
|
|
|
| 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=== %s : baseline combine %.4f ===" % (lg.upper(), base[2]))
|
|
|
|
|
| up = Counter(); lo = Counter(); endp = 0; ntot = 0
|
| for l in open("/root/devhard/train_%s_min.jsonl" % lg, encoding="utf-8"):
|
| t = json.loads(l).get("text", "").strip()
|
| if not t:
|
| continue
|
| ntot += 1
|
| if t.rstrip()[-1:] in ".!?":
|
| endp += 1
|
| ws = t.split()
|
| for i, w in enumerate(ws):
|
| if i == 0 or ws[i - 1].rstrip()[-1:] in ".!?":
|
| continue
|
| k = key(w)
|
| if not k:
|
| continue
|
| (up if w[:1].isupper() else lo)[k] += 1
|
| print(" train : %.1f%% des phrases finissent par une ponctuation" % (100.0 * endp / max(ntot, 1)))
|
|
|
| force_up = {k for k in up if up[k] >= N and up[k] >= R * lo.get(k, 0)}
|
| force_lo = {k for k in lo if lo[k] >= N and lo[k] >= R * up.get(k, 0)}
|
| print(" mots TOUJOURS majuscule en interne : %d | TOUJOURS minuscule : %d"
|
| % (len(force_up), len(force_lo)))
|
|
|
| def fix_case(t):
|
| ws = t.split()
|
| out = [ws[0]] if ws else []
|
| for i, w in enumerate(ws[1:], 1):
|
| if ws[i - 1].rstrip()[-1:] in ".!?":
|
| out.append(w); continue
|
| k = key(w)
|
| if k in force_up and not w[:1].isupper():
|
| w = w[:1].upper() + w[1:]
|
| elif k in force_lo and w[:1].isupper():
|
| w = w[:1].lower() + w[1:]
|
| out.append(w)
|
| return " ".join(out)
|
|
|
| def fix_endp(t, keep):
|
| """si le train met rarement une ponctuation finale, on la retire"""
|
| if keep:
|
| return t
|
| return t.rstrip(".!?").rstrip() if t.rstrip()[-1:] in ".!?" else t
|
|
|
| keep_final = (endp / max(ntot, 1)) >= 0.5
|
| for lbl, fn in (("casse seule", lambda t: fix_case(t)),
|
| ("ponct finale seule", lambda t: fix_endp(t, keep_final)),
|
| ("les deux", lambda t: fix_endp(fix_case(t), keep_final))):
|
| o = [fn(h) for h in hyps]
|
| m = sc(refs, o)
|
| nch = sum(1 for a, b in zip(hyps, o) if a != b)
|
| print(" %-20s WER %.4f CER %.4f combine %.4f (%+.4f) %d clips%s"
|
| % (lbl, m[0], m[1], m[2], m[2] - base[2], nch,
|
| " <-- GAIN" if m[2] < base[2] - 0.0005 else ""))
|
| print("\nCASE_RULES_DONE")
|
|
|