| |
| """SHONA v2 — pousser le levier qui a payé (+0.001127 au LB). |
| Constat v1 : seul un SPÉCIALISTE shona marche comme rescoreur (sna_r2 −0.0029 ; cont2 −0.0002 ; |
| cont +0.0011 = nuit). Or 4 autres modèles shona n'ont JAMAIS été testés. |
| Oracle 10-best = 0.0949 (marge −0.0332) : la marge est dans la SÉLECTION. |
| Ici : (1) tous les rescoreurs shona en solo, (2) N-best élargi 25, (3) combinaison des 2 meilleurs |
| (avec garde-fou : on ne retient la combinaison que si elle bat nettement le meilleur solo, |
| sinon sur-apprentissage sur 433 clips — leçon §4c). |
| """ |
| import json, os, pickle |
| import jiwer, numpy as np, soundfile as sf, torch |
| from multiprocessing import Pool |
| from pyctcdecode import build_ctcdecoder |
| from transformers import AutoModelForCTC, AutoProcessor |
|
|
| M1 = "/root/models/sna_ps_best" |
| R = "/scratch/restore" |
| NBEST = int(os.environ.get("NBEST", "25")) |
| AUD = "/root/devhard_audio" |
| CANDS_R = [("sna_r2", R + "/sna_r2_best"), ("sna_r", R + "/sna_r_best"), |
| ("sna_s1", R + "/sna_s1_best"), ("sna_s2", R + "/sna_s2_best"), |
| ("sna_ws", R + "/sna_ws_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] |
| return 0.5 * jiwer.wer(a, b) + 0.5 * jiwer.cer(a, b) |
|
|
|
|
| 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) |
| return -float(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)) |
|
|
|
|
| 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"] == "sna"] |
| for r in sub: |
| r["audio"] = os.path.join(AUD, os.path.basename(r["audio"])) |
| sub = [r for r in sub if os.path.exists(r["audio"])] |
| refs = [r["text"] for r in sub] |
| print("devhard-sna %d clips | NBEST=%d" % (len(sub), NBEST), flush=True) |
|
|
| CACHE = "/scratch/lm/logits_sna.pkl" |
| L1 = pickle.load(open(CACHE, "rb")) if os.path.exists(CACHE) else compute_logits(M1, sub)[1] |
| 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] |
| REF = comb(refs, greedy) |
| print("baseline greedy %.4f" % REF, flush=True) |
|
|
| dec = build_ctcdecoder(lab) |
| with Pool(8) as p: |
| allbeams = dec.decode_beams_batch(p, L1, beam_width=128) |
| with Pool(8) as p: |
| db = [" ".join(x.split()) for x in dec.decode_batch(p, L1, beam_width=128)] |
|
|
| cands, AC1, 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]] |
| for extra in (db[i], greedy[i]): |
| if extra and extra not in c: |
| c.append(extra) |
| a.append(ctc_score(L1[i], encode_for(tok, extra), tok.pad_token_id)) |
| cands.append(c); AC1.append(np.array(a)) |
| NW.append(np.array([float(len(x.split())) for x in c])) |
| orc = [min(cands[i], key=lambda h: comb([refs[i]], [h]) if refs[i].strip() else 0) |
| for i in range(len(cands))] |
| print("ORACLE %d-best %.4f (marge %+.4f)" % (NBEST, comb(refs, orc), comb(refs, orc) - REF), flush=True) |
|
|
| SC = {} |
| for tag, mdl in CANDS_R: |
| if not os.path.isdir(mdl): |
| print("%-8s ABSENT" % tag, 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 OK" % tag, flush=True) |
|
|
| def ev(W): |
| hyps = [] |
| for i in range(len(cands)): |
| tot = AC1[i].copy() |
| for t, w in W.items(): |
| if w: |
| tot = tot + w * SC[t][i] |
| hyps.append(cands[i][int(np.argmax(tot))]) |
| return comb(refs, hyps) |
|
|
| print("\n--- solo (ref %.4f) ---" % REF, flush=True) |
| solo = {} |
| for t in SC: |
| bb = (9.0, 0.0) |
| for w in (0.3, 0.5, 1.0, 1.5, 2.5, 4.0): |
| m = ev({t: w}) |
| if m < bb[0]: |
| bb = (m, w) |
| solo[t] = bb |
| print(" %-8s %.4f (w=%.1f) %+.4f" % (t, bb[0], bb[1], bb[0] - REF), flush=True) |
|
|
| ranked = sorted(solo, key=lambda t: solo[t][0]) |
| best_solo = solo[ranked[0]] |
| print("\n--- combinaison des 2 meilleurs (%s + %s) ---" % (ranked[0], ranked[1]), flush=True) |
| bc = (9.0, None, None) |
| for w1 in (0.5, 1.0, 1.5, 2.5): |
| for w2 in (0.0, 0.3, 0.5, 1.0, 1.5): |
| m = ev({ranked[0]: w1, ranked[1]: w2}) |
| if m < bc[0]: |
| bc = (m, w1, w2) |
| print(" best %.4f (%s=%.1f %s=%.1f) %+.4f vs solo" % (bc[0], ranked[0], bc[1], ranked[1], bc[2], bc[0] - best_solo[0]), flush=True) |
| keep_combo = bc[0] < best_solo[0] - 0.0015 |
| print("\nRETENU : %s" % ("COMBINAISON" if keep_combo else "SOLO %s w=%.1f" % (ranked[0], best_solo[1])), flush=True) |
| json.dump({"ref": REF, "solo": {k: list(v) for k, v in solo.items()}, |
| "combo": list(bc), "keep_combo": bool(keep_combo), "nbest": NBEST}, |
| open("/root/sna_v2.json", "w")) |
| print("SNA_V2_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|