Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
4.52 kB
#!/usr/bin/env python3
"""RESCORING N-BEST v2 : sweep etendu de LAMBDA + l'hypothese decode_batch injectee dans le
pool de candidats (decode_beams renvoie un 1-best different, souvent moins bon).
Ajoute aussi joint_best comme 3e rescoreur optionnel (MU).
score(h) = ac_cont(h) + lm_kenlm(h) + LAMBDA*ac_cont2(h) + MU*ac_jbest(h)
Les scores acoustiques manquants (hypothese venue de decode_batch) sont calcules par CTC.
"""
import json
import pickle
import jiwer
import numpy as np
import torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoProcessor
M1 = "/root/models/joint_cont_best"
ARPA = "/scratch/lm/lin_5g.arpa"
NBEST = 10
def comb(refs, hyps):
pr = [(r, h) for r, h in zip(refs, hyps) if r.strip()]
a = [x for x, _ in pr]
b = [y for _, y in pr]
w = jiwer.wer(a, b)
c = jiwer.cer(a, b)
return w, c, 0.5 * w + 0.5 * c
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():
rows = [json.loads(l) for l in open("/root/devhard/devhard_linsna.jsonl", encoding="utf-8")]
sub = [r for r in rows if r["lang"] == "lin"]
refs = [r["text"] for r in sub]
L1 = pickle.load(open("/scratch/lm/logits_lin.pkl", "rb"))
L2 = pickle.load(open("/scratch/lm/logits_lin_cont2.pkl", "rb"))
tok = AutoProcessor.from_pretrained(M1).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] = ""
greedy = [" ".join(tok.decode(l.argmax(-1)).replace("|", " ").split()) for l in L1]
dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=0.5, beta=0.5,
lm_score_boundary=False)
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)]
def cc(h, g):
return (g[:1] + h[1:]) if (h and g) else h
_, _, m_db = comb(refs, [cc(h, g) for h, g in zip(db, greedy)])
_, _, m_bm = comb(refs, [cc(bs[0][0], g) for bs, g in zip(allbeams, greedy)])
print("reference decode_batch : %.4f" % m_db, flush=True)
print("reference decode_beams : %.4f" % m_bm, flush=True)
REF = min(m_db, m_bm)
prepared = []
for i, bs in enumerate(allbeams):
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]]
# injecte l'hypothese decode_batch si absente (scores calcules par CTC + LM du beam le plus proche)
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(L2[i], cands, tok)
prepared.append((cands, np.array(ac1), np.array(lm), np.array(ac2)))
if (i + 1) % 200 == 0:
print(" prepare %d/%d" % (i + 1, len(allbeams)), flush=True)
print("\n--- sweep LAMBDA etendu ---", flush=True)
best = (9.0, None)
for lam in (0.0, 0.3, 0.5, 1.0, 1.5, 2.0, 3.0, 5.0, 8.0, 12.0):
hyps = []
for (cands, ac1, lm, ac2), g in zip(prepared, greedy):
tot = ac1 + lm + lam * ac2
hyps.append(cc(cands[int(np.argmax(tot))], g))
_, _, m = comb(refs, hyps)
mark = " <<<" if m < best[0] else ""
print(" lambda=%5.1f : %.4f (vs ref %.4f : %+.4f)%s" % (lam, m, REF, m - REF, mark), flush=True)
if m < best[0]:
best = (m, lam)
print("\nBEST_V2 %.4f lambda=%.1f (ref %.4f, gain %+.4f)"
% (best[0], best[1], REF, best[0] - REF), flush=True)
json.dump({"combine": best[0], "lambda": best[1], "ref": REF},
open("/root/nbest_v2.json", "w"))
print("NBEST_V2_DONE", flush=True)
if __name__ == "__main__":
main()