Files
K3/kda/training/success.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

87 lines
2.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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