#!/usr/bin/env python3 """RESCORING v3 : ajoute (a) les monolingues lin comme rescoreurs (bons chemins), (b) un TERME DE LONGUEUR — les scores etant des sommes de log-probs, ils penalisent structurellement les hypotheses longues, ce qui peut expliquer une part de l'ecart a l'oracle. score(h) = ac_cont(h) + lm_kenlm(h) + somme_i w_i*ac_i(h) + gamma*nb_mots(h) """ import json import os import pickle import jiwer import numpy as np import soundfile as sf import torch from multiprocessing import Pool from pyctcdecode import build_ctcdecoder from transformers import AutoModelForCTC, AutoProcessor M1 = "/root/models/joint_cont_best" ARPA = "/scratch/lm/lin_5g.arpa" NBEST = 10 R = "/scratch/restore" RESCORERS = [ ("cont2", "/root/models/joint_cont2_best"), ("lin_s4", "/root/models/lin_s4_best"), ("lin_r2", R + "/lin_r2_best"), ] 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 encode_for(tok, text): v = tok.get_vocab() delim = getattr(tok, "word_delimiter_token", "|") s = text.replace(" ", delim) keep = "".join(c for c in s if c in v) if not keep: keep = "".join(c for c in text.lower().replace(" ", delim) if c in v) return [v[c] for c in keep if v[c] != tok.pad_token_id] def ctc_score(logp, ids, blank): T = logp.shape[0] if not ids or len(ids) > T: return -1e9 lp = torch.from_numpy(logp).unsqueeze(1) loss = torch.nn.functional.ctc_loss( lp, torch.tensor(ids).unsqueeze(0), torch.tensor([T]), torch.tensor([len(ids)]), blank=blank, reduction="sum", zero_infinity=True) return -float(loss) def compute_logits(model_dir, rows): proc = AutoProcessor.from_pretrained(model_dir) m = AutoModelForCTC.from_pretrained(model_dir, dtype=torch.float32).cuda().eval() out = [] with torch.inference_mode(): for i in range(0, len(rows), 4): b = rows[i:i + 4] au = [sf.read(r["audio"], dtype="float32")[0] for r in b] x = proc(au, sampling_rate=16000, return_tensors="pt", padding=True) x = {k: v.cuda() for k, v in x.items()} lg = m(**x).logits.log_softmax(-1).float().cpu().numpy() for j in range(len(b)): out.append(lg[j]) del m torch.cuda.empty_cache() return proc, out 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] L1 = pickle.load(open("/scratch/lm/logits_lin.pkl", "rb")) tok = AutoProcessor.from_pretrained(M1).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 L1] def cc(h, g): return (g[:1] + h[1:]) if (h and g) else h dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=0.5, beta=0.5, lm_score_boundary=False) with Pool(8) as p: allbeams = dec.decode_beams_batch(p, L1, beam_width=64) with Pool(8) as p: db = [" ".join(x.split()) for x in dec.decode_batch(p, L1, beam_width=64)] _, _, REF = comb(refs, [cc(h, g) for h, g in zip(db, greedy)]) print("reference: %.4f" % REF, flush=True) cands, AC1, LM, NW = [], [], [], [] for i, bs in enumerate(allbeams): c = [" ".join(b[0].split()) for b in bs[:NBEST]] a = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]] l = [((b[4] - b[3]) if len(b) > 4 else 0.0) for b in bs[:NBEST]] if db[i] not in c: c.append(db[i]) a.append(ctc_score(L1[i], encode_for(tok, db[i]), tok.pad_token_id)) l.append(float(np.mean(l)) if l else 0.0) cands.append(c) AC1.append(np.array(a)) LM.append(np.array(l)) NW.append(np.array([float(len(x.split())) for x in c])) SC = {} for tag, mdl in RESCORERS: if not os.path.isdir(mdl): print("%-8s ABSENT %s" % (tag, mdl), flush=True) continue proc, LG = compute_logits(mdl, sub) t2 = proc.tokenizer SC[tag] = [np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in cands[i]]) for i in range(len(cands))] print("%-8s scores OK (|V|=%d)" % (tag, len(t2.get_vocab())), flush=True) def evaluate(W, gamma=0.0): hyps = [] for i in range(len(cands)): tot = AC1[i] + LM[i] + gamma * NW[i] for t, w in W.items(): if w: tot = tot + w * SC[t][i] hyps.append(cc(cands[i][int(np.argmax(tot))], greedy[i])) return comb(refs, hyps)[2] print("\n--- (a) rescoreurs solo ---", flush=True) for tag in SC: bb = (9.0, 0.0) for w in (0.3, 0.5, 1.0, 1.5, 2.5, 4.0): m = evaluate({tag: w}) if m < bb[0]: bb = (m, w) print(" %-8s %.4f (w=%.1f) %+.4f" % (tag, bb[0], bb[1], bb[0] - REF), flush=True) print("\n--- (b) terme de longueur seul (gamma), sans rescoreur ---", flush=True) for gm in (-2.0, -1.0, 0.0, 1.0, 2.0, 4.0, 8.0): print(" gamma=%+5.1f : %.4f (%+.4f)" % (gm, evaluate({}, gm), evaluate({}, gm) - REF), flush=True) print("\n--- (c) cont2 + longueur ---", flush=True) best = (9.0, None, None) for w in (1.0, 1.5, 2.5): for gm in (0.0, 1.0, 2.0, 4.0, 8.0): m = evaluate({"cont2": w}, gm) if m < best[0]: best = (m, w, gm) print(" cont2=%.1f gamma=%+5.1f : %.4f (%+.4f)" % (w, gm, m, m - REF), flush=True) print("\nBEST_V3 %.4f cont2=%s gamma=%s (ref %.4f, gain %+.4f)" % (best[0], best[1], best[2], REF, best[0] - REF), flush=True) json.dump({"combine": best[0], "cont2": best[1], "gamma": best[2], "ref": REF}, open("/root/multi_v3.json", "w")) print("V3_DONE", flush=True) if __name__ == "__main__": main()