waxal2026-backup / archive /scripts_B /perlang_devhard.py
Pricile's picture
Upload folder using huggingface_hub
c887738 verified
Raw
History Blame Contribute Delete
2.01 kB
"""Decode le devhard avec chaque modele candidat -> combine PAR LANGUE.
But : trouver un HYBRIDE par langue (meilleur modele pour lin, meilleur pour sna)."""
import json, os, soundfile as sf, torch, jiwer
from transformers import AutoModelForCTC, AutoProcessor
ROWS=[json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl")]
def norm(s): return " ".join(str(s).replace("|"," ").split())
def comb(R,H):
pr=[(r,h) for r,h in zip(R,H) if r.strip()]
r=[x for x,_ in pr]; h=[x for _,x in pr]
w=jiwer.wer(r,h); c=jiwer.cer(r,h); return w,c,0.5*w+0.5*c
CANDS=["joint_cont_best","joint_cont2_best","joint_best","mms1b_lin_best","sna_ps_best","lin_s4_best"]
allhyps={}
for name in CANDS:
M=f"/root/models/{name}"
try:
proc=AutoProcessor.from_pretrained(M)
model=AutoModelForCTC.from_pretrained(M,dtype=torch.bfloat16).cuda().eval()
except Exception as e:
print(f"{name}: SKIP ({str(e)[:60]})",flush=True); continue
hy={}
with torch.inference_mode():
for i in range(0,len(ROWS),8):
b=ROWS[i:i+8]
au=[sf.read(r["audio"],dtype="float32")[0] for r in b]
x=proc(au,sampling_rate=16000,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=model(**x).logits.float().argmax(-1).cpu().numpy()
for r,h in zip(b,proc.batch_decode(pid)): hy[r["id"]]=norm(h)
del model; torch.cuda.empty_cache()
allhyps[name]=hy
line=[name]
for lang in ["lin","sna"]:
sub=[r for r in ROWS if r["lang"]==lang]
w,c,m=comb([r["text"] for r in sub],[hy[r["id"]] for r in sub])
line.append(f"{lang}={m:.4f}")
sub=ROWS
w,c,m=comb([r["text"] for r in sub],[hy[r["id"]] for r in sub])
line.append(f"ALL={m:.4f}")
print(" ".join(line),flush=True)
json.dump(allhyps,open("/root/devhard_allhyps.json","w",encoding="utf-8"),ensure_ascii=False)
print("PERLANG_DONE",flush=True)