File size: 6,094 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
#!/usr/bin/env python3
"""Etape 1/2 : produit les listes N-BEST des deux langues, avec leurs scores
actuels, pour pouvoir les re-classer ensuite par un LM MASQUE BIDIRECTIONNEL.

Pourquoi cette piste : arXiv 2606.23306 teste 11 methodes INTERNES au CTC pour
recuperer l'ecart a l'oracle (MBR sur posterieurs, MC-dropout, decodage
contrastif, rescoreurs entraines) -- aucune n'est significative. La seule qui
marche est un posterieur issu d'un MLM bidirectionnel EXTERNE. Cela explique nos
propres refutations (moyennage de logits, ensembles medoide, LM neuronal causal).
Et notre plus gros ecart a l'oracle est cote shona (oracle 10-best 0.0949 vs
0.1281 retenu) : le goulot y est la SELECTION, exactement ce qu'un rescoreur corrige.

lin : beam KenLM alpha=0.6 beta=1.0 lsb=True beam=64  (config du RECORD)
sna : beam CTC pur 256 + rescorage sna_r2 w=4.0        (config du RECORD)
Les scores sauves sont EXACTEMENT ceux qui produisent le record, pour que
lambda=0 reproduise le record bit pour bit (controle de non-regression).
"""
import csv, json, os, pickle, sys
import numpy as np, soundfile as sf, torch
from multiprocessing import Pool
from pyctcdecode import build_ctcdecoder
from transformers import AutoModelForCTC, AutoProcessor

sys.path.insert(0, "/root")
from gen_sna_rescore import ctc_score, encode_for, norm

BASE = "/root/sub_LMA06.csv"
NBEST = 24
OUT = "/scratch/nbest_all.pkl"

base = {r["ID"]: r["Target"] for r in csv.DictReader(open(BASE, encoding="utf-8"))}
lang = json.load(open("/root/test_lang.json"))
lin_ids = [k for k in base if lang.get(k) == "lin"]
sna_ids = [k for k in base if lang.get(k) == "sna"]
print("lin %d | sna %d" % (len(lin_ids), len(sna_ids)), flush=True)

store = {}

# ---------------------------------------------------------------- LINGALA
with open("/scratch/lm/logits_test_lin.pkl", "rb") as f:
    keys, logs, greedy = pickle.load(f)
tokL = AutoProcessor.from_pretrained("/root/models/joint_cont_best").tokenizer
v = tokL.get_vocab()
labL = [None] * len(v)
for t, i in v.items():
    labL[i] = t
labL[tokL.word_delimiter_token_id] = " "
labL[tokL.unk_token_id] = "⁇"
labL[tokL.pad_token_id] = ""
dec = build_ctcdecoder(labL, kenlm_model_path="/scratch/lm/lin_5g.arpa",
                       alpha=0.6, beta=1.0, lm_score_boundary=True)
with Pool(8) as p:
    beams = dec.decode_beams_batch(p, logs, beam_width=64, prune_history=True)  # comme decode(), sinon le top ne reproduit pas le record
gmap = dict(zip(keys, greedy))
lin = {}
for k, bs in zip(keys, beams):
    cands, scores = [], []
    for b in bs[:NBEST]:
        h = norm(b[0])
        g = gmap[k]
        if h and g:
            h = g[:1] + h[1:]        # casse du 1er caractere, comme dans le record
        if h and h not in cands:
            cands.append(h); scores.append(float(b[3]))   # batch : b[3] = score combine (acoustique + alpha.LM + beta.mots)
    lin[k] = (cands, scores)
n_ok = sum(1 for k in lin if lin[k][0] and lin[k][0][int(np.argmax(lin[k][1]))] == base[k])
print("lin n-best OK | argmax reproduit le record sur %d/%d clips" % (n_ok, len(lin)), flush=True)
store["lin"] = lin
del logs, beams, dec

# ---------------------------------------------------------------- SHONA
files = [os.path.join("/scratch/p2_16k", k + ".wav") for k in sna_ids]
proc = AutoProcessor.from_pretrained("/root/models/sna_ps_best")
m = AutoModelForCTC.from_pretrained("/root/models/sna_ps_best", dtype=torch.float32).cuda().eval()
L1 = []
with torch.inference_mode():
    for i in range(0, len(files), 4):
        au = [sf.read(f, dtype="float32")[0] for f in files[i:i + 4]]
        x = proc(au, sampling_rate=16000, return_tensors="pt", padding=True)
        x = {kk: vv.cuda() for kk, vv in x.items()}
        lg = m(**x).logits.log_softmax(-1).float().cpu().numpy()
        for j in range(len(au)):
            L1.append(lg[j])
del m; torch.cuda.empty_cache()
tokS = proc.tokenizer
v = tokS.get_vocab()
labS = [None] * len(v)
for t, i in v.items():
    labS[i] = t
labS[tokS.word_delimiter_token_id] = " "
labS[tokS.unk_token_id] = "⁇"
labS[tokS.pad_token_id] = ""
greedyS = [norm(tokS.decode(l.argmax(-1))) for l in L1]
decS = build_ctcdecoder(labS)                 # CTC pur, aucun LM (config du record)
with Pool(8) as p:
    allb = decS.decode_beams_batch(p, L1, beam_width=256)
with Pool(8) as p:
    dbS = [norm(x) for x in decS.decode_batch(p, L1, beam_width=256)]
print("beam sna OK", flush=True)

procR = AutoProcessor.from_pretrained("/scratch/restore/sna_r2_best")
mR = AutoModelForCTC.from_pretrained("/scratch/restore/sna_r2_best", dtype=torch.float32).cuda().eval()
LG = []
with torch.inference_mode():
    for i in range(0, len(files), 4):
        au = [sf.read(f, dtype="float32")[0] for f in files[i:i + 4]]
        x = procR(au, sampling_rate=16000, return_tensors="pt", padding=True)
        x = {kk: vv.cuda() for kk, vv in x.items()}
        lg = mR(**x).logits.log_softmax(-1).float().cpu().numpy()
        for j in range(len(au)):
            LG.append(lg[j])
del mR; torch.cuda.empty_cache()
t2 = procR.tokenizer
print("rescoreur sna OK", flush=True)

sna = {}
for i, k in enumerate(sna_ids):
    bs = allb[i]
    c = [norm(b[0]) for b in bs[:50]]
    a = [(b[2] if len(b) > 3 else 0.0) for b in bs[:50]]   # sna sans LM : b[2] = acoustique
    for extra in (dbS[i], greedyS[i]):
        if extra and extra not in c:
            c.append(extra)
            a.append(ctc_score(L1[i], encode_for(tokS, extra), tokS.pad_token_id))
    resc = np.array([ctc_score(LG[i], encode_for(t2, x), t2.pad_token_id) for x in c])
    tot = np.array(a) + 4.0 * resc
    order = np.argsort(-tot)[:NBEST]
    sna[k] = ([c[j] for j in order], [float(tot[j]) for j in order])
n_ok = sum(1 for k in sna if sna[k][0] and sna[k][0][0] == base[k])
print("sna n-best OK | argmax reproduit le record sur %d/%d clips" % (n_ok, len(sna)), flush=True)
store["sna"] = sna

with open(OUT, "wb") as f:
    pickle.dump(store, f, protocol=4)
tot_h = sum(len(v[0]) for L in store.values() for v in L.values())
print("NBEST_DONE %s | %d hypotheses au total" % (OUT, tot_h), flush=True)