Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
3.69 kB
#!/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()