#!/usr/bin/env python3 """POOL INTER-MODELES sur le SHONA — transposition de ce qui vient de marcher sur le lingala (`sub_XC010.csv` = 0.762499646, +0.000126 sur le record). Un juge CLASSE, il ne PROPOSE pas : tous les candidats shona sortaient du beam de `sna_ps_best` seul. On ajoute ici les hypotheses d un DONNEUR (autre distribution). BEAUCOUP plus simple que le lingala : le shona decode en beam PUR (aucun KenLM, le LM degradait le shona). Le score du beam est donc du CTC acoustique seul => un candidat externe se score EXACTEMENT sur la meme echelle avec ctc_score(), sans formule de calibration ni offset. C est precisement la calibration qui avait casse le 1er essai lingala (-0.0091). Ici ce risque n existe pas. Mecanisme deja valide dans le pipeline du record : `gen_sna_rescore.py` ajoute deja le greedy et le decode-batch comme candidats supplementaires, notes par ctc_score. GARDE-FOU DE LONGUEUR conserve (mots ET caracteres) : sur le lingala, l absence de borne en CARACTERES suffisait a tout perdre (memes mots, formes 40 % plus longues). """ import csv import json import os import sys import numpy as np import torch from multiprocessing import Pool from pyctcdecode import build_ctcdecoder sys.path.insert(0, "/root") from gen_sna_rescore import compute_logits, ctc_score, encode_for, norm # noqa: E402 BASE = os.environ.get("BASE", "/root/sub_XC010.csv") # record 0.762499646 LANGF = os.environ.get("LANGF", "/root/test_lang.json") SNAM = os.environ.get("SNA_MODEL", "/root/models/sna_ps_best") J1 = os.environ.get("J1", "/root/models/sna_r2_best") J2 = os.environ.get("J2", "/scratch/runs/xlsrlong/checkpoint-5600") J3 = os.environ.get("J3", "/root/models/mmsjoint_best") DONOR = os.environ.get("DONOR", "/root/models/mmsjoint_best") W1 = float(os.environ.get("W1", "6.0")) W2 = float(os.environ.get("W2", "4.0")) W3 = float(os.environ.get("W3", "2.0")) MARGINS = [float(x) for x in os.environ.get("MARGINS", "1e9,10,5,0").split(",")] DON_NBEST = int(os.environ.get("DON_NBEST", "5")) RLO = float(os.environ.get("RATIO_LO", "0.85")) RHI = float(os.environ.get("RATIO_HI", "1.15")) GAMMA = float(os.environ.get("GAMMA", "0.0")) NBEST = int(os.environ.get("NBEST", "50")) BW = int(os.environ.get("BW", "256")) TAG = os.environ.get("TAG", "SX") AUD = os.environ.get("AUDIO_DIR", "/scratch/p2_16k") def main(): base = {r["ID"]: r["Target"] for r in csv.DictReader(open(BASE, encoding="utf-8"))} lang = json.load(open(LANGF)) sna_ids = [k for k in base if lang.get(k) == "sna"] print("base %d | sna %d | lin intouches %d" % (len(base), len(sna_ids), len(base) - len(sna_ids)), flush=True) files = [os.path.join(AUD, k + ".wav") for k in sna_ids] assert all(os.path.exists(f) for f in files), "audio manquant" proc, L1 = compute_logits(SNAM, files) tok = proc.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 = [norm(tok.decode(l.argmax(-1))) for l in L1] dec = build_ctcdecoder(lab) # beam PUR, aucun LM with Pool(8) as p: allbeams = dec.decode_beams_batch(p, L1, beam_width=BW) with Pool(8) as p: db = [norm(x) for x in dec.decode_batch(p, L1, beam_width=BW)] cands, AC1, NW, ISEXT = [], [], [], [] for i, bs in enumerate(allbeams): c = [norm(b[0]) 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(list(a)) ISEXT.append([False] * len(c)) print("n-best du modele principal construit", flush=True) # ---------- candidats DONNEURS ---------- dproc, DL = compute_logits(DONOR, files) dtok = dproc.tokenizer dv = dtok.get_vocab() dlab = [None] * len(dv) for t, i in dv.items(): dlab[i] = t dlab[dtok.word_delimiter_token_id] = " " dlab[dtok.unk_token_id] = "⁇" dlab[dtok.pad_token_id] = "" ddec = build_ctcdecoder(dlab) with Pool(8) as p: dbeams = ddec.decode_beams_batch(p, DL, beam_width=64, prune_history=True) del ddec, DL torch.cuda.empty_cache() nadd = 0 for i in range(len(sna_ids)): seen = set(cands[i]) for b in dbeams[i][:DON_NBEST]: h = norm(b[0]) if h and h not in seen: seen.add(h) cands[i].append(h) # meme echelle EXACTE : le beam shona est du CTC pur, sans LM AC1[i].append(ctc_score(L1[i], encode_for(tok, h), tok.pad_token_id)) ISEXT[i].append(True) nadd += 1 del dbeams, L1 print("donneur %s : %d hypotheses NOUVELLES (%.1f/clip)" % (os.path.basename(DONOR), nadd, nadd / max(len(sna_ids), 1)), flush=True) AC1 = [np.array(x) for x in AC1] ISEXT = [np.array(x) for x in ISEXT] NW = [np.array([float(len(x.split())) for x in c]) for c in cands] def judge(path): pr, LG = compute_logits(path, files) t = pr.tokenizer R = [np.array([ctc_score(LG[i], encode_for(t, x), t.pad_token_id) for x in cands[i]]) for i in range(len(sna_ids))] del LG torch.cuda.empty_cache() print("juge %s OK" % os.path.basename(path), flush=True) return R R1, R2, R3 = judge(J1), judge(J2), judge(J3) from huggingface_hub import HfApi api = HfApi(token=open(os.path.expanduser("~/.cache/huggingface/token")).read().strip()) for mg in MARGINS: out = dict(base) nchg = next_ = 0 for i, k in enumerate(sna_ids): tot = AC1[i] + W1 * R1[i] + W2 * R2[i] + W3 * R3[i] + GAMMA * NW[i] bm = ~ISEXT[i] jb = int(np.argmax(np.where(bm, tot, -1e18))) nref, cref = len(cands[i][jb].split()), len(cands[i][jb]) ct = tot.copy() ct[ISEXT[i]] -= mg for j2 in np.flatnonzero(ISEXT[i]): r = len(cands[i][j2].split()) / max(nref, 1) rc = len(cands[i][j2]) / max(cref, 1) if r < RLO or r > RHI or rc < RLO or rc > RHI: ct[j2] = -1e18 j = int(np.argmax(ct)) if ISEXT[i][j]: next_ += 1 pick = cands[i][j] or greedy[i] or "a" if pick != base[k]: nchg += 1 out[k] = pick empt = sum(1 for x in out.values() if not str(x).strip()) dlin = sum(1 for k in base if lang.get(k) != "sna" and out[k] != base[k]) tag = "%s%s" % (TAG, ("CTL" if mg > 1e8 else "%03d" % round(mg))) OUT = "/root/sub_%s.csv" % tag with open(OUT, "w", newline="", encoding="utf-8") as f: wr = csv.writer(f) wr.writerow(["ID", "Target"]) for k in base: wr.writerow([k, out[k] or "a"]) assert len(out) == 892 and empt == 0 and dlin == 0, "%s INVALIDE" % tag if mg <= 1e8: api.upload_file(path_or_fileobj=OUT, path_in_repo="phase2_corrected/sub_%s.csv" % tag, repo_id="Pricile/waxal2026-backup", repo_type="model") flag = " <-- CONTROLE : doit etre 0" if mg > 1e8 else "" print("%-8s MARGIN=%8.1f | sna modifies %3d/%d | externes retenus %3d | lin intouche%s" % (tag, mg, nchg, len(sna_ids), next_, flag), flush=True) print("CROSSPOOL_SNA_DONE", flush=True) if __name__ == "__main__": main()