"""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()