Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
87 lines
2.3 KiB
Python
87 lines
2.3 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(hyp: str, ref: str) -> float:
|
||
"""chrF++ in 0–100. Falls back to char unigram F 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:
|
||
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)
|
||
|
||
|
||
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
|