File size: 5,148 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 | #!/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")
|