Use configured train device
This commit is contained in:
@@ -93,13 +93,13 @@ Useful train controls:
|
||||
|
||||
- `--resume`: resume from `<checkpoint-dir>/latest.pt`.
|
||||
- `--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
|
||||
values as YAML, e.g. `--set traversal.num_workers=4` or
|
||||
`--set run.max_hours=null`.
|
||||
|
||||
Common `--set` overrides:
|
||||
|
||||
- `--set run.device=cuda`: set the trainer device.
|
||||
- `--set checkpoint.exact_resume=true`: require checkpoint config compatibility.
|
||||
- `--set checkpoint.save_latest=false --set checkpoint.save_every_iteration=false
|
||||
--set checkpoint.save_iteration_interval=0`: disable checkpoint writes.
|
||||
|
||||
@@ -105,7 +105,7 @@ def train_command(args: argparse.Namespace) -> None:
|
||||
trainer = DeepCFRTrainer(
|
||||
config,
|
||||
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,
|
||||
)
|
||||
if resume_path:
|
||||
@@ -204,7 +204,6 @@ def main(argv: list[str] | None = None) -> None:
|
||||
train = subparsers.add_parser("train")
|
||||
train.add_argument("--config")
|
||||
train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None)
|
||||
train.add_argument("--device")
|
||||
train.add_argument(
|
||||
"--set",
|
||||
action="append",
|
||||
|
||||
Reference in New Issue
Block a user