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.
This commit is contained in:
@@ -0,0 +1,86 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user