#!/usr/bin/env python3 """POOL INTER-MODELES sur le LINGALA : les juges peuvent enfin CHOISIR une hypothese que le modele principal n'a jamais propose. Limite du mecanisme actuel : un juge CLASSE des candidats, il n'en PROPOSE aucun. Or tous les candidats sortent du beam de `joint_cont` seul. Si MMS-1B transcrit mieux un clip, sa transcription n'est meme pas dans la liste. Et ce n'est pas theorique : MMS-1B bat le champion sur son propre eval (0.2950 < 0.2966). /!\\ A NE PAS CONFONDRE avec l'elargissement de pool deja refute (-0.0023) : celui-la ajoutait PLUS DE CANDIDATS DU MEME MODELE (meme distribution, donc surtout de mauvais candidats en plus). Ici on ajoute des candidats d'une AUTRE distribution. CALIBRATION (le point delicat) : les candidats du beam portent le score interne de pyctcdecode (acoustique + LM + longueur). Un candidat externe n'en a pas. On le place sur la meme echelle avec une formule uniforme S(h) = ctc(joint_cont, h) + ALPHA*kenlm(h) + BETA*mots(h) recalee par clip : offset_k = mediane sur les candidats du beam de (score_beam - S). => les candidats du beam GARDENT leur score exact du record => MARGIN=+inf reproduit le record a l'identique (CONTROLE) MARGIN : un candidat externe doit depasser le meilleur candidat du beam de MARGIN pour etre retenu. Balayage de MARGIN = prudence decroissante. """ import csv import json import os import pickle import sys import kenlm import numpy as np import torch from multiprocessing import Pool from pyctcdecode import build_ctcdecoder from transformers import AutoProcessor 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_SN3020.csv") LINM = "/root/models/joint_cont_best" ARPA = os.environ.get("ARPA", "/scratch/lm/lin_5g.arpa") J1 = os.environ.get("J1", "/scratch/runs/xlsrlong/checkpoint-5600") J2 = os.environ.get("J2", "/root/models/mms1b_lin_best") # plusieurs donneurs (separes par des virgules) : leurs hypotheses s ACCUMULENT # dans le meme pool au lieu de se remplacer. Indispensable pour empiler MMS et # AfriHuBERT : repartir du beam seul JETAIT les candidats MMS deja gagnants # (controle rouge a 28 clips). DONORS = [x for x in os.environ.get("DONOR", "/root/models/mms1b_lin_best").split(",") if x] # 3e JUGE OPTIONNEL, note les candidats SANS reconstruire le pool (contrairement a # lin_judge3.py, qui repart du beam et PERD les hypotheses externes deja gagnantes). # C est le seul moyen de tester un juge lingala en gardant le crosspool intact. J3 = os.environ.get("J3", "") W3 = float(os.environ.get("W3", "0")) W1 = float(os.environ.get("W1", "1.0")) W2 = float(os.environ.get("W2", "1.0")) MARGINS = [float(x) for x in os.environ.get("MARGINS", "1e9,20,10,5,0").split(",")] DON_NBEST = int(os.environ.get("DON_NBEST", "5")) # GARDE-FOU DE LONGUEUR (ajoute apres le refus de XP020 a -0.0091). # Diagnostic : les candidats externes retenus etaient 1.56x PLUS LONGS (mediane 174 -> # 265 car.). La formule de calibration contient un bonus de longueur (BETA=1.0/mot) # calibre sur des candidats de longueur homogene ; elle s emballe des qu un candidat # sort de cette plage => insertions massives. On n accepte donc un externe que si sa # longueur en mots reste proche du meilleur candidat du beam. RLO = float(os.environ.get("RATIO_LO", "0.85")) RHI = float(os.environ.get("RATIO_HI", "1.15")) TAG = os.environ.get("TAG", "XP") AUD = "/scratch/p2_16k" # ALPHA (poids du LM) et BETA (prime dinsertion de MOTS) etaient figes depuis le # debut, regles dans un systeme sans crosspool ni juges MMS ni ponctuation. # Le leaderboard du 9 aout montre que notre CER est le 2e meilleur des 4 premiers # mais notre WER le PIRE : 66 % du retard vient du niveau MOT. BETA agit exactement # la — il controle combien de mots le beam accepte de produire. ALPHA = float(os.environ.get("ALPHA", "0.6")) BETA = float(os.environ.get("BETA", "1.0")) LSB, BW, NBEST = True, 64, 24 def main(): base = {r["ID"]: r["Target"] for r in csv.DictReader(open(BASE, encoding="utf-8"))} lang = json.load(open("/root/test_lang.json")) lin_ids = [k for k in base if lang.get(k) == "lin"] print("base %d | lin %d | sna intouches %d" % (len(base), len(lin_ids), len(base) - len(lin_ids)), flush=True) # ---------- 1. candidats du modele principal (config EXACTE du record) ---------- with open("/scratch/lm/logits_test_lin.pkl", "rb") as f: keys, logs, greedy = pickle.load(f) proc = AutoProcessor.from_pretrained(LINM) 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] = "" dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=ALPHA, beta=BETA, lm_score_boundary=LSB) with Pool(8) as p: beams = dec.decode_beams_batch(p, logs, beam_width=BW, prune_history=True) del dec gmap = dict(zip(keys, greedy)) LGmap = dict(zip(keys, logs)) CAND, SC, NBEAM = {}, {}, {} for k, bs in zip(keys, beams): c, s = [], [] for b in bs[:NBEST]: h = norm(b[0]) g = gmap[k] if h and g: h = g[:1] + h[1:] h = h[:1].upper() + h[1:] if h else h if h and h not in c: c.append(h) s.append(float(b[3])) CAND[k], SC[k], NBEAM[k] = c, s, len(c) del beams print("beam principal : %d clips" % len(CAND), flush=True) order = [k for k in lin_ids if CAND.get(k)] fmap = {k: os.path.join(AUD, k + ".wav") for k in order} files = [fmap[k] for k in order] # ---------- 2. candidats DONNEURS (autre modele, autre distribution) ---------- EXTRA = {k: [] for k in order} for DONOR in DONORS: 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() nd = 0 for i, k in enumerate(order): seen = set(CAND[k]) | set(EXTRA[k]) for b in dbeams[i][:DON_NBEST]: h = norm(b[0]) h = h[:1].upper() + h[1:] if h else h if h and h not in seen: seen.add(h) EXTRA[k].append(h) nd += 1 del dbeams print("donneur %s : %d hypotheses NOUVELLES" % (os.path.basename(DONOR), nd), flush=True) nadd = sum(len(x) for x in EXTRA.values()) ncl = sum(1 for x in EXTRA.values() if x) print("POOL EXTERNE TOTAL : %d hypotheses sur %d clips (%.1f/clip)" % (nadd, ncl, nadd / max(ncl, 1)), flush=True) # ---------- 3. calibration des externes sur l'echelle du beam ---------- lm = kenlm.Model(ARPA) def s_uniform(k, h): ids = encode_for(tok, h) return (ctc_score(LGmap[k], ids, tok.pad_token_id) + ALPHA * lm.score(h, bos=LSB, eos=LSB) + BETA * len(h.split())) OFF, resid = {}, [] for k in order: d = [SC[k][j] - s_uniform(k, CAND[k][j]) for j in range(len(CAND[k]))] OFF[k] = float(np.median(d)) resid.append(float(np.std(d))) print("calibration : ecart-type residuel median %.3f (offset median %.1f)" % (float(np.median(resid)), float(np.median(list(OFF.values())))), flush=True) ALL, ALLS, ISEXT = {}, {}, {} for k in order: c = list(CAND[k]) + EXTRA[k] s = list(SC[k]) + [s_uniform(k, h) + OFF[k] for h in EXTRA[k]] ALL[k], ALLS[k] = c, np.array(s) ISEXT[k] = np.array([False] * len(CAND[k]) + [True] * len(EXTRA[k])) # ---------- 4. juges (sur TOUS les candidats, externes compris) ---------- def judge(path): pr, LG = compute_logits(path, files) t = pr.tokenizer R = {k: np.array([ctc_score(LG[i], encode_for(t, x), t.pad_token_id) for x in ALL[k]]) for i, k in enumerate(order)} del LG torch.cuda.empty_cache() print("juge %s OK" % os.path.basename(path), flush=True) return R R1, R2 = judge(J1), judge(J2) R3 = judge(J3) if (J3 and W3) else None # ---------- 5. selection avec marge de prudence ---------- 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 k in order: tot = ALLS[k] + W1 * R1[k] + W2 * R2[k] if R3 is not None: tot = tot + W3 * R3[k] beam_mask = ~ISEXT[k] best_beam = float(np.max(tot[beam_mask])) nref = len(ALL[k][int(np.argmax(np.where(beam_mask, tot, -1e18)))].split()) cand_tot = tot.copy() cand_tot[ISEXT[k]] -= mg # penalise les externes de MARGIN cref = len(ALL[k][int(np.argmax(np.where(beam_mask, tot, -1e18)))]) for j2 in np.flatnonzero(ISEXT[k]): # garde-fou de longueur r = len(ALL[k][j2].split()) / max(nref, 1) rc = len(ALL[k][j2]) / max(cref, 1) # en MOTS (protege le WER) ET en CARACTERES (protege le CER) : # le 1er correctif n avait borne que les mots -> ratio caracteres # encore a 1.40 (memes mots, formes 40 % plus longues). if r < RLO or r > RHI or rc < RLO or rc > RHI: cand_tot[j2] = -1e18 j = int(np.argmax(cand_tot)) if ISEXT[k][j] and tot[j] > best_beam: next_ += 1 h = ALL[k][j] or base[k] if h != base[k]: nchg += 1 out[k] = h empt = sum(1 for x in out.values() if not str(x).strip()) dsna = 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 dsna == 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("%-9s MARGIN=%8.1f | lin modifies %3d/%d | externes retenus %3d | sna intouche%s" % (tag, mg, nchg, len(lin_ids), next_, flag), flush=True) print("CROSSPOOL_DONE", flush=True) if __name__ == "__main__": main()