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:
@@ -67,14 +67,18 @@ def evaluate_pairs(
|
|||||||
want = "zh" if target_lang.startswith("zh") else "en"
|
want = "zh" if target_lang.startswith("zh") else "en"
|
||||||
lang_ok += int(_detect_lang(hyp) == want)
|
lang_ok += int(_detect_lang(hyp) == want)
|
||||||
chrf_sum += _chrf(hyp, ref)
|
chrf_sum += _chrf(hyp, ref)
|
||||||
corpus = {}
|
corpus = {"chrf": chrf_sum / max(n, 1), "bleu": None}
|
||||||
try:
|
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)
|
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)
|
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
|
||||||
except Exception:
|
except Exception:
|
||||||
corpus["chrf"] = chrf_sum / max(n, 1)
|
|
||||||
corpus["bleu"] = None
|
corpus["bleu"] = None
|
||||||
return {
|
return {
|
||||||
"n": n,
|
"n": n,
|
||||||
@@ -131,6 +135,8 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||||
print(json.dumps(printable, ensure_ascii=False, indent=2))
|
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:
|
elif not args.prefix:
|
||||||
raise SystemExit("pass --prefix and/or --src + --ref")
|
raise SystemExit("pass --prefix and/or --src + --ref")
|
||||||
|
|
||||||
|
|||||||
+27
-11
@@ -28,8 +28,33 @@ def _detect_lang(text: str) -> str | None:
|
|||||||
return tag[:2]
|
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:
|
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:
|
if not hyp or not ref:
|
||||||
return 0.0
|
return 0.0
|
||||||
try:
|
try:
|
||||||
@@ -37,16 +62,7 @@ def _chrf(hyp: str, ref: str) -> float:
|
|||||||
|
|
||||||
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
||||||
except Exception:
|
except Exception:
|
||||||
hyp_c, ref_c = list(hyp), list(ref)
|
return _chrf_ngram(hyp, 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)
|
|
||||||
|
|
||||||
|
|
||||||
def translation_success(
|
def translation_success(
|
||||||
|
|||||||
@@ -28,6 +28,16 @@ def test_good_zh2en_passes():
|
|||||||
assert translation_success(src, hyp, ref, target_lang="en") is True
|
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():
|
def test_container_help_exits_2():
|
||||||
import importlib.util
|
import importlib.util
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|||||||
@@ -339,6 +339,8 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
printable = {k: v for k, v in out.items() if k != "hyps"}
|
printable = {k: v for k, v in out.items() if k != "hyps"}
|
||||||
print(printable)
|
print(printable)
|
||||||
|
for i, hyp in enumerate((out.get("hyps") or [])[:2]):
|
||||||
|
print(f" hyp[{i}] {hyp}")
|
||||||
if tracker is not None:
|
if tracker is not None:
|
||||||
tracker.log(
|
tracker.log(
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user