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")