Add --resume-from for warm-starting training from a checkpoint

Lets c6+ layer new exploration hyperparams on top of c5's learned
value head instead of restarting from random init. Saves ~60min per
cycle while preserving VPE-down trajectory observed in c5.
This commit is contained in:
2026-05-11 10:41:08 +09:00
parent 0d35341bbe
commit 33c44c708e
@@ -79,6 +79,19 @@ def train_command(args: argparse.Namespace) -> None:
device=config.run.device, device=config.run.device,
tracker=tracker, tracker=tracker,
) )
if args.resume_from:
import torch
ckpt = torch.load(args.resume_from, map_location=trainer.device, weights_only=False)
trainer.network.load_state_dict(ckpt["network"])
if "optimizer" in ckpt:
trainer.optimizer.load_state_dict(ckpt["optimizer"])
print(
f"[resume] loaded network + optimizer from {args.resume_from} "
f"(prior iteration={ckpt.get('iteration', '?')}); "
f"new run starts at iteration 1 with current config",
flush=True,
)
try: try:
trainer.train() trainer.train()
finally: finally:
@@ -111,6 +124,11 @@ def main(argv: list[str] | None = None) -> None:
train.add_argument("--wandb-job-type") train.add_argument("--wandb-job-type")
train.add_argument("--wandb-tag", action="append", default=[]) train.add_argument("--wandb-tag", action="append", default=[])
train.add_argument("--wandb-notes") train.add_argument("--wandb-notes")
train.add_argument(
"--resume-from",
default=None,
help="Path to a .pt checkpoint to warm-start network + optimizer state.",
)
train.set_defaults(func=train_command) train.set_defaults(func=train_command)
from .eval_checkpoint import add_eval_args, run_eval from .eval_checkpoint import add_eval_args, run_eval