waxal2026-backup / scripts /mms_lid_check.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
2.97 kB
#!/usr/bin/env python3
"""Arbitre LID indépendant : facebook/mms-lid-256 restreint à {lin,lug,sna}.
Compare sa distribution à celle de notre lid_best sur un échantillon Phase 2."""
import glob
import json
import numpy as np
import soundfile as sf
import torch
from transformers import (AutoFeatureExtractor, Wav2Vec2ForSequenceClassification,
Wav2Vec2BertForSequenceClassification, SeamlessM4TFeatureExtractor)
SR = 16000
N = 200
files = sorted(glob.glob("/root/phase2_audio/audio/*.wav"))[:N]
# --- MMS-LID-256 (arbitre externe) ---
fe = AutoFeatureExtractor.from_pretrained("facebook/mms-lid-256")
mms = Wav2Vec2ForSequenceClassification.from_pretrained("facebook/mms-lid-256", torch_dtype=torch.bfloat16).cuda().eval()
# label2id est None sur mms-lid-256 -> reconstruire depuis id2label (codes ISO)
i2l = mms.config.id2label
l2i = {str(v): int(k) for k, v in i2l.items()}
targets = {code: l2i[code] for code in ["lin", "lug", "sna"] if code in l2i}
print("indices MMS-LID pour lin/lug/sna:", targets)
assert len(targets) == 3, f"langues manquantes dans mms-lid-256: {set(['lin','lug','sna'])-set(targets)}"
idx = torch.tensor([targets["lin"], targets["lug"], targets["sna"]]).cuda()
order = ["lin", "lug", "sna"]
mms_pred = {}
with torch.inference_mode():
for f in files:
w = sf.read(f, dtype="float32")[0][:20 * SR]
x = fe([w], sampling_rate=SR, return_tensors="pt")
x = {k: v.to("cuda", dtype=torch.bfloat16 if v.dtype == torch.float32 else v.dtype) for k, v in x.items()}
logits = mms(**x).logits.float()[0]
sub = logits[idx] # restreint aux 3 langues
mms_pred[f] = order[int(sub.argmax())]
del mms; torch.cuda.empty_cache()
# --- notre lid_best ---
fe2 = SeamlessM4TFeatureExtractor.from_pretrained("/root/models/lid_best")
lid = Wav2Vec2BertForSequenceClassification.from_pretrained("/root/models/lid_best", torch_dtype=torch.bfloat16).cuda().eval()
id2label = json.load(open("/root/models/lid_best/config.json"))["id2label"]
our_pred = {}
with torch.inference_mode():
for f in files:
w = sf.read(f, dtype="float32")[0][:20 * SR]
x = fe2([w], 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()}
our_pred[f] = id2label[str(int(lid(**x).logits.float().argmax(-1)))]
from collections import Counter
print(f"\nN={len(files)}")
print(f"MMS-LID-256 (externe) distribution : {dict(Counter(mms_pred.values()))}")
print(f"notre lid_best distribution : {dict(Counter(our_pred.values()))}")
agree = sum(1 for f in files if mms_pred[f] == our_pred[f])
print(f"accord MMS-LID / notre LID : {agree}/{len(files)} ({100*agree//len(files)}%)")
lugs = [f for f in files if our_pred[f] == "lug"]
print(f"\nParmi les {len(lugs)} clips que NOTRE LID dit 'lug', MMS-LID dit: {dict(Counter(mms_pred[f] for f in lugs))}")
print("LIDCHECK_DONE")