diff --git a/train_k3.py b/train_k3.py index 924e731..ccfe310 100644 --- a/train_k3.py +++ b/train_k3.py @@ -66,7 +66,21 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None: group["lr"] = lr -def _init_swanlab(cfg: K3Config, args: argparse.Namespace): +def _swanlab_run_id(run) -> str | None: + for attr in ("id", "run_id"): + val = getattr(run, attr, None) + if isinstance(val, str) and val: + return val + public = getattr(run, "public", None) + if public is not None: + for attr in ("cloud_run_id", "run_id", "id"): + val = getattr(public, attr, None) + if isinstance(val, str) and val: + return val + return None + + +def _init_swanlab(cfg: K3Config, args: argparse.Namespace, resume_id: str | None = None): """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" key = os.environ.get("SWANLAB_API_KEY") if not key: @@ -80,7 +94,8 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace): swanlab.login(api_key=key, save=False) # swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init. project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" - return swanlab.init( + run_id = resume_id or os.environ.get("SWANLAB_RUN_ID") + init_kw = dict( project=project, name=f"{args.preset}-{cfg.attnres}", config={ @@ -105,6 +120,16 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace): "gen_every": args.gen_every, }, ) + if run_id: + init_kw["id"] = run_id + init_kw["resume"] = True + print(f"swanlab resume id={run_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 @@ -136,6 +161,7 @@ def _payload( "batch": args.batch, "seq_len": args.seq_len, "grad_acc": args.grad_acc, + "swanlab_id": getattr(args, "swanlab_id", None), } @@ -239,6 +265,16 @@ def main() -> None: ) p.add_argument("--heldout-frac", type=float, default=0.01) p.add_argument("--resume", default=None, help="checkpoint to continue from") + p.add_argument( + "--swanlab-id", + default=None, + help="resume this SwanLab run (URL /runs/); default: id stored in ckpt", + ) + p.add_argument( + "--swanlab-new", + action="store_true", + help="start a new SwanLab run even when --resume", + ) p.add_argument("--gen-prefix", action="append", default=None) p.add_argument("--device", default="auto") p.add_argument( @@ -358,11 +394,17 @@ def main() -> None: tokens = int(payload.get("tokens", 0)) chunk_index = int(payload.get("chunk_index", 0)) best_heldout = float(payload.get("best_heldout", best_heldout)) + if not args.swanlab_new and not args.swanlab_id: + args.swanlab_id = payload.get("swanlab_id") or args.swanlab_id else: model = CausalLM(cfg).to(device) _apply_moe_coefs(model, cfg) - tracker = _init_swanlab(cfg, args) + tracker = _init_swanlab( + cfg, + args, + resume_id=None if args.swanlab_new else args.swanlab_id, + ) n = sum(p.numel() for p in model.parameters()) print( f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} "