diff --git a/tests/integration/test_sft_ckpt.py b/tests/integration/test_sft_ckpt.py new file mode 100644 index 0000000..90060e2 --- /dev/null +++ b/tests/integration/test_sft_ckpt.py @@ -0,0 +1,13 @@ +from train_sft import _is_better_eval, _sibling + + +def test_sibling_last_best(): + assert _sibling("ckpts/k3_sft.pt", "_last") == "ckpts/k3_sft_last.pt" + assert _sibling("ckpts/k3_sft.pt", "_best") == "ckpts/k3_sft_best.pt" + + +def test_best_prefers_success_then_chrf(): + assert _is_better_eval(1.0, 70.0, 0.95, 90.0) + assert not _is_better_eval(0.95, 99.0, 1.0, 70.0) + assert _is_better_eval(1.0, 91.0, 1.0, 81.0) + assert not _is_better_eval(1.0, 70.0, 1.0, 81.0) diff --git a/train_sft.py b/train_sft.py index 41e7e37..656eca5 100644 --- a/train_sft.py +++ b/train_sft.py @@ -4,6 +4,7 @@ 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 @@ -23,6 +24,7 @@ from kda.training.data import ( ) 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 @@ -31,20 +33,43 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None: group["lr"] = lr -def _init_swanlab(args: argparse.Namespace): +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) - project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" - return swanlab.init( + init_kw = dict( project=project, - name=f"sft-{os.path.basename(args.ckpt)}", + name=f"sft-{os.path.basename(args.ckpt or args.out)}", config={ "ckpt": args.ckpt, "data": args.data, @@ -52,8 +77,20 @@ def _init_swanlab(args: argparse.Namespace): "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 @@ -69,9 +106,50 @@ def _read_lines(path: str) -> list[str]: ] +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", required=True) + 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", @@ -98,12 +176,26 @@ def main() -> None: 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" @@ -111,10 +203,13 @@ def main() -> None: if device == "cuda" and not use_bf16: raise SystemExit("KDA training needs bf16") - model, cfg = load_ckpt(args.ckpt) + start_path = resume_path or args.ckpt + model, cfg = load_ckpt(start_path) model.to(device) - payload = torch.load(args.ckpt, map_location="cpu", weights_only=False) - tok_src = args.tokenizer or payload.get("tokenizer") + 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) @@ -125,7 +220,6 @@ def main() -> None: ) if not rows: raise SystemExit(f"no SFT rows from {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 @@ -138,87 +232,157 @@ def main() -> None: 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) - tracker = _init_swanlab(args) - model.train() step = 0 opt_step = 0 - best = float("inf") + 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}" + ) - 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): - 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: - best = raw - 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, - ) - 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, + tracker = _init_swanlab( + args, + resume_id=loaded.get("swanlab_id") if resume_path is not None else None, ) - print(f"best sft loss {best:.4f}; checkpoint -> {args.out}") - if tracker is not None: - tracker.finish() + 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 _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 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__":