Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
4.95 kB
#!/usr/bin/env python3
"""Generateur de soumission parametrable : lin = <modele> + beam KenLM, sna = sna_ps greedy.
Casse du 1er caractere toujours copiee du greedy (gain valide).
Usage : MODEL=... ARPA=... ALPHA=.. BETA=.. LSB=0|1 OUT=... python gen_sub.py
"""
import csv
import glob
import json
import os
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"
LINM = os.environ.get("MODEL", "/root/models/joint_cont_best")
ARPA = os.environ.get("ARPA", "/scratch/lm/lin_5g.arpa")
ALPHA = float(os.environ.get("ALPHA", "0.5"))
BETA = float(os.environ.get("BETA", "0.5"))
LSB = os.environ.get("LSB", "0") == "1"
OUT = os.environ.get("OUT", "/root/sub_gen.csv")
def norm(s):
return " ".join(str(s).replace("|", " ").split())
def batches(sel, budget):
durs = {f: sf.info(f).duration for f in sel}
sel = sorted(sel, key=lambda f: -durs[f])
bs, cur, acc = [], [], 0.0
for f in sel:
if cur and acc + durs[f] > budget:
bs.append(cur)
cur, acc = [], 0.0
cur.append(f)
acc += durs[f]
if cur:
bs.append(cur)
return bs
def main():
lang = json.load(open(os.environ.get("LANGF", "/root/test_lang.json")))
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 (%s + KenLM a=%.2f b=%.2f lsb=%s) | sna=%d (sna_ps greedy)"
% (len(lin), os.path.basename(LINM), ALPHA, BETA, LSB, len(sna)), flush=True)
out = {}
# ---- lin : beam + KenLM ----
proc = AutoProcessor.from_pretrained(LINM)
tok = proc.tokenizer
m = AutoModelForCTC.from_pretrained(LINM, dtype=torch.float32).cuda().eval()
logs, order, greedy = [], [], []
with torch.inference_mode():
for b in batches(lin, 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):
logs.append(lg[j])
order.append(f)
greedy.append(norm(tok.decode(lg[j].argmax(-1))))
del m
torch.cuda.empty_cache()
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:
hyps = dec.decode_batch(p, logs, beam_width=64)
for f, h, g in zip(order, hyps, greedy):
h = norm(h)
if h and g:
h = g[:1] + h[1:] # casse du 1er caractere = celle du modele acoustique
out[os.path.splitext(os.path.basename(f))[0]] = h
print("lin OK", flush=True)
# ---- sna : greedy (le LM degrade, mesure) ----
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)
# ---- ecriture, en comblant les vides depuis la meilleure soumission connue ----
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"])
empty = sum(1 for i in ids if not out.get(i, "").strip())
print("GEN_DONE %s | %d IDs | vides=%d | combles=%d" % (OUT, len(ids), empty, filled), flush=True)
if __name__ == "__main__":
main()