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,139 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user