waxal2026-backup / code /ensemble_eval.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
3.35 kB
#!/usr/bin/env python3
"""Ensembles de checkpoints par langue : moyenne des log-probs -> greedy -> score val complet."""
import json
import jiwer
import numpy as np
import soundfile as sf
import torch
from transformers import Wav2Vec2BertForCTC, Wav2Vec2BertProcessor
SR = 16000
COMBOS = {
"lin": {"single_s4": ["/root/models/lin_s4_best"],
"s2+s4": ["/root/models/lin_s2_best", "/root/models/lin_s4_best"],
"s2+s3+s4": ["/root/models/lin_s2_best", "/root/models/lin_s3_best", "/root/models/lin_s4_best"]},
"lug": {"single_s2": ["/root/models/lug_s2_best"],
"s1+s2": ["/root/models/lug_s1_best", "/root/models/lug_s2_best"]},
"sna": {"single_s2": ["/root/models/sna_s2_best"],
"s1+s2": ["/root/models/sna_s1_best", "/root/models/sna_s2_best"]},
}
def read_jsonl(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 batch_logits(model, proc, audio):
feats = proc.feature_extractor(audio, sampling_rate=SR, return_tensors="pt", padding=True)
feats = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in feats.items()}
return torch.log_softmax(model(**feats).logits.float(), dim=-1)
results = {}
for lang, combos in COMBOS.items():
val = read_jsonl(f"/scratch/prep/manifests/waxal_{lang}_validation.jsonl")
val.sort(key=lambda r: -r["duration"])
batches, cur, bud = [], [], 0.0
for r in val:
if cur and bud + r["duration"] > 120.0:
batches.append(cur); cur, bud = [], 0.0
cur.append(r); bud += r["duration"]
if cur: batches.append(cur)
model_dirs = sorted({d for dirs in combos.values() for d in dirs})
loaded = {}
for d in model_dirs:
proc = Wav2Vec2BertProcessor.from_pretrained(d)
loaded[d] = (proc, Wav2Vec2BertForCTC.from_pretrained(d, torch_dtype=torch.bfloat16).cuda().eval())
per_model_ids = {d: {} for d in model_dirs}
logits_cache = {}
for b in batches:
audio = [sf.read(r["audio"], dtype="float32")[0] for r in b]
lg = {d: batch_logits(m, p, audio) for d, (p, m) in loaded.items()}
for name, dirs in combos.items():
avg = torch.stack([lg[d] for d in dirs]).mean(0)
ids = avg.argmax(-1).cpu().numpy()
tok = loaded[dirs[0]][0].tokenizer
for r, s in zip(b, tok.batch_decode(ids)):
logits_cache.setdefault(name, {})[r["id"]] = " ".join(s.replace("|", " ").split())
refs = {r["id"]: r["text"] for r in val}
for name in combos:
order = [r["id"] for r in val]
hyps = [logits_cache[name][i] for i in order]
rr = [refs[i] for i in order]
wer, cer = jiwer.wer(rr, hyps), jiwer.cer(rr, hyps)
comb = 0.5 * wer + 0.5 * cer
results[f"{lang}/{name}"] = {"wer": round(wer, 4), "cer": round(cer, 4), "combine": round(comb, 4)}
print(f"{lang}/{name}: WER {wer:.4f} CER {cer:.4f} combine {comb:.4f}", flush=True)
for d in list(loaded):
del loaded[d]
torch.cuda.empty_cache()
json.dump(results, open("/root/models/ensemble_results.json", "w"), indent=1)
print("ENSEMBLE_DONE", flush=True)