| |
| """RESCORING MULTI-MODELES sur le lingala. |
| |
| Point cle : pour du rescoring, un modele n'a PAS besoin du meme vocabulaire ni de la meme |
| frequence de trames que le decodeur — il doit seulement savoir evaluer log P(texte | ses logits). |
| On peut donc recruter des modeles ecartes comme decodeurs (MMS, monolingues) comme rescoreurs. |
| |
| score(h) = ac_cont(h) + lm_kenlm(h) + somme_i w_i * ac_modele_i(h) |
| Recherche des poids par montee de coordonnees sur devhard-lin. |
| """ |
| 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"), |
| ("jbest", "/root/models/joint_best"), |
| ("lin_s4", R + "/lin_s4_best"), |
| ("mmsjoint", "/root/models/mmsjoint_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): |
| """Encode en restant dans le vocabulaire du modele : les caracteres absents sont retires. |
| Permet de recruter des modeles sans casse/ponctuation (ils jugent alors le contenu seul).""" |
| 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: |
| low = text.lower().replace(" ", delim) |
| keep = "".join(c for c in low if c in v) |
| ids = [v[c] for c in keep] |
| return [i for i in ids if i != 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 (meilleur 1-best connu): %.4f" % REF, flush=True) |
|
|
| |
| cands, AC1, LM = [], [], [] |
| 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)) |
| print("taille moyenne du pool: %.1f hypotheses" % np.mean([len(c) for c in cands]), flush=True) |
|
|
| |
| SC = {} |
| chars = set("".join(refs)) |
| for tag, mdl in RESCORERS: |
| if not os.path.isdir(mdl): |
| print("%-9s ABSENT (%s)" % (tag, mdl), flush=True) |
| continue |
| try: |
| proc, LG = compute_logits(mdl, sub) |
| except Exception as e: |
| print("%-9s ERREUR %s" % (tag, str(e)[:60]), flush=True) |
| continue |
| t2 = proc.tokenizer |
| vv = set(t2.get_vocab()) |
| miss = sorted(c for c in chars if c not in vv and c != " ") |
| s = [] |
| for i, c in enumerate(cands): |
| s.append(np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in c])) |
| SC[tag] = s |
| print("%-9s |V|=%d, caract. refs absents=%d -> scores calcules" |
| % (tag, len(vv), len(miss)), flush=True) |
|
|
| def evaluate(weights): |
| hyps = [] |
| for i in range(len(cands)): |
| tot = AC1[i] + LM[i] |
| for tag, w in weights.items(): |
| if w: |
| tot = tot + w * SC[tag][i] |
| hyps.append(cc(cands[i][int(np.argmax(tot))], greedy[i])) |
| return comb(refs, hyps)[2] |
|
|
| print("\n--- rescoreurs pris un par un ---", flush=True) |
| solo = {} |
| for tag in SC: |
| bb = (9.0, 0.0) |
| for w in (0.3, 0.5, 1.0, 1.5, 2.5): |
| m = evaluate({tag: w}) |
| if m < bb[0]: |
| bb = (m, w) |
| solo[tag] = bb |
| print(" %-9s meilleur %.4f (w=%.1f) vs ref %.4f : %+.4f" |
| % (tag, bb[0], bb[1], REF, bb[0] - REF), flush=True) |
|
|
| print("\n--- montee de coordonnees (combinaison) ---", flush=True) |
| W = {t: 0.0 for t in SC} |
| order = sorted(solo, key=lambda t: solo[t][0]) |
| for t in order: |
| W[t] = solo[t][1] |
| cur = evaluate(W) |
| print(" depart (chacun a son optimum solo): %.4f %s" % (cur, W), flush=True) |
| for it in range(3): |
| improved = False |
| for t in order: |
| base_w = W[t] |
| for w in (0.0, 0.3, 0.5, 1.0, 1.5, 2.5): |
| W[t] = w |
| m = evaluate(W) |
| if m < cur - 1e-6: |
| cur, base_w, improved = m, w, True |
| W[t] = base_w |
| print(" passe %d: %.4f %s" % (it + 1, cur, {k: v for k, v in W.items() if v}), flush=True) |
| if not improved: |
| break |
|
|
| print("\nBEST_MULTI %.4f poids=%s (ref %.4f, gain %+.4f)" |
| % (cur, {k: v for k, v in W.items() if v}, REF, cur - REF), flush=True) |
| json.dump({"combine": cur, "weights": W, "ref": REF, "solo": {k: list(v) for k, v in solo.items()}}, |
| open("/root/multi_rescore.json", "w")) |
| print("MULTI_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|