FonBench / repro /evaluate.py
Kimyayd's picture
Publish the reproduction kit: scoring code, standalone evaluator, training script
d9983a9 verified
Raw
History Blame Contribute Delete
10.2 kB
#!/usr/bin/env python3
"""FonBench — évaluation reproductible, sans dépendance à notre infrastructure.
Ce script produit exactement les mêmes chiffres que le classement. Il ne
parle à aucun serveur, ne lit aucune clé, et n'a besoin que de Hugging Face.
C'est volontaire : un leaderboard qu'on ne peut pas vérifier ne vaut rien.
pip install torch transformers datasets av jiwer
python evaluate.py --model chrisjay/fonxlsr --dataset alaleye/fon
Le corpus de test principal est privé (il perdrait sa valeur en circulant),
mais il reste vérifiable : qui obtient l'accès au dépôt `JMLdata/fon-test-v1`
peut rejouer n'importe quelle ligne du classement avec ce script et doit
retrouver les mêmes chiffres, à la décimale. Passer le jeton par --token ou
par la variable HF_TOKEN. Le README.md de ce dossier donne des valeurs de
référence à retrouver exactement.
"""
from __future__ import annotations
import argparse
import io
import json
import os
import sys
import time
import numpy as np
import torch
from datasets import Audio, load_dataset
sys.path.insert(0, ".")
sys.path.insert(0, "..")
from fonbench_eval import ( # noqa: E402
FONBENCH_EVAL_VERSION,
accumulate,
finalize,
new_counters,
)
# Architectures autorégressives : découpage à 30 s obligatoire, sinon les
# énoncés longs font échouer Whisper sur une incompatibilité de dimensions.
SEQ2SEQ_TYPES = {
"whisper", "speech_to_text", "speech-encoder-decoder",
"speech_encoder_decoder", "seamless_m4t", "seamless_m4t_v2",
}
def decode_audio(cell, target_sr: int = 16000):
"""Décode un audio en mono float32 16 kHz.
PyAV plutôt que soundfile : les corpus fongbe contiennent des conteneurs
que libsndfile ne sait pas ouvrir (WebM, Opus) et des fichiers dont
l'en-tête le fait échouer (« array is too big »).
"""
import av
raw = cell["bytes"] if isinstance(cell, dict) else cell
if raw is None and isinstance(cell, dict) and cell.get("path"):
with open(cell["path"], "rb") as f:
raw = f.read()
with av.open(io.BytesIO(raw)) as container:
stream = container.streams.audio[0]
resampler = av.audio.resampler.AudioResampler(
format="flt", layout="mono", rate=target_sr
)
chunks: list = []
def _emit(frame):
res = resampler.resample(frame)
for rf in res if isinstance(res, list) else ([res] if res else []):
chunks.append(rf.to_ndarray().reshape(-1))
for frame in container.decode(stream):
_emit(frame)
_emit(None) # flush du resampler
if not chunks:
return np.zeros(1, dtype="float32")
return np.concatenate(chunks).astype("float32")
class Transcriber:
"""Charge un modèle du Hub et transcrit des tableaux 16 kHz."""
def __init__(self, model_id: str, revision: str, device: str):
from transformers import AutoConfig
self.device = device
# Jamais True, quelle que soit l'erreur rencontrée : c'est ce qui
# garantit qu'aucun code du dépôt évalué ne s'exécute.
cfg = AutoConfig.from_pretrained(model_id, revision=revision,
trust_remote_code=False)
self.architecture = cfg.model_type
self.seq2seq = cfg.model_type in SEQ2SEQ_TYPES or any(
"ConditionalGeneration" in a or "Seq2Seq" in a
for a in (getattr(cfg, "architectures", None) or [])
)
self.decoder_type = "encoder-decoder" if self.seq2seq else "ctc"
if self.seq2seq:
from transformers import (AutoModelForSpeechSeq2Seq, AutoProcessor,
pipeline)
self.processor = AutoProcessor.from_pretrained(
model_id, revision=revision, trust_remote_code=False)
self.model = AutoModelForSpeechSeq2Seq.from_pretrained(
model_id, revision=revision, trust_remote_code=False,
).to(device).eval()
self.pipe = pipeline(
"automatic-speech-recognition", model=self.model,
tokenizer=self.processor.tokenizer,
feature_extractor=self.processor.feature_extractor,
chunk_length_s=30, device=device,
)
else:
from transformers import AutoModelForCTC, AutoProcessor
self.processor = AutoProcessor.from_pretrained(
model_id, revision=revision, trust_remote_code=False)
self.model = AutoModelForCTC.from_pretrained(
model_id, revision=revision, trust_remote_code=False,
).to(device).eval()
self.pipe = None
# MMS multilingue : les poids fongbe vivent dans un adaptateur
# séparé, sans quoi le modèle transcrit une autre langue.
if getattr(self.model.config, "adapter_attn_dim", None):
try:
self.model.load_adapter("fon")
self.processor.tokenizer.set_target_lang("fon")
print(" adaptateur MMS « fon » chargé")
except Exception as exc: # noqa: BLE001
print(f" pas d'adaptateur fon ({type(exc).__name__})")
self.model_params = sum(p.numel() for p in self.model.parameters())
def __call__(self, arrays: list, batch: int) -> list[str]:
if self.seq2seq:
return [(self.pipe(a)["text"] or "").strip() for a in arrays]
out: list[str] = []
for i in range(0, len(arrays), batch):
inputs = self.processor(arrays[i:i + batch], sampling_rate=16000,
return_tensors="pt", padding=True)
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.inference_mode():
logits = self.model(**inputs).logits
ids = torch.argmax(logits, dim=-1).cpu()
out.extend(t.strip() for t in self.processor.batch_decode(ids))
return out
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--model", required=True, help="dépôt HF du modèle")
ap.add_argument("--dataset", required=True, help="dépôt HF du corpus")
ap.add_argument("--split", default="test")
ap.add_argument("--revision", default=None,
help="révision du CORPUS (fige le résultat)")
ap.add_argument("--model-revision", default=None,
help="révision du MODÈLE (défaut : dernière)")
ap.add_argument("--text-column", default=None,
help="colonne de référence (auto-détectée sinon)")
ap.add_argument("--limit", type=int, default=0,
help="n'évaluer que les N premiers énoncés")
ap.add_argument("--batch", type=int, default=8)
ap.add_argument("--token", default=None,
help="jeton HF, pour un corpus à accès restreint "
"(défaut : variable d'environnement HF_TOKEN)")
ap.add_argument("--device", default="cuda" if torch.cuda.is_available()
else "cpu")
args = ap.parse_args()
from huggingface_hub import HfApi
token = args.token or os.environ.get("HF_TOKEN") or None
# Le jeton ne sert QU'au corpus : le modèle évalué doit être public,
# sans quoi personne d'autre ne pourrait refaire la mesure.
sha = args.model_revision or HfApi(token=False).model_info(args.model).sha
print(f"modèle : {args.model} @ {sha[:12]}")
print(f"corpus : {args.dataset} [{args.split}]"
+ (f" @ {args.revision[:12]}" if args.revision else ""))
ds = load_dataset(args.dataset, split=args.split, revision=args.revision,
token=token)
if args.limit:
ds = ds.select(range(min(args.limit, len(ds))))
col = args.text_column
if col is None:
for candidat in ("transcription", "sentence", "text", "target"):
if candidat in ds.column_names:
col = candidat
break
if col is None:
sys.exit(f"colonne de référence introuvable parmi {ds.column_names} — "
"préciser --text-column")
print(f"colonne : {col} | {len(ds)} énoncés | {args.device}")
ds = ds.cast_column("audio", Audio(decode=False))
tr = Transcriber(args.model, sha, args.device)
counters = new_counters()
audio_s = compute_s = 0.0
skipped = 0
CHUNK = 50
for i in range(0, len(ds), CHUNK):
rows = ds.select(range(i, min(i + CHUNK, len(ds))))
arrays, refs = [], []
for row in rows:
try:
arrays.append(decode_audio(row["audio"]))
refs.append(row[col])
except Exception: # noqa: BLE001 — énoncé illisible, pas la faute du modèle
skipped += 1
if not arrays:
continue
t0 = time.time()
hyps = tr(arrays, args.batch)
compute_s += time.time() - t0
audio_s += sum(len(a) for a in arrays) / 16000.0
accumulate(counters, refs, hyps)
print(f" {min(i + CHUNK, len(ds))}/{len(ds)}", end="\r", flush=True)
if skipped:
print(f"\n {skipped} énoncés illisibles écartés")
metrics = finalize(counters)
metrics.update({
"model_id": args.model,
"model_revision": sha,
"dataset": f"{args.dataset}[{args.split}]",
"architecture": tr.architecture,
"decoder_type": tr.decoder_type,
"model_params": tr.model_params,
"rtfx": round(audio_s / compute_s, 3) if compute_s else None,
"eval_seconds": round(compute_s, 1),
"device": args.device,
"eval_version": FONBENCH_EVAL_VERSION,
})
print()
print(json.dumps(metrics, indent=2, ensure_ascii=False))
return 0
if __name__ == "__main__":
raise SystemExit(main())