waxal2026-backup / code /eval_devhard_generic.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
1.51 kB
import sys, json, os
import soundfile as sf, torch, jiwer
from transformers import AutoModelForCTC, AutoProcessor
lang, path = sys.argv[1], sys.argv[2]
SR=16000; dh=f"/scratch/prep/manifests/waxal_{lang}_devhard.jsonl"
rows=[]
for l in open(dh,encoding="utf-8"):
r=json.loads(l); r["text"]=" ".join(r.get("text","").replace("|"," ").split())
if r["text"] and os.path.exists(r["audio"]): rows.append(r)
proc=AutoProcessor.from_pretrained(path)
m=AutoModelForCTC.from_pretrained(path,torch_dtype=torch.bfloat16).cuda().eval()
rows.sort(key=lambda r:-r["duration"]); out={}; bs=[]; cur=[]; bud=0.0
for r in rows:
if cur and bud+r["duration"]>120: bs.append(cur); cur=[]; bud=0.0
cur.append(r); bud+=r["duration"]
if cur: bs.append(cur)
with torch.inference_mode():
for b in bs:
au=[sf.read(r["audio"],dtype="float32")[0] for r in b]
f=proc(au,sampling_rate=SR,return_tensors="pt",padding=True)
f={k:v.to("cuda",dtype=torch.bfloat16 if v.dtype==torch.float32 else v.dtype) for k,v in f.items()}
ids=m(**f).logits.float().argmax(-1).cpu().numpy()
for r,s in zip(b,proc.batch_decode(ids)): out[r["id"]]=" ".join(s.replace("|"," ").split())
KNOWN={"lin":0.2970,"sna":0.2238,"lug":0.0996}
refs=[r["text"] for r in rows]; hs=[out.get(r["id"],"") for r in rows]
wer=jiwer.wer(refs,hs); cer=jiwer.cer(refs,hs); comb=0.5*wer+0.5*cer
print(f"{lang} devhard: WER={wer:.4f} CER={cer:.4f} combine={comb:.4f} (champion ~{KNOWN[lang]}) delta={comb-KNOWN[lang]:+.4f}")