| |
| """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) |
| 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() |
|
|