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:
+45
-3
@@ -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/<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("--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} "
|
||||
|
||||
Reference in New Issue
Block a user