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.
103 lines
2.9 KiB
Python
103 lines
2.9 KiB
Python
"""Frozen translation success() — SFT eval and RL reward must call this."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import re
|
||
|
||
CHRF_MIN = 40.0
|
||
COPY_CHRF_MAX = 80.0
|
||
_CJK = re.compile(r"[\u4e00-\u9fff]")
|
||
|
||
|
||
def _detect_lang(text: str) -> str | None:
|
||
sample = text.strip()
|
||
if not sample:
|
||
return None
|
||
try:
|
||
from langdetect import detect
|
||
|
||
tag = detect(sample)
|
||
except Exception:
|
||
if _CJK.search(sample):
|
||
return "zh"
|
||
if any(c.isascii() and c.isalpha() for c in sample):
|
||
return "en"
|
||
return None
|
||
if tag.startswith("zh"):
|
||
return "zh"
|
||
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 count-based char n-grams if sacrebleu is missing."""
|
||
if not hyp or not ref:
|
||
return 0.0
|
||
try:
|
||
from sacrebleu.metrics import CHRF
|
||
|
||
return float(CHRF(word_order=2).sentence_score(hyp, [ref]).score)
|
||
except Exception:
|
||
return _chrf_ngram(hyp, ref)
|
||
|
||
|
||
def translation_success(
|
||
src: str,
|
||
hyp: str,
|
||
ref: str | None = None,
|
||
*,
|
||
target_lang: str,
|
||
chrf_min: float = CHRF_MIN,
|
||
copy_chrf_max: float = COPY_CHRF_MAX,
|
||
) -> bool:
|
||
"""Binary task success for zh↔en instruction translation.
|
||
|
||
1. non-empty hyp, no instruction leak prefix
|
||
2. langid(hyp) matches target_lang (zh / en)
|
||
3. hyp is not a copy of src
|
||
4. if ref is given, chrF(hyp, ref) >= chrf_min
|
||
"""
|
||
hyp = hyp.strip()
|
||
src = src.strip()
|
||
if not hyp:
|
||
return False
|
||
leak = ("翻译如下", "translate to", "translation:", "译文:")
|
||
head = hyp[:40].lower()
|
||
if any(p in head or p in hyp[:20] for p in leak):
|
||
return False
|
||
want = "zh" if target_lang.startswith("zh") else "en"
|
||
got = _detect_lang(hyp)
|
||
if got != want:
|
||
return False
|
||
if src and _chrf(hyp, src) >= copy_chrf_max:
|
||
return False
|
||
if hyp == src:
|
||
return False
|
||
if ref is not None and _chrf(hyp, ref.strip()) < chrf_min:
|
||
return False
|
||
return True
|