"""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