Files
K3/kda/training/eval_mt.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

140 lines
4.4 KiB
Python

"""Greedy translation eval on line-aligned src/ref files.
python -m kda.training.eval_mt \\
--ckpt ckpts/k3_wiki.pt --src /data/eval/zh2en.src.txt \\
--ref /data/eval/zh2en.ref.txt --target-lang en
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from kda.training.data import eos_id, load_tokenizer
from kda.training.prompts import instruction_prompt
from kda.training.success import _chrf, _detect_lang, translation_success
from kda.training.toy import load_ckpt
def _read_lines(path: str) -> list[str]:
return [ln.strip() for ln in Path(path).read_text(encoding="utf-8").splitlines() if ln.strip()]
def _instruction(src: str, target_lang: str) -> str:
return instruction_prompt(src, target_lang)
@torch.inference_mode()
def decode_one(model, tok, prompt: str, device: str, max_new: int) -> str:
ids = tok.encode(prompt)
if not ids:
return ""
inp = torch.tensor([ids], dtype=torch.long, device=device)
out = model.generate(inp, max_new, eos_token_id=eos_id(tok))
gen = out[0, inp.size(1) :].tolist()
return tok.decode(gen).strip()
def evaluate_pairs(
model,
tok,
srcs: list[str],
refs: list[str],
*,
target_lang: str,
device: str,
max_new: int,
limit: int | None,
) -> dict:
n = len(srcs)
if limit is not None:
n = min(n, limit)
hyps: list[str] = []
wins = 0
copies = 0
lang_ok = 0
chrf_sum = 0.0
for i in range(n):
src, ref = srcs[i], refs[i]
hyp = decode_one(model, tok, _instruction(src, target_lang), device, max_new)
hyps.append(hyp)
ok = translation_success(src, hyp, ref, target_lang=target_lang)
wins += int(ok)
copies += int(_chrf(hyp, src) >= 80.0 or hyp == src)
want = "zh" if target_lang.startswith("zh") else "en"
lang_ok += int(_detect_lang(hyp) == want)
chrf_sum += _chrf(hyp, ref)
corpus = {}
try:
from sacrebleu.metrics import BLEU, CHRF
corpus["chrf"] = float(CHRF(word_order=2).corpus_score(hyps, [refs[:n]]).score)
corpus["bleu"] = float(BLEU().corpus_score(hyps, [refs[:n]]).score)
except Exception:
corpus["chrf"] = chrf_sum / max(n, 1)
corpus["bleu"] = None
return {
"n": n,
"success_rate": wins / max(n, 1),
"copy_rate": copies / max(n, 1),
"lang_ok": lang_ok / max(n, 1),
"chrf": corpus["chrf"],
"bleu": corpus["bleu"],
"hyps": hyps,
}
def main() -> None:
p = argparse.ArgumentParser(description=__doc__)
p.add_argument("--ckpt", required=True)
p.add_argument("--tokenizer", default=None, help="override ckpt tokenizer field")
p.add_argument("--src", default=None, help="one source sentence per line")
p.add_argument("--ref", default=None, help="one reference sentence per line")
p.add_argument("--target-lang", default="en", choices=["en", "zh"])
p.add_argument("--max-new", type=int, default=64)
p.add_argument("--limit", type=int, default=None)
p.add_argument("--prefix", default=None, help="single-prompt smoke decode")
p.add_argument("--device", default="auto")
args = p.parse_args()
device = args.device
if device == "auto":
device = "cuda" if torch.cuda.is_available() else "cpu"
model, _config = load_ckpt(args.ckpt)
model.to(device).eval()
payload = torch.load(args.ckpt, map_location="cpu", weights_only=False)
tok_src = args.tokenizer or payload.get("tokenizer")
if not tok_src:
raise SystemExit("need --tokenizer or a 'tokenizer' field in the checkpoint")
tok = load_tokenizer(tok_src)
if args.prefix:
print(decode_one(model, tok, args.prefix, device, args.max_new))
if args.src and args.ref:
srcs, refs = _read_lines(args.src), _read_lines(args.ref)
if len(srcs) != len(refs):
raise SystemExit(f"src/ref length mismatch: {len(srcs)} vs {len(refs)}")
out = evaluate_pairs(
model,
tok,
srcs,
refs,
target_lang=args.target_lang,
device=device,
max_new=args.max_new,
limit=args.limit,
)
printable = {k: v for k, v in out.items() if k != "hyps"}
print(json.dumps(printable, ensure_ascii=False, indent=2))
elif not args.prefix:
raise SystemExit("pass --prefix and/or --src + --ref")
if __name__ == "__main__":
main()