waxal2026-backup / code /devhard_ps2.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
5.15 kB
#!/usr/bin/env python3
"""Compare sna_ps (tour 1) et sna_ps2 (tour 2) sur devhard-sna : 444 clips a
LOCUTEURS DISJOINTS du train. C'est le seul jeu ou l'apport des 295 h de
locuteurs NOUVEAUX peut se manifester -- la validation officielle partage ses
locuteurs avec le train et est restee plate (0.1743 -> 0.1736).
Pipeline identique a la soumission : beam CTC pur 256, N-best 50, rescorage
par sna_r2_best avec w=4.0. On mesure aussi le greedy, pour separer l'effet
du modele acoustique de l'effet du decodage.
"""
import json, os, sys
import numpy as np, torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
sys.path.insert(0, "/root")
from gen_sna_rescore import compute_logits, ctc_score, encode_for, norm
DEV = "/root/devhard/devhard_sna.jsonl"
M1 = "/root/models/sna_ps_best"
M2 = "/scratch/runs/sna_ps2/best"
RESC = "/scratch/restore/sna_r2_best"
W, NBEST, BW = 4.0, 50, 256
def lev(a, b):
if a == b:
return 0
if not a:
return len(b)
if not b:
return len(a)
prev = list(range(len(b) + 1))
for i, ca in enumerate(a, 1):
cur = [i]
for j, cb in enumerate(b, 1):
cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb)))
prev = cur
return prev[-1]
def score(hyps, refs, lower=False):
"""WER/CER agreges (somme des erreurs / somme des longueurs), + erreurs par clip."""
we = wl = ce = cl = 0
per = []
for h, r in zip(hyps, refs):
if lower:
h, r = h.lower(), r.lower()
hw, rw = h.split(), r.split()
e1 = lev(hw, rw)
e2 = lev(h, r)
we += e1; wl += len(rw); ce += e2; cl += len(r)
per.append(e1 / max(len(rw), 1) * 0.5 + e2 / max(len(r), 1) * 0.5)
wer, cer = we / max(wl, 1), ce / max(cl, 1)
return wer, cer, 0.5 * wer + 0.5 * cer, np.array(per)
def decode(model_dir, files, LG, t2, tag):
proc, L1 = compute_logits(model_dir, files)
tok = proc.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 = [norm(tok.decode(l.argmax(-1))) for l in L1]
print(" [%s] logits + greedy OK" % tag, flush=True)
dec = build_ctcdecoder(lab)
with Pool(8) as p:
allbeams = dec.decode_beams_batch(p, L1, beam_width=BW)
with Pool(8) as p:
db = [norm(x) for x in dec.decode_batch(p, L1, beam_width=BW)]
print(" [%s] beam %d OK" % (tag, BW), flush=True)
picked = []
for i, bs in enumerate(allbeams):
c = [norm(b[0]) for b in bs[:NBEST]]
a = [(b[3] if len(b) > 3 else 0.0) for b in bs[:NBEST]]
for extra in (db[i], greedy[i]):
if extra and extra not in c:
c.append(extra)
a.append(ctc_score(L1[i], encode_for(tok, extra), tok.pad_token_id))
sc = np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in c])
tot = np.array(a) + W * sc
picked.append(c[int(np.argmax(tot))] or greedy[i] or "a")
print(" [%s] rescorage OK" % tag, flush=True)
return greedy, picked
rows = [json.loads(l) for l in open(DEV, encoding="utf-8")]
files, refs = [], []
for r in rows:
p = r["audio"]
if not os.path.exists(p):
p = "/root/devhard_audio/%s.flac" % r["id"]
if os.path.exists(p):
files.append(p); refs.append(r["text"])
print("devhard-sna : %d clips utilisables sur %d" % (len(files), len(rows)), flush=True)
assert len(files) > 400, "trop d'audio manquant"
proc2, LG = compute_logits(RESC, files)
t2 = proc2.tokenizer
print("rescoreur %s OK" % RESC, flush=True)
g1, p1 = decode(M1, files, LG, t2, "tour1")
g2, p2 = decode(M2, files, LG, t2, "tour2")
print("\n================ devhard-sna, %d clips, LOCUTEURS DISJOINTS ================" % len(files))
print("%-34s %8s %8s %9s" % ("", "WER", "CER", "combine"))
res = {}
for nom, h in (("tour1 greedy", g1), ("tour1 pipeline complet", p1),
("tour2 greedy", g2), ("tour2 pipeline complet", p2)):
w, c, k, per = score(h, refs)
res[nom] = per
print("%-34s %8.4f %8.4f %9.4f" % (nom, w, c, k))
print("\n--- insensible a la casse (la casse compte-t-elle dans l'ecart ?) ---")
for nom, h in (("tour1 pipeline complet", p1), ("tour2 pipeline complet", p2)):
w, c, k, _ = score(h, refs, lower=True)
print("%-34s %8.4f %8.4f %9.4f" % (nom, w, c, k))
d = res["tour2 pipeline complet"] - res["tour1 pipeline complet"]
mieux = int((d < -1e-9).sum()); pire = int((d > 1e-9).sum()); egal = int((abs(d) <= 1e-9).sum())
print("\n--- par clip (pipeline complet) ---")
print("tour2 MEILLEUR sur %d clips | PIRE sur %d | identique sur %d" % (mieux, pire, egal))
print("delta moyen (tour2 - tour1) : %+.5f (negatif = tour2 gagne)" % d.mean())
rng = np.random.default_rng(0)
bs = np.array([d[rng.integers(0, len(d), len(d))].mean() for _ in range(4000)])
print("bootstrap 95%% : [%+.5f , %+.5f] | P(tour2 gagne) = %.1f%%"
% (np.percentile(bs, 2.5), np.percentile(bs, 97.5), (bs < 0).mean() * 100))
print("DEVHARD_PS2_DONE")