| |
| """INSERTION DE VIRGULES guidée par notre KenLM in-domain (texte WAXAL train, virgules incluses). |
| Déficit mesuré : nos sorties lin ont 0.29 virgule/1000 car. contre 5.22 dans les références |
| => coût +0.0144 sur le combine lin (mesuré). La tentative §5 avait échoué avec un mBERT |
| multilingue qui RÉÉCRIVAIT tout (et détruisait les points corrects à 94%). |
| Ici : on n'AJOUTE que des virgules, jamais on ne touche au reste, et le juge est le LM |
| entraîné sur le corpus qui définit la convention. |
| |
| Pour chaque position inter-mots : score LM de "... wi wi+1 ..." vs "... wi, wi+1 ...". |
| On insère si le gain dépasse un seuil (balayé), avec au plus MAXINS virgules par phrase. |
| Gate : devhard-lin. |
| """ |
| import json, math, os, re, sys |
| import jiwer, kenlm |
|
|
| ARPA = os.environ.get("ARPA", "/scratch/lm/lin_5g.arpa") |
| LM = kenlm.Model(ARPA) |
|
|
|
|
| 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 sc(text): |
| return LM.score(text, bos=True, eos=True) |
|
|
|
|
| def insert_commas(text, thr, maxins=6): |
| """Ajoute des virgules là où le LM y gagne le plus, de façon gloutonne.""" |
| w = text.split() |
| if len(w) < 3: |
| return text |
| cur = list(w) |
| n = 0 |
| while n < maxins: |
| base = sc(" ".join(cur)) |
| best = (0.0, -1) |
| for i in range(len(cur) - 1): |
| if cur[i].endswith(",") or cur[i].endswith(".") or cur[i].endswith("!"): |
| continue |
| cand = list(cur) |
| cand[i] = cand[i] + "," |
| g = sc(" ".join(cand)) - base |
| if g > best[0]: |
| best = (g, i) |
| if best[1] < 0 or best[0] < thr: |
| break |
| cur[best[1]] = cur[best[1]] + "," |
| n += 1 |
| return " ".join(cur) |
|
|
|
|
| def main(): |
| D = json.load(open("/root/devhard_allhyps.json", encoding="utf-8")) |
| rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")] |
| sub = [r for r in rows if r["lang"] == "lin"] |
| refs = [r["text"] for r in sub] |
| H = D["joint_cont_best"] |
| hyps = [H.get(r["id"], "") for r in sub] |
| base = comb(refs, hyps) |
| print("baseline lin (joint_cont greedy) : WER %.4f CER %.4f combine %.4f" % base, flush=True) |
| nref = sum(t.count(",") for t in refs) |
| nhyp = sum(t.count(",") for t in hyps) |
| print("virgules : refs %d | nos hyps %d" % (nref, nhyp), flush=True) |
|
|
| print("\nseuil | virgules ajoutees | WER CER combine delta") |
| best = (base[2], None) |
| for thr in (0.2, 0.5, 1.0, 1.5, 2.0, 3.0, 5.0): |
| out = [insert_commas(h, thr) for h in hyps] |
| added = sum(o.count(",") for o in out) - nhyp |
| m = comb(refs, out) |
| d = m[2] - base[2] |
| if m[2] < best[0]: |
| best = (m[2], thr) |
| print("%5.1f | %5d | %.4f %.4f %.4f %+.4f%s" |
| % (thr, added, m[0], m[1], m[2], d, " <-- MIEUX" if d < 0 else ""), flush=True) |
| print("\nMEILLEUR : combine %.4f (seuil %s) vs baseline %.4f" % (best[0], best[1], base[2]), flush=True) |
| json.dump({"baseline": base[2], "best": best[0], "thr": best[1]}, |
| open("/root/comma_kenlm.json", "w")) |
| print("COMMA_KENLM_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|