FonBench / fonbench_eval.py
Kimyayd's picture
FonBench : leaderboard + evaluateur CPU autonome
e70a0ad verified
Raw
History Blame Contribute Delete
9.26 kB
"""FonBench — normalisation du texte fongbe et métriques ASR.
Copie conforme du module de référence utilisé par le worker GPU : les
scores publiés par le Space doivent être comparables au chiffre près à
ceux déjà au classement.
Métriques (un seul passage) :
· Qualité : WER, CER, MER, WIL (jiwer)
· Segmentale : WER_seg / CER_seg — texte dé-tonalisé (accents retirés).
Non biaisée : comparable même entre corpus qui ne marquent
pas les tons (ALFFA n'en met aucune, JML les met toutes).
· Tonale : WER_ton — taux d'erreur sur les seules marques tonales.
None si le corpus n'annote pas les tons (sinon biaisé).
· Phare : T-WER = WER_seg + 2·WER_ton (double pénalité tonale).
S'y ajoute ici, par rapport au module du worker, une variante **par
tranches** (`accumulate` / `finalize`). Une évaluation sur CPU dure des
heures et le Space redémarre : il faut pouvoir reprendre. Mais stocker les
transcriptions déjà produites reviendrait à recopier le corpus privé hors
du Space. On ne garde donc que des compteurs d'erreurs, dont la somme
redonne *exactement* les mêmes scores — les alignements jiwer étant
indépendants d'un énoncé à l'autre, c'est une identité, pas une
approximation. `test_accumulation.py` le vérifie.
"""
from __future__ import annotations
import re
import unicodedata
import jiwer
FONBENCH_EVAL_VERSION = "0.3.0"
_WHITESPACE = re.compile(r"\s+")
# Voyelles fongbe (base, après minuscules/NFD).
_VOWELS = set("aeiouɛɔ")
# Marques tonales combinantes → symbole de ton.
_TONE_MARKS = {
"́": "H", # accent aigu — ton haut
"̀": "L", # accent grave — ton bas
"̌": "R", # caron — ton montant
"̂": "F", # circonflexe — ton descendant
"̄": "M", # macron — ton moyen
}
# Seuil : en dessous, le corpus n'annote pas vraiment les tons → WER_ton biaisé.
_TONE_ANNOTATION_THRESHOLD = 0.10
def normalize_fon(text: str) -> str:
"""Normalise un texte fongbe avant le calcul des métriques."""
text = unicodedata.normalize("NFC", text)
text = text.lower()
text = "".join(
ch
for ch in text
if unicodedata.category(ch)[0] in ("L", "M", "N") or ch.isspace()
)
return _WHITESPACE.sub(" ", text).strip()
def strip_tones(text: str) -> str:
"""Retire les marques tonales, garde les lettres fongbe.
ɖ, ɛ, ɔ, ŋ… sont des caractères atomiques (non décomposables) et restent ;
seuls les accents combinants (á→a, ɔ́→ɔ, ě→e) tombent.
"""
decomposed = unicodedata.normalize("NFD", text)
without_marks = "".join(
ch for ch in decomposed if unicodedata.category(ch) != "Mn"
)
return unicodedata.normalize("NFC", without_marks)
def tone_sequence(text: str) -> tuple[list[str], int]:
"""Séquence de tons (un par voyelle) + nombre de voyelles marquées."""
d = unicodedata.normalize("NFD", text)
seq: list[str] = []
marked = 0
i = 0
while i < len(d):
ch = d[i]
if ch in _VOWELS:
tone = "."
j = i + 1
while j < len(d) and unicodedata.category(d[j]) == "Mn":
if d[j] in _TONE_MARKS:
tone = _TONE_MARKS[d[j]]
j += 1
if tone != ".":
marked += 1
seq.append(tone)
i = j
else:
i += 1
return seq, marked
def levenshtein(a: list[str], b: list[str]) -> int:
"""Distance d'édition entre deux séquences (espace linéaire)."""
if not a:
return len(b)
if not b:
return len(a)
prev = list(range(len(b) + 1))
for i, ca in enumerate(a, 1):
cur = [i]
for j, cb in enumerate(b, 1):
cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb)))
prev = cur
return prev[-1]
# --- calcul en un passage (référence) ----------------------------------
def compute_metrics(references: list[str], hypotheses: list[str]) -> dict:
"""Toutes les métriques FonBench, sur des listes parallèles réf/hyp."""
counters = new_counters()
accumulate(counters, references, hypotheses)
if not counters["n_scored"]:
raise ValueError("aucune référence non vide après normalisation")
return finalize(counters)
# --- calcul par tranches (reprise après redémarrage) --------------------
COUNTER_KEYS = (
"w_s", "w_d", "w_i", "w_h", # mots, tons compris
"c_s", "c_d", "c_i", "c_h", # caractères, tons compris
"ws_s", "ws_d", "ws_i", "ws_h", # mots dé-tonalisés
"cs_s", "cs_d", "cs_i", "cs_h", # caractères dé-tonalisés
"tone_dist", "tone_units", "tone_marked",
"n_scored",
)
def new_counters() -> dict:
return {k: 0 for k in COUNTER_KEYS}
def accumulate(counters: dict, references: list[str], hypotheses: list[str]) -> dict:
"""Ajoute une tranche aux compteurs. Modifie et renvoie `counters`."""
if len(references) != len(hypotheses):
raise ValueError(
f"références ({len(references)}) et hypothèses ({len(hypotheses)}) "
"doivent avoir la même longueur"
)
pairs = [
(normalize_fon(ref), normalize_fon(hyp))
for ref, hyp in zip(references, hypotheses)
]
# Une référence vide rendrait le WER indéfini : on écarte la paire.
pairs = [(r, h) for r, h in pairs if r]
if not pairs:
return counters
refs = [r for r, _ in pairs]
hyps = [h for _, h in pairs]
refs_seg = [strip_tones(r) for r in refs]
hyps_seg = [strip_tones(h) for h in hyps]
for prefix, out in (
("w", jiwer.process_words(refs, hyps)),
("c", jiwer.process_characters(refs, hyps)),
("ws", jiwer.process_words(refs_seg, hyps_seg)),
("cs", jiwer.process_characters(refs_seg, hyps_seg)),
):
counters[f"{prefix}_s"] += out.substitutions
counters[f"{prefix}_d"] += out.deletions
counters[f"{prefix}_i"] += out.insertions
counters[f"{prefix}_h"] += out.hits
for r, h in zip(refs, hyps):
rs, rm = tone_sequence(r)
hs, _ = tone_sequence(h)
counters["tone_dist"] += levenshtein(rs, hs)
counters["tone_units"] += len(rs)
counters["tone_marked"] += rm
counters["n_scored"] += len(refs)
return counters
def _rate(errors: int, total: int) -> float | None:
return round(errors / total, 4) if total else None
def finalize(counters: dict) -> dict:
"""Métriques finales à partir des compteurs cumulés."""
c = {k: int(counters.get(k, 0)) for k in COUNTER_KEYS}
ref_words = c["w_h"] + c["w_s"] + c["w_d"]
hyp_words = c["w_h"] + c["w_s"] + c["w_i"]
wer_errors = c["w_s"] + c["w_d"] + c["w_i"]
wer = _rate(wer_errors, ref_words)
mer = _rate(wer_errors, wer_errors + c["w_h"])
if not ref_words:
wil = None
elif not hyp_words:
# Modèle muet : plus aucune information transmise (jiwer pose wip = 0).
wil = 1.0
else:
wil = round(1 - (c["w_h"] / ref_words) * (c["w_h"] / hyp_words), 4)
cer = _rate(c["c_s"] + c["c_d"] + c["c_i"], c["c_h"] + c["c_s"] + c["c_d"])
wer_seg = _rate(
c["ws_s"] + c["ws_d"] + c["ws_i"], c["ws_h"] + c["ws_s"] + c["ws_d"]
)
cer_seg = _rate(
c["cs_s"] + c["cs_d"] + c["cs_i"], c["cs_h"] + c["cs_s"] + c["cs_d"]
)
units, marked = c["tone_units"], c["tone_marked"]
annotated = units > 0 and (marked / units) >= _TONE_ANNOTATION_THRESHOLD
wer_ton = round(c["tone_dist"] / units, 4) if annotated else None
twer = (
round(wer_seg + 2 * wer_ton, 4)
if wer_ton is not None and wer_seg is not None
else None
)
return {
"wer": wer,
"cer": cer,
"mer": mer,
"wil": wil,
"wer_seg": wer_seg,
"cer_seg": cer_seg,
"wer_ton": wer_ton,
"twer": twer,
"tone_annotated": annotated,
"num_utterances_scored": c["n_scored"],
"version": FONBENCH_EVAL_VERSION,
}
if __name__ == "__main__":
# Auto-tests : python fonbench_eval.py
assert normalize_fon("Ɖò xwégbè, é ɖù nǔ!") == "ɖò xwégbè é ɖù nǔ"
assert strip_tones("étɛ́ ká ɖíe") == "etɛ ka ɖie", strip_tones("étɛ́ ká ɖíe")
seq, marked = tone_sequence("étɛ́ ká ɖíe") # é=H, ɛ́=H, á=H, í=H, e=.
assert seq == ["H", "H", "H", "H", "."], seq
assert marked == 4, marked
# Erreur purement tonale : segmental parfait, tonal fautif
m = compute_metrics(["étɛ́ ká ɖíe"], ["etɛ ka ɖie"])
assert m["wer_seg"] == 0.0, m
assert m["wer_ton"] and m["wer_ton"] > 0, m
assert m["twer"] == round(0 + 2 * m["wer_ton"], 4), m
# Corpus sans tons (type ALFFA) : wer_ton doit être None (non biaisé)
m2 = compute_metrics(["un yi axi me"], ["un yi axi me"])
assert m2["wer_ton"] is None and m2["twer"] is None, m2
# Modèle muet : WIL saturé, pas de division par zéro
m3 = compute_metrics(["étɛ́ ká ɖíe"], [""])
assert m3["wil"] == 1.0, m3
print("fonbench_eval OK —", compute_metrics(["étɛ́ ká ɖíe"], ["etɛ ka die"]))