Use configured train device

This commit is contained in:
2026-05-07 15:58:10 +09:00
parent 0fea13786f
commit c2b88d7c9c
2 changed files with 2 additions and 3 deletions
+1 -1
View File
@@ -93,13 +93,13 @@ Useful train controls:
- `--resume`: resume from `<checkpoint-dir>/latest.pt`. - `--resume`: resume from `<checkpoint-dir>/latest.pt`.
- `--resume PATH`: resume from a specific checkpoint. - `--resume PATH`: resume from a specific checkpoint.
- `--device DEVICE`: override the trainer device for this invocation.
- `--set PATH=VALUE`: override config fields. It is repeatable and parses - `--set PATH=VALUE`: override config fields. It is repeatable and parses
values as YAML, e.g. `--set traversal.num_workers=4` or values as YAML, e.g. `--set traversal.num_workers=4` or
`--set run.max_hours=null`. `--set run.max_hours=null`.
Common `--set` overrides: Common `--set` overrides:
- `--set run.device=cuda`: set the trainer device.
- `--set checkpoint.exact_resume=true`: require checkpoint config compatibility. - `--set checkpoint.exact_resume=true`: require checkpoint config compatibility.
- `--set checkpoint.save_latest=false --set checkpoint.save_every_iteration=false - `--set checkpoint.save_latest=false --set checkpoint.save_every_iteration=false
--set checkpoint.save_iteration_interval=0`: disable checkpoint writes. --set checkpoint.save_iteration_interval=0`: disable checkpoint writes.
@@ -105,7 +105,7 @@ def train_command(args: argparse.Namespace) -> None:
trainer = DeepCFRTrainer( trainer = DeepCFRTrainer(
config, config,
config.rules.to_lost_cities_config(seed=config.run.seed), config.rules.to_lost_cities_config(seed=config.run.seed),
device=args.device or config.run.device, device=config.run.device,
extra_trackers=extra_trackers or None, extra_trackers=extra_trackers or None,
) )
if resume_path: if resume_path:
@@ -204,7 +204,6 @@ def main(argv: list[str] | None = None) -> None:
train = subparsers.add_parser("train") train = subparsers.add_parser("train")
train.add_argument("--config") train.add_argument("--config")
train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None) train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None)
train.add_argument("--device")
train.add_argument( train.add_argument(
"--set", "--set",
action="append", action="append",