"""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/toy.jsonl uv run python train_sft.py --ckpt ckpts/k3_0.5b_best.pt --data opus-100 \\ --limit 100000 --seq-len 512 --batch 4 --lr 5e-5 --epochs 2 uv run python train_sft.py --resume --out ckpts/k3_sft.pt """ from __future__ import annotations import argparse import os from dataclasses import asdict import torch from kda.layers.latent_moe import moe_router_losses from kda.training.data import ( IGNORE_INDEX, iter_sft_batches, resolve_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.swanlab_env import prepare_swanlab_env, swanlab_run_id 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 _sibling(path: str, suffix: str) -> str: root, ext = os.path.splitext(path) return f"{root}{suffix}{ext}" def _save(path: str, payload: dict) -> None: os.makedirs(os.path.dirname(path) or ".", exist_ok=True) tmp = path + ".tmp" torch.save(payload, tmp) os.replace(tmp, path) def _is_better_eval( success: float, chrf: float, best_success: float, best_chrf: float ) -> bool: if success > best_success: return True if success == best_success and chrf > best_chrf: return True return False def _init_swanlab(args: argparse.Namespace, resume_id: str | None = None): key = os.environ.get("SWANLAB_API_KEY") if not key: return None project = prepare_swanlab_env() try: import swanlab except ImportError: print("SWANLAB_API_KEY set but swanlab is not installed") return None try: swanlab.login(api_key=key, save=False) init_kw = dict( project=project, name=f"sft-{os.path.basename(args.ckpt or args.out)}", config={ "ckpt": args.ckpt, "data": args.data, "lr": args.lr, "batch": args.batch, "seq_len": args.seq_len, "epochs": args.epochs, "grad_acc": args.grad_acc, "limit": args.limit, }, ) if resume_id: init_kw["id"] = resume_id init_kw["resume"] = True print(f"swanlab resume id={resume_id}") run = swanlab.init(**init_kw) got = swanlab_run_id(run) if got: args.swanlab_id = got print(f"swanlab run id {got}") return run 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 _payload( cfg, model, optim: torch.optim.Optimizer, args: argparse.Namespace, *, tok_src: str, step: int, opt_step: int, row_index: int, best_success: float, best_chrf: float, best_train: float, ): return { "config": asdict(cfg), "model_state": model.state_dict(), "optimizer_state": optim.state_dict(), "tokenizer": tok_src, "sft_data": args.data, "pretrained_ckpt": args.ckpt, "sft_step": step, "opt_step": opt_step, "row_index": row_index, "best_success": best_success, "best_chrf": best_chrf, "best_train": best_train, "batch": args.batch, "seq_len": args.seq_len, "grad_acc": args.grad_acc, "swanlab_id": getattr(args, "swanlab_id", None), } def main() -> None: p = argparse.ArgumentParser(description=__doc__) p.add_argument("--ckpt", default=None, help="pretrained (or SFT) checkpoint to start from") p.add_argument( "--resume", nargs="?", const="__last__", default=None, help="resume SFT; default path is _last", ) p.add_argument( "--data", default="opus-100", help="local jsonl/tsv, or 'opus-100' to stream Helsinki-NLP/opus-100 en-zh", ) p.add_argument( "--limit", type=int, default=100_000, help="OPUS source pairs to pull (each becomes zh2en + en2zh unless --one-dir)", ) p.add_argument( "--one-dir", action="store_true", help="only zh→en rows when pulling OPUS", ) 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( "--ckpt-every", type=int, default=500, help="write _last this many steps (0 = only interrupt + end)", ) 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() resume_path = args.resume if resume_path == "__last__": resume_path = _sibling(args.out, "_last") if resume_path is None and not args.ckpt: raise SystemExit("need --ckpt or --resume") if resume_path is not None and not os.path.isfile(resume_path): raise SystemExit(f"resume checkpoint not found: {resume_path}") 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") start_path = resume_path or args.ckpt model, cfg = load_ckpt(start_path) model.to(device) loaded = torch.load(start_path, map_location="cpu", weights_only=False) if args.ckpt is None: args.ckpt = loaded.get("pretrained_ckpt") tok_src = args.tokenizer or loaded.get("tokenizer") if not tok_src: raise SystemExit("need --tokenizer or a tokenizer field in the checkpoint") tok = load_tokenizer(tok_src) rows = resolve_sft_rows( args.data, limit=args.limit, both_dirs=not args.one_dir, ) if not rows: raise SystemExit(f"no SFT rows from {args.data}") 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, ) print( f"SFT {len(rows)} rows from {args.data}; model {cfg.__class__.__name__} " f"{max_micro} steps ({args.epochs} epoch, batch {args.batch}) " f"eval/{args.eval_every} ckpt/{args.ckpt_every}" ) optim = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) step = 0 opt_step = 0 row_start = 0 best_train = float("inf") best_success = -1.0 best_chrf = -1.0 if resume_path is not None: if loaded.get("optimizer_state"): optim.load_state_dict(loaded["optimizer_state"]) step = int(loaded.get("sft_step", 0)) opt_step = int(loaded.get("opt_step", 0)) row_start = int(loaded.get("row_index", 0)) best_train = float(loaded.get("best_train", best_train)) best_success = float(loaded.get("best_success", best_success)) best_chrf = float(loaded.get("best_chrf", best_chrf)) print( f"resume {resume_path} step {step} opt {opt_step} " f"row {row_start} best success {best_success:.2f} chrf {best_chrf:.2f}" ) tracker = _init_swanlab( args, resume_id=loaded.get("swanlab_id") if resume_path is not None else None, ) last_path = _sibling(args.out, "_last") best_path = _sibling(args.out, "_best") model.train() next_row = row_start def dump() -> dict: return _payload( cfg, model, optim, args, tok_src=tok_src, step=step, opt_step=opt_step, row_index=next_row, best_success=best_success, best_chrf=best_chrf, best_train=best_train, ) try: for row_index, x, y in iter_sft_batches( rows, tok, args.batch, args.seq_len, start=row_start ): 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): task = model(x, labels=y, ignore_index=IGNORE_INDEX) aux, z_loss = moe_router_losses(model) loss = (task + aux + z_loss) / 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 = float(task.detach()) if raw < best_train: best_train = raw next_row = row_index + args.batch if tracker is not None: tracker.log( { "sft/loss": raw, "sft/lr": optim.param_groups[0]["lr"], "moe/aux": float(aux.detach()), "moe/z": float(z_loss.detach()), }, step=step, ) eval_now = step % args.eval_every == 0 or step == max_micro - 1 ckpt_now = args.ckpt_every > 0 and step > 0 and ( step % args.ckpt_every == 0 or step == max_micro - 1 ) if eval_now: print( f"step {step:4d} sft loss {raw:.4f} " f"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, ) if step > 0 and _is_better_eval( printable["success_rate"], printable["chrf"], best_success, best_chrf, ): best_success = float(printable["success_rate"]) best_chrf = float(printable["chrf"]) _save(best_path, dump()) print( f" best success {best_success:.2f} " f"chrf {best_chrf:.2f} -> {best_path}" ) model.train() if device == "cuda": torch.cuda.empty_cache() if ckpt_now: _save(last_path, dump()) print(f" last -> {last_path}") step += 1 payload = dump() _save(last_path, payload) _save(args.out, payload) print( f"best train {best_train:.4f} best success {best_success:.2f} " f"chrf {best_chrf:.2f}; checkpoint -> {args.out}" ) except KeyboardInterrupt: print("interrupt; writing last checkpoint") _save(last_path, dump()) print(f" last -> {last_path}") if os.path.isfile(best_path): print(f" best remains {best_path}") raise SystemExit(130) from None finally: if tracker is not None: tracker.finish() if __name__ == "__main__": main()