waxal2026-backup / scripts /finale2_multi.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
2.3 kB
#!/usr/bin/env python3
"""Finale 2 = modèle MULTILINGUE joint_cont2 sur TOUS les clips Phase 2, SANS routage.
Robuste à la confusion lug/sna du LID (dégrade gracieusement)."""
import csv, glob, os
import soundfile as sf, torch
from transformers import AutoModelForCTC, AutoProcessor
SR = 16000
MODEL = "/root/models/joint_cont2_best"
AUD = "/root/phase2_audio/audio"
OUT = "/root/sub_finale2_multi.csv"
TESTCSV = "/root/Test_phase2.csv"
def norm(s):
return " ".join(str(s).replace("|", " ").split())
def main():
test_ids = [r["ID"] for r in csv.DictReader(open(TESTCSV, encoding="utf-8"))]
files = sorted(glob.glob(os.path.join(AUD, "*.wav")))
durs = {f: sf.info(f).duration for f in files}
files.sort(key=lambda f: -durs[f])
proc = AutoProcessor.from_pretrained(MODEL)
m = AutoModelForCTC.from_pretrained(MODEL, torch_dtype=torch.bfloat16).cuda().eval()
# batching par budget de durée
bs, cur, bud = [], [], 0.0
for f in files:
if cur and bud + durs[f] > 140:
bs.append(cur); cur, bud = [], 0.0
cur.append(f); bud += durs[f]
if cur:
bs.append(cur)
out = {}
with torch.inference_mode():
for j, b in enumerate(bs):
au = [sf.read(f, dtype="float32")[0] for f in b]
x = proc(au, sampling_rate=SR, return_tensors="pt", padding=True)
x = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in x.items()}
ids = m(**x).logits.float().argmax(-1).cpu().numpy()
for f, s in zip(b, proc.batch_decode(ids)):
out[os.path.splitext(os.path.basename(f))[0]] = norm(s)
if (j + 1) % 20 == 0:
print(f"{j+1}/{len(bs)} batches", flush=True)
empt = sum(1 for t in test_ids if not out.get(t, "").strip())
with open(OUT, "w", newline="", encoding="utf-8") as fo:
w = csv.writer(fo); w.writerow(["ID", "Target"])
for t in test_ids:
w.writerow([t, out.get(t) or "a"])
print(f"MULTI_DONE {OUT} | {len(test_ids)} IDs | couvre={set(test_ids)==set(out)|(set(test_ids)>=set(out))} | vides={empt}", flush=True)
# apercu
for t in test_ids[:4]:
print(f" {t}: {out.get(t,'')[:80]}", flush=True)
if __name__ == "__main__":
main()