"""Instruction SFT for zh↔en translation. Prompt template matches eval_mt. 用法: uv run python train_sft.py --ckpt ckpts/k3_wiki.pt --data data/sft/train.jsonl uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data data/sft/opus.jsonl \\ --seq-len 512 --batch 4 --lr 5e-5 --epochs 2 """ from __future__ import annotations import argparse import os from dataclasses import asdict import torch from kda.training.data import ( IGNORE_INDEX, iter_sft_batches, load_sft_rows, load_tokenizer, ) from kda.training.eval_mt import evaluate_pairs from kda.training.schedule import lr_scale, total_opt_steps from kda.training.toy import load_ckpt def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None: for group in optim.param_groups: group["lr"] = lr def _init_swanlab(args: argparse.Namespace): key = os.environ.get("SWANLAB_API_KEY") if not key: return None try: import swanlab except ImportError: return None try: swanlab.login(api_key=key, save=False) project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" return swanlab.init( project=project, name=f"sft-{os.path.basename(args.ckpt)}", config={ "ckpt": args.ckpt, "data": args.data, "lr": args.lr, "batch": args.batch, "seq_len": args.seq_len, "epochs": args.epochs, }, ) except Exception as exc: print(f"swanlab init failed ({exc}); continuing without cloud monitor") return None def _read_lines(path: str) -> list[str]: from pathlib import Path return [ ln.strip() for ln in Path(path).read_text(encoding="utf-8").splitlines() if ln.strip() ] def main() -> None: p = argparse.ArgumentParser(description=__doc__) p.add_argument("--ckpt", required=True) p.add_argument("--data", required=True, help="jsonl {src,tgt,target_lang} or TSV") p.add_argument("--out", default="ckpts/k3_sft.pt") p.add_argument("--tokenizer", default=None) p.add_argument("--batch", type=int, default=4) p.add_argument("--seq-len", type=int, default=256) p.add_argument("--lr", type=float, default=1e-3) p.add_argument("--warmup", type=int, default=20) p.add_argument("--epochs", type=int, default=2) p.add_argument("--max-steps", type=int, default=None) p.add_argument("--grad-acc", type=int, default=1) p.add_argument("--eval-every", type=int, default=50) p.add_argument("--src", default=None, help="frozen eval src (not used as train)") p.add_argument("--ref", default=None) p.add_argument("--target-lang", default="en", choices=["en", "zh"]) 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" use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported() if device == "cuda" and not use_bf16: raise SystemExit("KDA training needs bf16") model, cfg = load_ckpt(args.ckpt) model.to(device) 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) rows = load_sft_rows(args.data) if not rows: raise SystemExit(f"no SFT rows in {args.data}") print(f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__}") steps_per_epoch = max((len(rows) + args.batch - 1) // args.batch, 1) max_micro = args.max_steps if max_micro is None: max_micro = steps_per_epoch * args.epochs horizon = total_opt_steps( max_tokens=None, max_micro=max_micro, batch=args.batch, seq_len=args.seq_len, grad_acc=args.grad_acc, ) optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) tracker = _init_swanlab(args) model.train() step = 0 opt_step = 0 best = float("inf") for _, x, y in iter_sft_batches(rows, tok, args.batch, args.seq_len): if step >= max_micro: break x, y = x.to(device), y.to(device) _set_lr(optim, args.lr * lr_scale(opt_step, args.warmup, horizon)) with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16): loss = model(x, labels=y, ignore_index=IGNORE_INDEX) / args.grad_acc loss.backward() if (step + 1) % args.grad_acc == 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optim.step() optim.zero_grad(set_to_none=True) opt_step += 1 raw = loss.item() * args.grad_acc if raw < best: best = raw if tracker is not None: tracker.log( {"sft/loss": raw, "sft/lr": optim.param_groups[0]["lr"]}, step=step, ) if step % args.eval_every == 0 or step == max_micro - 1: print(f"step {step:4d} sft loss {raw:.4f} lr {optim.param_groups[0]['lr']:.2e}") if args.src and args.ref: model.eval() srcs, refs = _read_lines(args.src), _read_lines(args.ref) out = evaluate_pairs( model, tok, srcs, refs, target_lang=args.target_lang, device=device, max_new=64, limit=None, ) printable = {k: v for k, v in out.items() if k != "hyps"} print(printable) if tracker is not None: tracker.log( { "eval/success_rate": printable["success_rate"], "eval/chrf": printable["chrf"], "eval/copy_rate": printable["copy_rate"], }, step=step, ) model.train() step += 1 os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) torch.save( { "config": asdict(cfg), "model_state": model.state_dict(), "optimizer_state": optim.state_dict(), "tokenizer": tok_src, "sft_data": args.data, "pretrained_ckpt": args.ckpt, }, args.out, ) print(f"best sft loss {best:.4f}; checkpoint -> {args.out}") if tracker is not None: tracker.finish() if __name__ == "__main__": main()