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:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user