File size: 10,173 Bytes
d9983a9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 | #!/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())
|