Fix chrF fallback so English wiki garbage fails translation_success

Sacrebleu errors used to fall back to set-overlap unigrams, so any
English hyp scored ~70–80 against English refs and SFT reported
success 1.0. Use count-based char n-grams, keep BLEU failures from
clobbering chrF, and print a few hyps during eval.
This commit is contained in:
dela
2026-08-26 14:23:22 +08:00
parent 5a7d949b01
commit 5cc0555563
4 changed files with 48 additions and 14 deletions
+9 -3
View File
@@ -67,14 +67,18 @@ def evaluate_pairs(
want = "zh" if target_lang.startswith("zh") else "en"
lang_ok += int(_detect_lang(hyp) == want)
chrf_sum += _chrf(hyp, ref)
corpus = {}
corpus = {"chrf": chrf_sum / max(n, 1), "bleu": None}
try:
from sacrebleu.metrics import BLEU, CHRF
from sacrebleu.metrics import CHRF
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
except Exception:
pass
try:
from sacrebleu.metrics import BLEU
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
except Exception:
corpus["chrf"] = chrf_sum / max(n, 1)
corpus["bleu"] = None
return {
"n": n,
@@ -131,6 +135,8 @@ def main() -> None:
)
printable = {k: v for k, v in out.items() if k != "hyps"}
print(json.dumps(printable, ensure_ascii=False, indent=2))
for i, hyp in enumerate(out["hyps"][: min(5, out["n"])]):
print(f" [{i}] {hyp}")
elif not args.prefix:
raise SystemExit("pass --prefix and/or --src + --ref")
+27 -11
View File
@@ -28,8 +28,33 @@ def _detect_lang(text: str) -> str | None:
return tag[:2]
def _chrf_ngram(hyp: str, ref: str, max_n: int = 4) -> float:
"""Count-based char n-gram F (β=2), 0–100. Not set-overlap unigrams."""
from collections import Counter
hyp, ref = hyp.strip(), ref.strip()
if not hyp or not ref:
return 0.0
scores: list[float] = []
for n in range(1, max_n + 1):
if len(hyp) < n or len(ref) < n:
scores.append(0.0)
continue
hc = Counter(hyp[i : i + n] for i in range(len(hyp) - n + 1))
rc = Counter(ref[i : i + n] for i in range(len(ref) - n + 1))
overlap = sum((hc & rc).values())
prec = overlap / max(sum(hc.values()), 1)
rec = overlap / max(sum(rc.values()), 1)
if prec + rec == 0:
scores.append(0.0)
continue
beta2 = 4.0
scores.append((1.0 + beta2) * prec * rec / (beta2 * prec + rec))
return 100.0 * sum(scores) / max(len(scores), 1)
def _chrf(hyp: str, ref: str) -> float:
"""chrF++ in 0–100. Falls back to char unigram F if sacrebleu is missing."""
"""chrF++ in 0–100. Falls back to count-based char n-grams if sacrebleu is missing."""
if not hyp or not ref:
return 0.0
try:
@@ -37,16 +62,7 @@ def _chrf(hyp: str, ref: str) -> float:
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
except Exception:
hyp_c, ref_c = list(hyp), list(ref)
if not hyp_c:
return 0.0
ref_set = set(ref_c)
overlap = sum(1 for c in hyp_c if c in ref_set)
prec = overlap / len(hyp_c)
rec = overlap / max(len(ref_c), 1)
if prec + rec == 0:
return 0.0
return 100.0 * 2 * prec * rec / (prec + rec)
return _chrf_ngram(hyp, ref)
def translation_success(