File size: 1,709 Bytes
6eed659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
import json,pickle,jiwer
from transformers import AutoProcessor
from pyctcdecode import build_ctcdecoder
from multiprocessing import Pool
ROWS=[json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl")]
SUB=[r for r in ROWS if r["lang"]=="lin"]; refs=[r["text"] for r in SUB]
LOG=pickle.load(open("/scratch/lm/logits_lin.pkl","rb"))
tok=AutoProcessor.from_pretrained("/root/models/joint_cont_best").tokenizer
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]=""
g=[" ".join(tok.decode(l.argmax(-1)).replace("|"," ").split()) for l in LOG]
def comb(a,b):
    pr=[(x,y) for x,y in zip(a,b) if x.strip()]
    A=[x for x,_ in pr]; B=[y for _,y in pr]
    return 0.5*jiwer.wer(A,B)+0.5*jiwer.cer(A,B)
def run(a,b,bw=64):
    dec=build_ctcdecoder(lab,kenlm_model_path="/scratch/lm/lin_5g.arpa",alpha=a,beta=b,lm_score_boundary=False)
    with Pool(8) as p: h=[" ".join(x.split()) for x in dec.decode_batch(p,LOG,beam_width=bw)]
    h=[(gg[:1]+x[1:] if x and gg else x) for x,gg in zip(h,g)]
    return comb(refs,h)
print("=== lm_score_boundary=False : re-optimisation alpha x beta ===",flush=True)
best=(9,None)
for a in [0.4,0.5,0.6]:
    cells=[]
    for b in [0.0,0.25,0.5,0.75,1.0]:
        m=run(a,b); cells.append("b%.2f=%.4f"%(b,m))
        if m<best[0]: best=(m,(a,b))
    print("  a=%.1f  %s"%(a," ".join(cells)),flush=True)
print("  >>> best %.4f %s"%(best[0],best[1]),flush=True)
a,b=best[1]
for bw in [128,256]:
    print("  beam=%d : %.4f"%(bw,run(a,b,bw)),flush=True)
json.dump({"combine":best[0],"alpha":a,"beta":b},open("/root/sweep_lsb.json","w"))
print("SWEEP_LSB_DONE",flush=True)