| |
| """Inference multilingue (whole-clip greedy) sur l'audio Phase 2, checkpoint parametrable. |
| Usage: multi_infer.py --model /root/models/joint_cont_best --out /root/sub_jointcont.csv""" |
| import argparse, csv, glob, os |
| import soundfile as sf, torch |
| from transformers import AutoModelForCTC, AutoProcessor |
| SR = 16000 |
|
|
|
|
| def norm(s): |
| return " ".join(str(s).replace("|", " ").split()) |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--model", required=True) |
| ap.add_argument("--audio_dir", default="/root/phase2_audio/audio") |
| ap.add_argument("--test_csv", default="/root/Test_phase2.csv") |
| ap.add_argument("--out", required=True) |
| a = ap.parse_args() |
| test_ids = [r["ID"] for r in csv.DictReader(open(a.test_csv, encoding="utf-8"))] |
| files = sorted(glob.glob(os.path.join(a.audio_dir, "*.wav"))) |
| durs = {f: sf.info(f).duration for f in files} |
| files.sort(key=lambda f: -durs[f]) |
| proc = AutoProcessor.from_pretrained(a.model) |
| m = AutoModelForCTC.from_pretrained(a.model, torch_dtype=torch.bfloat16).cuda().eval() |
| 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.to("cuda")) 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(a.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"DONE {a.out} | {len(test_ids)} IDs | vides={empt}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|