#!/usr/bin/env python3 """Sweep final du decodage lingala sur devhard (modele FIXE joint_cont, logits en cache). Compare : LM base (13960) vs LM enrichi SANS FUITE (15800, o5 et o6), grille alpha x beta, puis parametres avances pyctcdecode (unk_score_offset, lm_score_boundary) au meilleur point. Toutes les hypotheses recoivent la casse du 1er caractere du greedy (gain deja valide). """ import json import pickle import jiwer from multiprocessing import Pool from pyctcdecode import build_ctcdecoder from transformers import AutoProcessor MDL = "/root/models/joint_cont_best" LOGP = "/scratch/lm/logits_lin.pkl" LMS = [ ("base_5g", "/scratch/lm/lin_5g.arpa"), ("noleak_5g", "/scratch/lm/lin_noleak_5g.arpa"), ("noleak_6g", "/scratch/lm/lin_noleak_6g.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 main(): 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] log = pickle.load(open(LOGP, "rb")) tok = AutoProcessor.from_pretrained(MDL).tokenizer v = tok.get_vocab() lab = [None] * len(v) for t, i in v.items(): lab[i] = t lab[tok.word_delimiter_token_id] = " " lab[tok.unk_token_id] = "⁇" lab[tok.pad_token_id] = "" greedy = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in log] _, _, g0 = comb(refs, greedy) print("greedy: %.4f" % g0, flush=True) def run(arpa, alpha, beta, bw=64, **kw): dec = build_ctcdecoder(lab, kenlm_model_path=arpa, alpha=alpha, beta=beta, **kw) with Pool(8) as p: h = dec.decode_batch(p, log, beam_width=bw) h = [" ".join(x.split()) for x in h] h = [(g[:1] + x[1:] if x and g else x) for x, g in zip(h, greedy)] return comb(refs, h) best = (9.0, None) print("--- LM x alpha x beta ---", flush=True) for tag, arpa in LMS: for alpha in (0.4, 0.5, 0.6): cells = [] for beta in (0.0, 0.5, 1.0): _, _, m = run(arpa, alpha, beta) cells.append("b%.1f=%.4f" % (beta, m)) if m < best[0]: best = (m, (tag, arpa, alpha, beta)) print(" %-11s a=%.1f %s" % (tag, alpha, " ".join(cells)), flush=True) print(" >>> meilleur: %.4f %s" % (best[0], best[1][0::1][:1] + best[1][2:]), flush=True) tag, arpa, alpha, beta = best[1] print("--- parametres avances au meilleur point (%s a=%.1f b=%.1f) ---" % (tag, alpha, beta), flush=True) for uso in (-10.0, 0.0, 10.0): _, _, m = run(arpa, alpha, beta, unk_score_offset=uso) print(" unk_score_offset=%+.0f : %.4f%s" % (uso, m, " <<<" if m < best[0] else ""), flush=True) if m < best[0]: best = (m, (tag, arpa, alpha, beta, {"unk_score_offset": uso})) for lsb in (True, False): _, _, m = run(arpa, alpha, beta, lm_score_boundary=lsb) print(" lm_score_boundary=%s : %.4f%s" % (lsb, m, " <<<" if m < best[0] else ""), flush=True) if m < best[0]: best = (m, (tag, arpa, alpha, beta, {"lm_score_boundary": lsb})) print("BEST_FINAL %.4f cfg=%s (greedy %.4f, gain %+.4f)" % (best[0], best[1], g0, g0 - best[0]), flush=True) json.dump({"combine": best[0], "cfg": [str(x) for x in best[1]], "greedy": g0}, open("/root/sweep_final.json", "w")) print("SWEEP_FINAL_DONE", flush=True) if __name__ == "__main__": main()