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
|
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} "
|
||||||
|
|||||||
Reference in New Issue
Block a user