Resume the same SwanLab run from the id stored in the ckpt

swanlab.init always opened a new experiment on --resume. Save the run
id in the checkpoint and pass resume=True, id=... on the next start.
--swanlab-id overrides; --swanlab-new forces a fresh experiment.
This commit is contained in:
dela
2026-08-25 22:08:15 +08:00
parent 47c72e5bb8
commit 9a4862a866
+45 -3
View File
@@ -66,7 +66,21 @@ def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None:
group["lr"] = lr 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.""" """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op."""
key = os.environ.get("SWANLAB_API_KEY") key = os.environ.get("SWANLAB_API_KEY")
if not key: if not key:
@@ -80,7 +94,8 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
swanlab.login(api_key=key, save=False) swanlab.login(api_key=key, save=False)
# swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init. # swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init.
project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" 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, project=project,
name=f"{args.preset}-{cfg.attnres}", name=f"{args.preset}-{cfg.attnres}",
config={ config={
@@ -105,6 +120,16 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace):
"gen_every": args.gen_every, "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: except Exception as exc:
print(f"swanlab init failed ({exc}); continuing without cloud monitor") print(f"swanlab init failed ({exc}); continuing without cloud monitor")
return None return None
@@ -136,6 +161,7 @@ def _payload(
"batch": args.batch, "batch": args.batch,
"seq_len": args.seq_len, "seq_len": args.seq_len,
"grad_acc": args.grad_acc, "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("--heldout-frac", type=float, default=0.01)
p.add_argument("--resume", default=None, help="checkpoint to continue from") 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/<id>); 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("--gen-prefix", action="append", default=None)
p.add_argument("--device", default="auto") p.add_argument("--device", default="auto")
p.add_argument( p.add_argument(
@@ -358,11 +394,17 @@ def main() -> None:
tokens = int(payload.get("tokens", 0)) tokens = int(payload.get("tokens", 0))
chunk_index = int(payload.get("chunk_index", 0)) chunk_index = int(payload.get("chunk_index", 0))
best_heldout = float(payload.get("best_heldout", best_heldout)) 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: else:
model = CausalLM(cfg).to(device) model = CausalLM(cfg).to(device)
_apply_moe_coefs(model, cfg) _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()) n = sum(p.numel() for p in model.parameters())
print( print(
f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} " f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} "