#!/usr/bin/env python3 """Génère la soumission = RECORD 0.757995 avec les clips SHONA remplacés par la version rescorée (validée devhard-sna : 0.1281 -> 0.1252, sna_r2 w=1.5 gamma=0). Le LINGALA reste STRICTEMENT identique au record => Δsna isolé (règle d'or §4a). """ import csv, json, os import numpy as np, soundfile as sf, torch from multiprocessing import Pool from pyctcdecode import build_ctcdecoder from transformers import AutoModelForCTC, AutoProcessor BASE = os.environ.get("BASE", "/root/sub_p2corr_KENLM_lin_snaps.csv") LANGF = os.environ.get("LANGF", "/root/test_lang.json") OUT = os.environ.get("OUT", "/root/sub_SNARESC.csv") SNAM = "/root/models/sna_ps_best" RESC = "/scratch/restore/sna_r2_best" W = float(os.environ.get("W", "1.5")) GAMMA = float(os.environ.get("GAMMA", "0.0")) NBEST = 10 AUD = "/scratch/p2_16k" def norm(s): return " ".join(str(s).replace("|", " ").split()) 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, files): 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(files), 4): b = files[i:i + 4] au = [sf.read(f, dtype="float32")[0] for f 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(): 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 clips | sna a re-decoder : %d | lin inchanges : %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] missing = [f for f in files if not os.path.exists(f)] assert not missing, "audio manquant: %s" % missing[:3] 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] print("logits sna_ps OK", flush=True) dec = build_ctcdecoder(lab) # CTC pur, aucun LM (il degrade le shona) with Pool(8) as p: allbeams = dec.decode_beams_batch(p, L1, beam_width=64) with Pool(8) as p: db = [norm(x) for x in dec.decode_batch(p, L1, beam_width=64)] cands, AC1, NW = [], [], [] 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(np.array(a)) NW.append(np.array([float(len(x.split())) for x in c])) proc2, LG = compute_logits(RESC, files) t2 = proc2.tokenizer print("logits rescoreur sna_r2 OK", flush=True) out = dict(base); nchg = 0 for i, k in enumerate(sna_ids): sc = np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in cands[i]]) tot = AC1[i] + W * sc + GAMMA * NW[i] pick = cands[i][int(np.argmax(tot))] 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()) with open(OUT, "w", newline="", encoding="utf-8") as f: w = csv.writer(f); w.writerow(["ID", "Target"]) for k in base: w.writerow([k, out[k] or "a"]) print("SNARESC_DONE %s | %d lignes | sna modifies %d/%d | vides %d" % (OUT, len(out), nchg, len(sna_ids), empt), flush=True) if __name__ == "__main__": main()