| |
| """Soumission avec RESCORING N-BEST sur le lingala : |
| lin : beams KenLM de joint_cont (+ hypothese decode_batch) reordonnes par |
| score = ac_cont + lm_kenlm + LAMBDA * ac_cont2 (LAMBDA par env, defaut 1.5) |
| sna : sna_ps greedy |
| Casse du 1er caractere copiee du greedy. Routage par LANGF. |
| """ |
| import csv |
| import glob |
| import json |
| import os |
|
|
| import numpy as np |
| import soundfile as sf |
| import torch |
| from multiprocessing import Pool |
| from pyctcdecode import build_ctcdecoder |
| from transformers import AutoModelForCTC, AutoProcessor |
|
|
| SR = 16000 |
| CACHE = "/scratch/p2_16k" |
| M1 = "/root/models/joint_cont_best" |
| M2 = "/root/models/joint_cont2_best" |
| ARPA = os.environ.get("ARPA", "/scratch/lm/lin_5g.arpa") |
| LAMBDA = float(os.environ.get("LAMBDA", "1.5")) |
| GAMMA = float(os.environ.get("GAMMA", "2.0")) |
| NBEST = 10 |
| LANGF = os.environ.get("LANGF", "/root/test_lang_gpulid.json") |
| OUT = os.environ.get("OUT", "/root/sub_rescore.csv") |
|
|
|
|
| def norm(s): |
| return " ".join(str(s).replace("|", " ").split()) |
|
|
|
|
| def batches(sel, budget): |
| d = {f: sf.info(f).duration for f in sel} |
| sel = sorted(sel, key=lambda f: -d[f]) |
| bs, cur, acc = [], [], 0.0 |
| for f in sel: |
| if cur and acc + d[f] > budget: |
| bs.append(cur) |
| cur, acc = [], 0.0 |
| cur.append(f) |
| acc += d[f] |
| if cur: |
| bs.append(cur) |
| return bs |
|
|
|
|
| def logits_for(model_dir, files, dtype=torch.float32): |
| proc = AutoProcessor.from_pretrained(model_dir) |
| m = AutoModelForCTC.from_pretrained(model_dir, dtype=dtype).cuda().eval() |
| res = {} |
| with torch.inference_mode(): |
| for b in batches(files, 90): |
| au = [sf.read(f, dtype="float32")[0] for f in b] |
| x = proc(au, sampling_rate=SR, 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, f in enumerate(b): |
| res[f] = lg[j] |
| del m |
| torch.cuda.empty_cache() |
| return proc, res |
|
|
|
|
| def ctc_scores(logp, texts, tok): |
| T = logp.shape[0] |
| lp = torch.from_numpy(logp).unsqueeze(1) |
| out = [] |
| for t in texts: |
| ids = [i for i in tok(t.replace(" ", "|")).input_ids if i != tok.pad_token_id] if t else [] |
| if not ids or len(ids) > T: |
| out.append(-1e9) |
| continue |
| loss = torch.nn.functional.ctc_loss( |
| lp, torch.tensor(ids).unsqueeze(0), torch.tensor([T]), torch.tensor([len(ids)]), |
| blank=tok.pad_token_id, reduction="sum", zero_infinity=True) |
| out.append(-float(loss)) |
| return out |
|
|
|
|
| def main(): |
| lang = json.load(open(LANGF)) |
| files = sorted(glob.glob(os.path.join(CACHE, "*.wav"))) |
| ids = [os.path.splitext(os.path.basename(f))[0] for f in files] |
| lin = [f for f in files if lang[os.path.splitext(os.path.basename(f))[0]] == "lin"] |
| sna = [f for f in files if lang[os.path.splitext(os.path.basename(f))[0]] == "sna"] |
| print("lin=%d (rescoring lambda=%.1f gamma=%.1f) | sna=%d (sna_ps greedy)" |
| % (len(lin), LAMBDA, GAMMA, len(sna)), flush=True) |
| out = {} |
|
|
| proc1, LG1 = logits_for(M1, lin) |
| print("logits joint_cont OK", flush=True) |
| _, LG2 = logits_for(M2, lin) |
| print("logits joint_cont2 OK", flush=True) |
|
|
| tok = proc1.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] = "" |
| _A=float(os.environ.get("ALPHA","0.5")); _B=float(os.environ.get("BETA","0.5")) |
| _LSB=os.environ.get("LSB","0") not in ("0","false","False") |
| print("DECODE alpha=%s beta=%s lsb=%s" % (_A,_B,_LSB), flush=True) |
| dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=_A, beta=_B, |
| lm_score_boundary=_LSB) |
|
|
| order = lin |
| L1 = [LG1[f] for f in order] |
| with Pool(8) as p: |
| allbeams = dec.decode_beams_batch(p, L1, beam_width=64) |
| with Pool(8) as p: |
| db = [" ".join(x.split()) for x in dec.decode_batch(p, L1, beam_width=64)] |
|
|
| for i, f in enumerate(order): |
| g = norm(tok.decode(L1[i].argmax(-1))) |
| bs = allbeams[i] |
| cands = [" ".join(b[0].split()) for b in bs[:NBEST]] |
| ac1 = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]] |
| lm = [((b[4] - b[3]) if len(b) > 4 else 0.0) for b in bs[:NBEST]] |
| if db[i] not in cands: |
| cands.append(db[i]) |
| ac1.append(ctc_scores(L1[i], [db[i]], tok)[0]) |
| lm.append(float(np.mean(lm)) if lm else 0.0) |
| ac2 = ctc_scores(LG2[f], cands, tok) |
| nw = np.array([float(len(x.split())) for x in cands]) |
| tot = np.array(ac1) + np.array(lm) + LAMBDA * np.array(ac2) + GAMMA * nw |
| h = norm(cands[int(np.argmax(tot))]) |
| if h and g: |
| h = g[:1] + h[1:] |
| out[os.path.splitext(os.path.basename(f))[0]] = h |
| if (i + 1) % 150 == 0: |
| print(" rescore %d/%d" % (i + 1, len(order)), flush=True) |
| print("lin OK", flush=True) |
|
|
| proc2 = AutoProcessor.from_pretrained("/root/models/sna_ps_best") |
| m2 = AutoModelForCTC.from_pretrained("/root/models/sna_ps_best", |
| dtype=torch.bfloat16).cuda().eval() |
| with torch.inference_mode(): |
| for b in batches(sna, 140): |
| au = [sf.read(f, dtype="float32")[0] for f in b] |
| x = proc2(au, sampling_rate=SR, return_tensors="pt", padding=True) |
| x = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) |
| for k, v in x.items()} |
| pid = m2(**x).logits.float().argmax(-1).cpu().numpy() |
| for f, s in zip(b, proc2.batch_decode(pid)): |
| out[os.path.splitext(os.path.basename(f))[0]] = norm(s) |
| del m2 |
| torch.cuda.empty_cache() |
| print("sna OK", flush=True) |
|
|
| fb = {} |
| ref = "/root/sub_p2_KENLM.csv" |
| if os.path.exists(ref): |
| fb = {r["ID"]: r["Target"] for r in csv.DictReader(open(ref, encoding="utf-8"))} |
| filled = 0 |
| for i in ids: |
| if not out.get(i, "").strip() and fb.get(i, "").strip(): |
| out[i] = fb[i] |
| filled += 1 |
| with open(OUT, "w", newline="", encoding="utf-8") as fo: |
| w = csv.writer(fo) |
| w.writerow(["ID", "Target"]) |
| for i in ids: |
| w.writerow([i, out.get(i) or "a"]) |
| print("RESCORE_GEN_DONE %s | %d IDs | vides=%d | combles=%d" |
| % (OUT, len(ids), sum(1 for i in ids if not out.get(i, "").strip()), filled), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|