#!/usr/bin/env python3 """(A) VÉRIFIE que le cache d'hypothèses shona n'est pas tronqué (asymétrie omissions suspecte). (B) Construit une table de règles de SEGMENTATION lingala à partir du corpus TRAIN (la spec) : la segmentation est 22.5% du budget d'erreur lin. Pour chaque forme collée que nous produisons, on compare dans le TRAIN la fréquence de la forme COLLÉE vs SÉPARÉE. On ne garde que les cas tranchés (ratio >= RATIO et effectif suffisant) => règle sûre. (C) Applique et mesure sur devhard-lin. Aucune soumission. """ import json, os, re from collections import Counter import difflib, jiwer def comb(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 norm(w): return re.sub(r"[^\w'ɛɔ]", "", w.lower()) rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] D = json.load(open("/root/devhard_allhyps.json", encoding="utf-8")) # ---------- (A) sanity du cache shona ---------- print("=== (A) VERIFICATION DU CACHE D'HYPOTHESES ===") for lg, mdl in (("lin", "joint_cont_best"), ("sna", "sna_ps_best")): sub = [r for r in rows if r["lang"] == lg] refs = [r["text"] for r in sub] hyps = [D[mdl].get(r["id"], "") for r in sub] nr = sum(len(t.split()) for t in refs); nh = sum(len(t.split()) for t in hyps) m = comb(refs, hyps) print(" %s : mots refs %d | mots hyps %d (ratio %.2f) | combine %.4f" % (lg, nr, nh, nh / max(nr, 1), m[2])) if nh < 0.8 * nr: print(" ⚠️ HYPOTHESES TRONQUEES — ce cache n'est pas exploitable pour %s" % lg) # ---------- (B) table de règles de segmentation lin ---------- print("\n=== (B) REGLES DE SEGMENTATION (corpus train = la spec) ===") LG = "lin" sub = [r for r in rows if r["lang"] == LG] refs = [r["text"] for r in sub] hyps = [D["joint_cont_best"].get(r["id"], "") for r in sub] base = comb(refs, hyps) print("baseline lin : combine %.4f" % base[2]) # corpus train : fréquences des formes collées et des bigrammes uni = Counter(); big = Counter() for l in open("/root/devhard/train_%s_min.jsonl" % LG, encoding="utf-8"): w = [norm(x) for x in json.loads(l).get("text", "").split()] w = [x for x in w if x] uni.update(w) big.update(zip(w, w[1:])) # candidats : formes collées que NOUS produisons et qui se découpent en 2 mots connus hv = Counter() for h in hyps: hv.update(norm(x) for x in h.split()) hv.pop("", None) RATIO = float(os.environ.get("RATIO", "5")) MINC = int(os.environ.get("MINC", "10")) split_rules = {} for w, n in hv.items(): if len(w) < 5 or uni.get(w, 0) > 0 and uni[w] >= MINC: pass best = None for k in range(2, len(w) - 1): a, b = w[:k], w[k:] pa = big.get((a, b), 0) if pa >= MINC and uni.get(a, 0) >= MINC and uni.get(b, 0) >= MINC: if pa > (best[2] if best else 0): best = (a, b, pa) if best: joined = uni.get(w, 0) if best[2] >= RATIO * max(joined, 1): split_rules[w] = (best[0] + " " + best[1], joined, best[2]) print("regles de SEPARATION retenues (ratio>=%g, effectif>=%d) : %d" % (RATIO, MINC, len(split_rules))) for w, (s, j, p) in sorted(split_rules.items(), key=lambda x: -x[1][2])[:15]: print(" %-16s -> %-20s (train : colle %d | separe %d)" % (w, s, j, p)) # règles de FUSION : bigrammes que nous produisons mais que le train colle join_rules = {} for h in hyps: w = [norm(x) for x in h.split()] for a, b in zip(w, w[1:]): if not a or not b: continue j = a + b if uni.get(j, 0) >= MINC and uni[j] >= RATIO * max(big.get((a, b), 0), 1): join_rules[(a, b)] = (j, big.get((a, b), 0), uni[j]) print("\nregles de FUSION retenues : %d" % len(join_rules)) for (a, b), (j, s, c) in sorted(join_rules.items(), key=lambda x: -x[1][2])[:10]: print(" %-20s -> %-16s (train : separe %d | colle %d)" % (a + " " + b, j, s, c)) def apply_rules(t, do_split=True, do_join=True): out = [] for tok in t.split(): k = norm(tok) if do_split and k in split_rules: rep = split_rules[k][0] # conserver la casse initiale et la ponctuation finale tail = tok[len(tok.rstrip(".,!?")):] if tok[:1].isupper(): rep = rep[:1].upper() + rep[1:] out.append(rep + tail) else: out.append(tok) if do_join: res = [] i = 0 toks = out while i < len(toks): if i + 1 < len(toks): a, b = norm(toks[i]), norm(toks[i + 1]) if (a, b) in join_rules: tail = toks[i + 1][len(toks[i + 1].rstrip(".,!?")):] m = join_rules[(a, b)][0] if toks[i][:1].isupper(): m = m[:1].upper() + m[1:] res.append(m + tail); i += 2; continue res.append(toks[i]); i += 1 out = res return " ".join(out) print("\n=== (C) EFFET SUR devhard-lin ===") for lbl, ds, dj in (("separation seule", True, False), ("fusion seule", False, True), ("les deux", True, True)): o = [apply_rules(h, ds, dj) for h in hyps] m = comb(refs, o) nch = sum(1 for a, b in zip(hyps, o) if a != b) print(" %-18s WER %.4f CER %.4f combine %.4f (%+.4f) %d clips modifies%s" % (lbl, m[0], m[1], m[2], m[2] - base[2], nch, " <-- GAIN" if m[2] < base[2] else "")) print("\nSEG_RULES_DONE")