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