waxal2026-backup / code /ensemble_decorrele.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
3.78 kB
#!/usr/bin/env python3
"""PISTE 3 - Ensemble de modeles DECORRELES : moyenne des log-probs de modeles
entraines differemment (standard / robuste / pseudo), pas de la meme lignee.
Evalue chaque combinaison sur la VALIDATION COMPLETE."""
import itertools
import json
import os
import jiwer
import numpy as np
import soundfile as sf
import torch
from transformers import Wav2Vec2BertForCTC, Wav2Vec2BertProcessor
SR = 16000
MANIF = "/scratch/prep/manifests"
POOL = {
"lin": ["lin_s4_best", "lin_r_best", "lin_r2_best"],
"lug": ["lug_r_best", "lug_ps_best", "lug_r2_best", "lug_s2_best"],
"sna": ["sna_ps_best", "sna_r_best", "sna_r2_best"],
}
def read(p):
rows = []
for l in open(p, encoding="utf-8"):
r = json.loads(l)
r["text"] = " ".join(r.get("text", "").replace("|", " ").split())
rows.append(r)
return [r for r in rows if r["text"]]
@torch.inference_mode()
def logits_of(path, rows):
proc = Wav2Vec2BertProcessor.from_pretrained(path)
m = Wav2Vec2BertForCTC.from_pretrained(path, torch_dtype=torch.bfloat16).cuda().eval()
rows_s = sorted(rows, key=lambda r: -r["duration"])
bs, cur, bud = [], [], 0.0
for r in rows_s:
if cur and bud + r["duration"] > 120:
bs.append(cur); cur, bud = [], 0.0
cur.append(r); bud += r["duration"]
if cur: bs.append(cur)
out = {}
for b in bs:
au = [sf.read(r["audio"], dtype="float32")[0] for r in b]
f = proc.feature_extractor(au, sampling_rate=SR, return_tensors="pt", padding=True)
f = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in f.items()}
lg = torch.log_softmax(m(**f).logits.float(), -1).cpu().numpy()
att = f.get("attention_mask")
for j, r in enumerate(b):
T = lg.shape[1]
if att is not None:
T = max(1, int(round(lg.shape[1] * att[j].sum().item() / att.shape[1])))
out[r["id"]] = lg[j, :T].astype(np.float32)
del m
torch.cuda.empty_cache()
return out, proc
def main():
print("=== PISTE 3 : ensembles decorreles (validation complete) ===", flush=True)
for lang, names in POOL.items():
names = [n for n in names if os.path.isdir(f"/root/models/{n}")]
if len(names) < 2:
print(f"{lang}: <2 modeles dispo, saute"); continue
val = read(f"{MANIF}/waxal_{lang}_validation.jsonl")
refs = [r["text"] for r in val]
order = [r["id"] for r in val]
cache, proc = {}, None
for n in names:
cache[n], proc = logits_of(f"/root/models/{n}", val)
hy = [" ".join(proc.tokenizer.decode(np.argmax(cache[n][i], -1)).replace("|", " ").split()) for i in order]
s = 0.5 * jiwer.wer(refs, hy) + 0.5 * jiwer.cer(refs, hy)
print(f" {lang}/{n} seul: {s:.4f}", flush=True)
best = None
for k in (2, 3):
for combo in itertools.combinations(names, k):
hy = []
for i in order:
L = min(cache[n][i].shape[0] for n in combo)
avg = np.mean([cache[n][i][:L] for n in combo], axis=0)
hy.append(" ".join(proc.tokenizer.decode(np.argmax(avg, -1)).replace("|", " ").split()))
s = 0.5 * jiwer.wer(refs, hy) + 0.5 * jiwer.cer(refs, hy)
print(f" {lang}/ENS[{'+'.join(c.replace('_best','') for c in combo)}]: {s:.4f}", flush=True)
if best is None or s < best[1]:
best = (combo, s)
print(f"{lang} -> MEILLEUR ENSEMBLE: {best[0]} = {best[1]:.4f}", flush=True)
cache.clear()
torch.cuda.empty_cache()
print("ENSEMBLE_DONE", flush=True)
if __name__ == "__main__":
main()