Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
140 lines
4.4 KiB
Python
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()
|