waxal2026-backup / phase2_corrected /code /multi_rescore.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
7.42 kB
#!/usr/bin/env python3
"""RESCORING MULTI-MODELES sur le lingala.
Point cle : pour du rescoring, un modele n'a PAS besoin du meme vocabulaire ni de la meme
frequence de trames que le decodeur — il doit seulement savoir evaluer log P(texte | ses logits).
On peut donc recruter des modeles ecartes comme decodeurs (MMS, monolingues) comme rescoreurs.
score(h) = ac_cont(h) + lm_kenlm(h) + somme_i w_i * ac_modele_i(h)
Recherche des poids par montee de coordonnees sur devhard-lin.
"""
import json
import os
import pickle
import jiwer
import numpy as np
import soundfile as sf
import torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoModelForCTC, AutoProcessor
M1 = "/root/models/joint_cont_best"
ARPA = "/scratch/lm/lin_5g.arpa"
NBEST = 10
R = "/scratch/restore"
RESCORERS = [
("cont2", "/root/models/joint_cont2_best"),
("jbest", "/root/models/joint_best"),
("lin_s4", R + "/lin_s4_best"),
("mmsjoint", "/root/models/mmsjoint_best"),
]
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 encode_for(tok, text):
"""Encode en restant dans le vocabulaire du modele : les caracteres absents sont retires.
Permet de recruter des modeles sans casse/ponctuation (ils jugent alors le contenu seul)."""
v = tok.get_vocab()
delim = getattr(tok, "word_delimiter_token", "|")
s = text.replace(" ", delim)
keep = "".join(c for c in s if c in v)
if not keep:
low = text.lower().replace(" ", delim)
keep = "".join(c for c in low if c in v)
ids = [v[c] for c in keep]
return [i for i in ids if i != tok.pad_token_id]
def ctc_score(logp, ids, blank):
T = logp.shape[0]
if not ids or len(ids) > T:
return -1e9
lp = torch.from_numpy(logp).unsqueeze(1)
loss = torch.nn.functional.ctc_loss(
lp, torch.tensor(ids).unsqueeze(0), torch.tensor([T]), torch.tensor([len(ids)]),
blank=blank, reduction="sum", zero_infinity=True)
return -float(loss)
def compute_logits(model_dir, rows):
proc = AutoProcessor.from_pretrained(model_dir)
m = AutoModelForCTC.from_pretrained(model_dir, dtype=torch.float32).cuda().eval()
out = []
with torch.inference_mode():
for i in range(0, len(rows), 4):
b = rows[i:i + 4]
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.cuda() for k, v in x.items()}
lg = m(**x).logits.log_softmax(-1).float().cpu().numpy()
for j in range(len(b)):
out.append(lg[j])
del m
torch.cuda.empty_cache()
return proc, 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"))
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]
def cc(h, g):
return (g[:1] + h[1:]) if (h and g) else h
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)]
_, _, REF = comb(refs, [cc(h, g) for h, g in zip(db, greedy)])
print("reference (meilleur 1-best connu): %.4f" % REF, flush=True)
# pool de candidats + scores du decodeur
cands, AC1, LM = [], [], []
for i, bs in enumerate(allbeams):
c = [" ".join(b[0].split()) for b in bs[:NBEST]]
a = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]]
l = [((b[4] - b[3]) if len(b) > 4 else 0.0) for b in bs[:NBEST]]
if db[i] not in c:
c.append(db[i])
a.append(ctc_score(L1[i], encode_for(tok, db[i]), tok.pad_token_id))
l.append(float(np.mean(l)) if l else 0.0)
cands.append(c)
AC1.append(np.array(a))
LM.append(np.array(l))
print("taille moyenne du pool: %.1f hypotheses" % np.mean([len(c) for c in cands]), flush=True)
# scores de chaque rescoreur
SC = {}
chars = set("".join(refs))
for tag, mdl in RESCORERS:
if not os.path.isdir(mdl):
print("%-9s ABSENT (%s)" % (tag, mdl), flush=True)
continue
try:
proc, LG = compute_logits(mdl, sub)
except Exception as e:
print("%-9s ERREUR %s" % (tag, str(e)[:60]), flush=True)
continue
t2 = proc.tokenizer
vv = set(t2.get_vocab())
miss = sorted(c for c in chars if c not in vv and c != " ")
s = []
for i, c in enumerate(cands):
s.append(np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in c]))
SC[tag] = s
print("%-9s |V|=%d, caract. refs absents=%d -> scores calcules"
% (tag, len(vv), len(miss)), flush=True)
def evaluate(weights):
hyps = []
for i in range(len(cands)):
tot = AC1[i] + LM[i]
for tag, w in weights.items():
if w:
tot = tot + w * SC[tag][i]
hyps.append(cc(cands[i][int(np.argmax(tot))], greedy[i]))
return comb(refs, hyps)[2]
print("\n--- rescoreurs pris un par un ---", flush=True)
solo = {}
for tag in SC:
bb = (9.0, 0.0)
for w in (0.3, 0.5, 1.0, 1.5, 2.5):
m = evaluate({tag: w})
if m < bb[0]:
bb = (m, w)
solo[tag] = bb
print(" %-9s meilleur %.4f (w=%.1f) vs ref %.4f : %+.4f"
% (tag, bb[0], bb[1], REF, bb[0] - REF), flush=True)
print("\n--- montee de coordonnees (combinaison) ---", flush=True)
W = {t: 0.0 for t in SC}
order = sorted(solo, key=lambda t: solo[t][0])
for t in order:
W[t] = solo[t][1]
cur = evaluate(W)
print(" depart (chacun a son optimum solo): %.4f %s" % (cur, W), flush=True)
for it in range(3):
improved = False
for t in order:
base_w = W[t]
for w in (0.0, 0.3, 0.5, 1.0, 1.5, 2.5):
W[t] = w
m = evaluate(W)
if m < cur - 1e-6:
cur, base_w, improved = m, w, True
W[t] = base_w
print(" passe %d: %.4f %s" % (it + 1, cur, {k: v for k, v in W.items() if v}), flush=True)
if not improved:
break
print("\nBEST_MULTI %.4f poids=%s (ref %.4f, gain %+.4f)"
% (cur, {k: v for k, v in W.items() if v}, REF, cur - REF), flush=True)
json.dump({"combine": cur, "weights": W, "ref": REF, "solo": {k: list(v) for k, v in solo.items()}},
open("/root/multi_rescore.json", "w"))
print("MULTI_DONE", flush=True)
if __name__ == "__main__":
main()