File size: 6,661 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 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | #!/usr/bin/env python3
"""Soumission avec RESCORING N-BEST sur le lingala :
lin : beams KenLM de joint_cont (+ hypothese decode_batch) reordonnes par
score = ac_cont + lm_kenlm + LAMBDA * ac_cont2 (LAMBDA par env, defaut 1.5)
sna : sna_ps greedy
Casse du 1er caractere copiee du greedy. Routage par LANGF.
"""
import csv
import glob
import json
import os
import numpy as np
import soundfile as sf
import torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoModelForCTC, AutoProcessor
SR = 16000
CACHE = "/scratch/p2_16k"
M1 = "/root/models/joint_cont_best"
M2 = "/root/models/joint_cont2_best"
ARPA = os.environ.get("ARPA", "/scratch/lm/lin_5g.arpa")
LAMBDA = float(os.environ.get("LAMBDA", "1.5"))
GAMMA = float(os.environ.get("GAMMA", "2.0")) # bonus par mot (compense le biais des sommes de log-probs)
NBEST = 10
LANGF = os.environ.get("LANGF", "/root/test_lang_gpulid.json")
OUT = os.environ.get("OUT", "/root/sub_rescore.csv")
def norm(s):
return " ".join(str(s).replace("|", " ").split())
def batches(sel, budget):
d = {f: sf.info(f).duration for f in sel}
sel = sorted(sel, key=lambda f: -d[f])
bs, cur, acc = [], [], 0.0
for f in sel:
if cur and acc + d[f] > budget:
bs.append(cur)
cur, acc = [], 0.0
cur.append(f)
acc += d[f]
if cur:
bs.append(cur)
return bs
def logits_for(model_dir, files, dtype=torch.float32):
proc = AutoProcessor.from_pretrained(model_dir)
m = AutoModelForCTC.from_pretrained(model_dir, dtype=dtype).cuda().eval()
res = {}
with torch.inference_mode():
for b in batches(files, 90):
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.cuda() for k, v in x.items()}
lg = m(**x).logits.log_softmax(-1).float().cpu().numpy()
for j, f in enumerate(b):
res[f] = lg[j]
del m
torch.cuda.empty_cache()
return proc, res
def ctc_scores(logp, texts, tok):
T = logp.shape[0]
lp = torch.from_numpy(logp).unsqueeze(1)
out = []
for t in texts:
ids = [i for i in tok(t.replace(" ", "|")).input_ids if i != tok.pad_token_id] if t else []
if not ids or len(ids) > T:
out.append(-1e9)
continue
loss = torch.nn.functional.ctc_loss(
lp, torch.tensor(ids).unsqueeze(0), torch.tensor([T]), torch.tensor([len(ids)]),
blank=tok.pad_token_id, reduction="sum", zero_infinity=True)
out.append(-float(loss))
return out
def main():
lang = json.load(open(LANGF))
files = sorted(glob.glob(os.path.join(CACHE, "*.wav")))
ids = [os.path.splitext(os.path.basename(f))[0] for f in files]
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("lin=%d (rescoring lambda=%.1f gamma=%.1f) | sna=%d (sna_ps greedy)"
% (len(lin), LAMBDA, GAMMA, len(sna)), flush=True)
out = {}
proc1, LG1 = logits_for(M1, lin)
print("logits joint_cont OK", flush=True)
_, LG2 = logits_for(M2, lin)
print("logits joint_cont2 OK", flush=True)
tok = proc1.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] = ""
_A=float(os.environ.get("ALPHA","0.5")); _B=float(os.environ.get("BETA","0.5"))
_LSB=os.environ.get("LSB","0") not in ("0","false","False")
print("DECODE alpha=%s beta=%s lsb=%s" % (_A,_B,_LSB), flush=True)
dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=_A, beta=_B,
lm_score_boundary=_LSB)
order = lin
L1 = [LG1[f] for f in order]
with Pool(8) as p:
allbeams = dec.decode_beams_batch(p, L1, beam_width=64)
with Pool(8) as p:
db = [" ".join(x.split()) for x in dec.decode_batch(p, L1, beam_width=64)]
for i, f in enumerate(order):
g = norm(tok.decode(L1[i].argmax(-1)))
bs = allbeams[i]
cands = [" ".join(b[0].split()) for b in bs[:NBEST]]
ac1 = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]]
lm = [((b[4] - b[3]) if len(b) > 4 else 0.0) for b in bs[:NBEST]]
if db[i] not in cands:
cands.append(db[i])
ac1.append(ctc_scores(L1[i], [db[i]], tok)[0])
lm.append(float(np.mean(lm)) if lm else 0.0)
ac2 = ctc_scores(LG2[f], cands, tok)
nw = np.array([float(len(x.split())) for x in cands])
tot = np.array(ac1) + np.array(lm) + LAMBDA * np.array(ac2) + GAMMA * nw
h = norm(cands[int(np.argmax(tot))])
if h and g:
h = g[:1] + h[1:]
out[os.path.splitext(os.path.basename(f))[0]] = h
if (i + 1) % 150 == 0:
print(" rescore %d/%d" % (i + 1, len(order)), flush=True)
print("lin OK", flush=True)
proc2 = AutoProcessor.from_pretrained("/root/models/sna_ps_best")
m2 = AutoModelForCTC.from_pretrained("/root/models/sna_ps_best",
dtype=torch.bfloat16).cuda().eval()
with torch.inference_mode():
for b in batches(sna, 140):
au = [sf.read(f, dtype="float32")[0] for f in b]
x = proc2(au, sampling_rate=SR, 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 = m2(**x).logits.float().argmax(-1).cpu().numpy()
for f, s in zip(b, proc2.batch_decode(pid)):
out[os.path.splitext(os.path.basename(f))[0]] = norm(s)
del m2
torch.cuda.empty_cache()
print("sna OK", flush=True)
fb = {}
ref = "/root/sub_p2_KENLM.csv"
if os.path.exists(ref):
fb = {r["ID"]: r["Target"] for r in csv.DictReader(open(ref, encoding="utf-8"))}
filled = 0
for i in ids:
if not out.get(i, "").strip() and fb.get(i, "").strip():
out[i] = fb[i]
filled += 1
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"])
print("RESCORE_GEN_DONE %s | %d IDs | vides=%d | combles=%d"
% (OUT, len(ids), sum(1 for i in ids if not out.get(i, "").strip()), filled), flush=True)
if __name__ == "__main__":
main()
|