| """SOUMISSION : lin -> joint_cont + KenLM beam(o5,a0.5,b1.0) ; sna -> sna_ps greedy. |
| Routage = test_lang.json (identique a DEMI_B qui a fait 0.7491).""" |
| import csv, glob, json, os, pickle, numpy as np, soundfile as sf, torch |
| from transformers import AutoModelForCTC, AutoProcessor |
| from pyctcdecode import build_ctcdecoder |
| from multiprocessing import Pool |
| SR=16000; CACHE="/scratch/p2_16k" |
| lang=json.load(open("/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] |
| def norm(s): return " ".join(str(s).replace("|"," ").split()) |
| out={} |
|
|
| def decode(mdl, sel, use_lm, arpa=None, alpha=None, beta=None): |
| proc=AutoProcessor.from_pretrained(mdl); tok=proc.tokenizer |
| m=AutoModelForCTC.from_pretrained(mdl,dtype=torch.float32 if use_lm else torch.bfloat16).cuda().eval() |
| durs={f:sf.info(f).duration for f in sel}; sel=sorted(sel,key=lambda f:-durs[f]) |
| bs,cur,bud=[],[],0.0 |
| for f in sel: |
| if cur and bud+durs[f]>(90 if use_lm else 140): bs.append(cur); cur,bud=[],0.0 |
| cur.append(f); bud+=durs[f] |
| if cur: bs.append(cur) |
| logs=[]; order=[] |
| with torch.inference_mode(): |
| for b in bs: |
| 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.to("cuda",dtype=torch.bfloat16 if (not use_lm and v.dtype==torch.float32) else v.dtype) for k,v in x.items()} |
| lg=m(**x).logits |
| if use_lm: |
| lg=lg.log_softmax(-1).float().cpu().numpy() |
| for j,f in enumerate(b): logs.append(lg[j]); order.append(f) |
| else: |
| pid=lg.float().argmax(-1).cpu().numpy() |
| for f,s in zip(b,proc.batch_decode(pid)): out[os.path.splitext(os.path.basename(f))[0]]=norm(s) |
| del m; torch.cuda.empty_cache() |
| if use_lm: |
| 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) |
| with Pool(8) as p: hyps=dec.decode_batch(p,logs,beam_width=64) |
| for f,h,lg in zip(order,hyps,logs): |
| h=norm(h) |
| g=norm(tok.decode(lg.argmax(-1))) |
| if h and g: h=g[:1]+h[1:] |
| out[os.path.splitext(os.path.basename(f))[0]]=h |
|
|
| 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(f"lin={len(lin)} (joint_cont + KenLM 5g a0.5 b1.0) | sna={len(sna)} (sna_ps greedy)",flush=True) |
| decode("/root/models/joint_cont_best",lin,True,os.environ.get("ARPA","/scratch/lm/lin_5g.arpa"),0.5,float(os.environ.get("BETA","0.5"))) |
| print("lin OK",flush=True) |
| decode("/root/models/sna_ps_best",sna,False) |
| print("sna OK",flush=True) |
| OUT=os.environ.get("OUT","/root/sub_p2_KENLM.csv") |
| 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"]) |
| emp=sum(1 for i in ids if not out.get(i,'').strip()) |
| print('KENLM_SUB_DONE %s | %d IDs | vides=%d' % (OUT,len(ids),emp),flush=True) |
| for i in ids[:3]: print(" ",i,":",out.get(i,"")[:70]) |
|
|