waxal2026-backup / phase2_corrected /code /gen_sna_rescore.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
4.86 kB
#!/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()