#!/usr/bin/env python3 """Pipeline Phase 2 : dossier d audios anonymes -> LID -> routage -> transcription -> CSV. Usage: phase2_run.py --audio_dir DIR --out sub.csv [--ids ID1.csv]""" import argparse, csv, glob, json, os, subprocess import numpy as np, soundfile as sf, torch from transformers import (AutoModelForCTC, AutoProcessor, Wav2Vec2BertForSequenceClassification, SeamlessM4TFeatureExtractor) SR = 16000 # lin = MMS-1B adaptateurs (GO) ; sna/lug = w2v-BERT champions. AutoModelForCTC gere les 2 archis. MODELS = {"lin": "/root/models/lin_s4_best", "lug": "/root/models/lug_ps_best", "sna": "/root/models/sna_ps_best"} LANGS = ["lin", "lug", "sna"] def to_wav(path): if path.endswith((".flac", ".wav")): wav, sr = sf.read(path, dtype="float32") if sr == SR: return wav p = subprocess.run(["ffmpeg", "-v", "error", "-i", path, "-f", "f32le", "-ac", "1", "-ar", str(SR), "pipe:1"], capture_output=True) return np.frombuffer(p.stdout, dtype=np.float32) def main(): ap = argparse.ArgumentParser() ap.add_argument("--audio_dir", required=True) ap.add_argument("--out", required=True) a = ap.parse_args() files = sorted(glob.glob(os.path.join(a.audio_dir, "*"))) files = [f for f in files if f.lower().endswith((".mp3", ".wav", ".flac", ".ogg", ".m4a"))] print(f"{len(files)} fichiers audio", flush=True) fe = SeamlessM4TFeatureExtractor.from_pretrained("/root/models/lid_best") lid = Wav2Vec2BertForSequenceClassification.from_pretrained("/root/models/lid_best", torch_dtype=torch.bfloat16).cuda().eval() routing, durs = {}, {} with torch.inference_mode(): for i in range(0, len(files), 16): batch = files[i:i+16] wavs = [to_wav(f) for f in batch] for f, w in zip(batch, wavs): durs[f] = len(w) / SR crops = [w[:20*SR] for w in wavs] feats = fe(crops, 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()} preds = lid(**feats).logits.float().argmax(-1).cpu().numpy() for f, p in zip(batch, preds): routing[f] = LANGS[int(p)] if (i // 16) % 30 == 0: print(f"LID {i+len(batch)}/{len(files)}", flush=True) del lid; torch.cuda.empty_cache() counts = {l: sum(1 for v in routing.values() if v == l) for l in LANGS} print("routage:", counts, flush=True) hyps = {} for lang in LANGS: group = [f for f in files if routing[f] == lang] if not group: continue proc = AutoProcessor.from_pretrained(MODELS[lang]) m = AutoModelForCTC.from_pretrained(MODELS[lang], torch_dtype=torch.bfloat16).cuda().eval() group.sort(key=lambda f: -durs[f]) batches, cur, bud = [], [], 0.0 for f in group: if cur and bud + durs[f] > 140.0: batches.append(cur); cur, bud = [], 0.0 cur.append(f); bud += durs[f] if cur: batches.append(cur) with torch.inference_mode(): for j, b in enumerate(batches): wavs = [to_wav(f) for f in b] feats = proc.feature_extractor(wavs, 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()} ids = m(**feats).logits.float().argmax(-1).cpu().numpy() for f, s in zip(b, proc.tokenizer.batch_decode(ids)): hyps[f] = " ".join(s.replace("|", " ").split()) or "a" if (j + 1) % 20 == 0: print(f"{lang} {j+1}/{len(batches)} batches", flush=True) del m; torch.cuda.empty_cache() print(f"{lang}: {len(group)} clips transcrits", flush=True) with open(a.out, "w", newline="", encoding="utf-8") as fo: w = csv.writer(fo); w.writerow(["ID", "Target"]) for f in files: w.writerow([os.path.splitext(os.path.basename(f))[0], hyps.get(f) or "a"]) print(f"PHASE2_DONE {a.out} ({len(files)} lignes)", flush=True) if __name__ == "__main__": main()