|
|
| """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 (
|
| FONBENCH_EVAL_VERSION,
|
| accumulate,
|
| finalize,
|
| new_counters,
|
| )
|
|
|
|
|
|
|
| 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)
|
|
|
| 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
|
|
|
|
|
| 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
|
|
|
|
|
| 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:
|
| 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
|
|
|
|
|
| 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:
|
| 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())
|
|
|