File size: 6,565 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
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
#!/usr/bin/env python3
"""BANC D'ESSAI LINGALA sur la validation officielle (1844 clips, 0% de fuite de
clips verifiee ; KenLM verifie propre a 0.1%). devhard est INUTILISABLE (100% de
ses clips sont dans le train).

But n1 : CALIBRER L'INSTRUMENT. Trois modeles ont un verdict LB connu :
   joint_cont  = CHAMPION
   lin_s4      = LB -0.0072  (doit finir DERRIERE joint_cont)
   joint_linps = LB -0.0062  (doit finir DERRIERE joint_cont)
Si la validation ne reproduit pas cet ordre, elle ne sert pas non plus a
selectionner un modele, et il faudra reentrainer une baseline B0 sur
train_lin_min pour retrouver un gate honnete.

But n2 : donner leur premiere chance equitable aux modeles jamais retenus
(mms1b_lin, mmsjoint, mmsjointpl, joint_cont2, joint_best) -- ils avaient ete
ecartes sur devhard, terrain ou nos propres modeles etaient gonfles par la
memorisation et eux non.
"""
import json, os, sys, time
import numpy as np, soundfile as sf, torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoModelForCTC, AutoProcessor

VAL = "/scratch/prep/manifests/waxal_lin_validation.jsonl"
ARPA = "/scratch/lm/lin_5g.arpa"
ALPHA, BETA, LSB, BW = 0.5, 1.0, True, 64

MODELS = [
    ("joint_cont   [CHAMPION]", "/root/models/joint_cont_best"),
    ("lin_s4       [LB -0.0072]", "/root/models/lin_s4_best"),
    ("joint_linps  [LB -0.0062]", "/scratch/runs/joint_linps/best"),
    ("joint_cont2", "/root/models/joint_cont2_best"),
    ("joint_best", "/root/models/joint_best"),
    ("mms1b_lin", "/root/models/mms1b_lin_best"),
    ("mmsjoint", "/root/models/mmsjoint_best"),
    ("mmsjointpl", "/root/models/mmsjointpl_best"),
]


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):
    we = wl = ce = cl = 0
    per = []
    for h, r in zip(hyps, refs):
        hw, rw = h.split(), r.split()
        e1, e2 = lev(hw, rw), lev(h, r)
        we += e1; wl += len(rw); ce += e2; cl += len(r)
        per.append(0.5 * e1 / max(len(rw), 1) + 0.5 * e2 / max(len(r), 1))
    wer, cer = we / max(wl, 1), ce / max(cl, 1)
    return wer, cer, 0.5 * wer + 0.5 * cer, np.array(per)


def norm(s):
    return " ".join(str(s).replace("|", " ").split())


def logits_for(model_dir, files):
    proc = AutoProcessor.from_pretrained(model_dir)
    m = AutoModelForCTC.from_pretrained(model_dir, dtype=torch.float32).cuda().eval()
    out, bs = [], 4
    with torch.inference_mode():
        i = 0
        while i < len(files):
            b = files[i:i + bs]
            try:
                au = [sf.read(f, dtype="float32")[0] for f in b]
                x = proc(au, sampling_rate=16000, return_tensors="pt", padding=True)
                x = {k: v.cuda() for k, v in x.items() if k in ("input_values", "input_features", "attention_mask")}
                lg = m(**x).logits.log_softmax(-1).float().cpu().numpy()
                for j in range(len(b)):
                    out.append(lg[j])
                i += bs
            except torch.cuda.OutOfMemoryError:
                torch.cuda.empty_cache()
                if bs == 1:
                    raise
                bs = max(1, bs // 2)
                print("      OOM -> batch %d" % bs, flush=True)
    del m; torch.cuda.empty_cache()
    return proc, out


rows = [json.loads(l) for l in open(VAL, encoding="utf-8")]
rows = [r for r in rows if os.path.exists(r["audio"])]
files = [r["audio"] for r in rows]
refs = [r["text"] for r in rows]
print("validation lingala : %d clips (0%% de fuite de clips, verifie)" % len(files), flush=True)
print("decodage KenLM : alpha=%.1f beta=%.1f lsb=%s beam=%d  (config exacte du champion)\n"
      % (ALPHA, BETA, LSB, BW), flush=True)

RES = {}
for nom, path in MODELS:
    if not os.path.isdir(path):
        print("%-26s  ABSENT (%s)" % (nom, path), flush=True); continue
    t0 = time.time()
    try:
        proc, L = logits_for(path, 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] = ""
        g = [norm(tok.decode(l.argmax(-1))) for l in L]
        wg, cg, kg, perg = score(g, refs)
        dec = build_ctcdecoder(lab, kenlm_model_path=ARPA, alpha=ALPHA, beta=BETA,
                               lm_score_boundary=LSB)
        with Pool(8) as p:
            hl = [norm(x) for x in dec.decode_batch(p, L, beam_width=BW)]
        wl_, cl_, kl, perl = score(hl, refs)
        RES[nom] = dict(greedy=(wg, cg, kg), kenlm=(wl_, cl_, kl), perg=perg, perl=perl)
        print("%-26s  greedy %.4f | KenLM %.4f   (%.0f s)" % (nom, kg, kl, time.time() - t0), flush=True)
        del L
    except Exception as e:
        print("%-26s  ECHEC : %s" % (nom, type(e).__name__ + " " + str(e)[:120]), flush=True)
    torch.cuda.empty_cache()

print("\n============== CLASSEMENT (validation lingala, combine, plus bas = mieux) ==============")
print("%-26s %8s %8s %9s   %8s %8s %9s" % ("", "WERg", "CERg", "GREEDY", "WERlm", "CERlm", "KenLM"))
for nom in sorted(RES, key=lambda k: RES[k]["kenlm"][2]):
    g, l = RES[nom]["greedy"], RES[nom]["kenlm"]
    print("%-26s %8.4f %8.4f %9.4f   %8.4f %8.4f %9.4f" % (nom, g[0], g[1], g[2], l[0], l[1], l[2]))

ref_key = [k for k in RES if k.startswith("joint_cont ")]
if ref_key:
    rk = ref_key[0]
    base = RES[rk]["perl"]
    rng = np.random.default_rng(0)
    print("\n--- comparaison appariee vs CHAMPION, decodage KenLM (negatif = bat le champion) ---")
    for nom in sorted(RES, key=lambda k: RES[k]["kenlm"][2]):
        if nom == rk:
            continue
        d = RES[nom]["perl"] - base
        bs = np.array([d[rng.integers(0, len(d), len(d))].mean() for _ in range(3000)])
        lo, hi = np.percentile(bs, 2.5), np.percentile(bs, 97.5)
        verdict = "BAT LE CHAMPION" if hi < 0 else ("perd" if lo > 0 else "dans le bruit")
        print("%-26s delta %+.5f  IC95 [%+.5f,%+.5f]  %s" % (nom, d.mean(), lo, hi, verdict))

print("\n>>> CALIBRATION : lin_s4 et joint_linps DOIVENT finir derriere joint_cont.")
print(">>> Si ce n'est pas le cas, la validation ne sert pas a selectionner un modele.")
print("BENCH_LIN_DONE")