diff --git a/kda/training/eval_mt.py b/kda/training/eval_mt.py index 2f7f22e..37b2cba 100644 --- a/kda/training/eval_mt.py +++ b/kda/training/eval_mt.py @@ -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") diff --git a/kda/training/success.py b/kda/training/success.py index 4422687..0db5f4a 100644 --- a/kda/training/success.py +++ b/kda/training/success.py @@ -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( diff --git a/tests/integration/test_eval_mt.py b/tests/integration/test_eval_mt.py index 1c2a795..1867d63 100644 --- a/tests/integration/test_eval_mt.py +++ b/tests/integration/test_eval_mt.py @@ -28,6 +28,16 @@ def test_good_zh2en_passes(): assert translation_success(src, hyp, ref, target_lang="en") is True +def test_english_wiki_garbage_does_not_pass_zh2en(): + from kda.training.success import _chrf + + src = "今天天气很好。" + hyp = "The first one's the time." + ref = "The weather is very nice today." + assert _chrf(hyp, ref) < 40.0 + assert translation_success(src, hyp, ref, target_lang="en") is False + + def test_container_help_exits_2(): import importlib.util from pathlib import Path diff --git a/train_sft.py b/train_sft.py index 56618bc..640e68d 100644 --- a/train_sft.py +++ b/train_sft.py @@ -339,6 +339,8 @@ def main() -> None: ) printable = {k: v for k, v in out.items() if k != "hyps"} print(printable) + for i, hyp in enumerate((out.get("hyps") or [])[:2]): + print(f" hyp[{i}] {hyp}") if tracker is not None: tracker.log( {