Use generic train config overrides
This commit is contained in:
@@ -66,8 +66,11 @@ Short fixed-iteration run:
|
||||
```bash
|
||||
uv run lost-cities-deep-cfr train \
|
||||
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
||||
--iterations 100 \
|
||||
--save-latest-only
|
||||
--set run.iterations=100 \
|
||||
--set run.max_hours=null \
|
||||
--set run.max_iterations=null \
|
||||
--set checkpoint.save_latest_only=true \
|
||||
--set checkpoint.save_every_iteration=false
|
||||
```
|
||||
|
||||
Use explicit run directories for experiments. Put Deep CFR runs under
|
||||
@@ -77,8 +80,8 @@ Use explicit run directories for experiments. Put Deep CFR runs under
|
||||
RUN_DIR="runs/deep_cfr/$(date +%Y-%m-%d_%H%M%S)_deep_cfr_experiment_name"
|
||||
uv run lost-cities-deep-cfr train \
|
||||
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
||||
--checkpoint-dir "$RUN_DIR" \
|
||||
--max-iterations 100
|
||||
--set checkpoint.directory="$RUN_DIR" \
|
||||
--set run.max_iterations=100
|
||||
```
|
||||
|
||||
Date-prefixed examples:
|
||||
@@ -86,16 +89,23 @@ Date-prefixed examples:
|
||||
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_100iter`
|
||||
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_unbounded`
|
||||
|
||||
Useful train overrides:
|
||||
Useful train controls:
|
||||
|
||||
- `--resume`: resume from `<checkpoint-dir>/latest.pt`.
|
||||
- `--resume PATH`: resume from a specific checkpoint.
|
||||
- `--exact-resume`: require checkpoint config compatibility for exact resume.
|
||||
- `--no-save`: disable checkpoint writes.
|
||||
- `--save-latest-only`: keep only `latest.pt`.
|
||||
- `--save-iteration-interval N`: archive every N iterations.
|
||||
- `--set PATH=VALUE`: override arbitrary config fields, e.g.
|
||||
`--set traversal.num_workers=4`.
|
||||
- `--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 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.
|
||||
- `--set checkpoint.save_latest_only=true --set checkpoint.save_every_iteration=false`:
|
||||
keep only `latest.pt`.
|
||||
- `--set checkpoint.save_iteration_interval=N`: archive every N iterations.
|
||||
|
||||
## Long Runs
|
||||
|
||||
@@ -147,8 +157,8 @@ checkpoint:
|
||||
```
|
||||
|
||||
`latest.pt` is updated continuously; archive checkpoints are written every 100
|
||||
iterations. If disk is tight, prefer `--save-latest-only` or increase
|
||||
`save_iteration_interval`.
|
||||
iterations. If disk is tight, prefer `--set checkpoint.save_latest_only=true` or
|
||||
increase `save_iteration_interval`.
|
||||
|
||||
## Evaluation And Analysis
|
||||
|
||||
|
||||
@@ -80,52 +80,6 @@ def _resolve_resume_path(config: DeepCFRConfig, resume: str | None) -> str | Non
|
||||
|
||||
def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {}
|
||||
if args.no_save and (args.save_latest_only or args.save_iteration_interval is not None):
|
||||
raise ValueError("--no-save cannot be combined with checkpoint save overrides")
|
||||
run_overrides = overrides.setdefault("run", {})
|
||||
if args.iterations is not None:
|
||||
run_overrides["iterations"] = args.iterations
|
||||
if args.max_hours is None and args.max_iterations is None:
|
||||
run_overrides["max_hours"] = None
|
||||
run_overrides["max_iterations"] = None
|
||||
if args.max_hours is not None:
|
||||
run_overrides["max_hours"] = args.max_hours
|
||||
if args.max_iterations is not None:
|
||||
run_overrides["max_iterations"] = args.max_iterations
|
||||
if args.seed is not None:
|
||||
run_overrides["seed"] = args.seed
|
||||
if args.traversals_per_iteration is not None:
|
||||
traversal_overrides = overrides.setdefault("traversal", {})
|
||||
traversal_overrides["traversals_per_iteration"] = args.traversals_per_iteration
|
||||
traversal_overrides["traversals_per_player"] = None
|
||||
if args.num_workers is not None:
|
||||
overrides.setdefault("traversal", {})["num_workers"] = args.num_workers
|
||||
if args.checkpoint_dir is not None:
|
||||
overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir
|
||||
if args.eval_every is not None:
|
||||
overrides.setdefault("evaluation", {})["eval_every"] = args.eval_every
|
||||
if args.eval_games is not None:
|
||||
overrides.setdefault("evaluation", {})["games"] = args.eval_games
|
||||
if args.regret_fallback is not None:
|
||||
overrides.setdefault("regret_matching", {})["all_negative_fallback"] = args.regret_fallback
|
||||
if args.training_weighting is not None:
|
||||
overrides.setdefault("training_weighting", {})["mode"] = args.training_weighting
|
||||
if args.no_save:
|
||||
checkpoint_overrides = overrides.setdefault("checkpoint", {})
|
||||
checkpoint_overrides["save_latest"] = False
|
||||
checkpoint_overrides["save_every_iteration"] = False
|
||||
checkpoint_overrides["save_iteration_interval"] = 0
|
||||
if args.save_latest_only:
|
||||
checkpoint_overrides = overrides.setdefault("checkpoint", {})
|
||||
checkpoint_overrides["save_latest"] = True
|
||||
checkpoint_overrides["save_latest_only"] = True
|
||||
checkpoint_overrides["save_every_iteration"] = False
|
||||
if args.save_iteration_interval is not None:
|
||||
overrides.setdefault("checkpoint", {})["save_iteration_interval"] = (
|
||||
args.save_iteration_interval
|
||||
)
|
||||
if args.exact_resume:
|
||||
overrides.setdefault("checkpoint", {})["exact_resume"] = True
|
||||
for assignment in getattr(args, "config_overrides", None) or ():
|
||||
_set_path_override(overrides, assignment)
|
||||
return overrides
|
||||
@@ -249,39 +203,19 @@ def main(argv: list[str] | None = None) -> None:
|
||||
|
||||
train = subparsers.add_parser("train")
|
||||
train.add_argument("--config")
|
||||
train.add_argument("--iterations", type=int)
|
||||
train.add_argument("--max-hours", type=float)
|
||||
train.add_argument("--max-iterations", type=int)
|
||||
train.add_argument("--traversals-per-iteration", type=int)
|
||||
train.add_argument("--num-workers")
|
||||
train.add_argument("--checkpoint-dir")
|
||||
train.add_argument("--resume", nargs="?", const=_RESUME_LATEST, default=None)
|
||||
train.add_argument("--exact-resume", action="store_true")
|
||||
train.add_argument("--device")
|
||||
train.add_argument("--eval-every", type=int)
|
||||
train.add_argument("--eval-games", type=int)
|
||||
train.add_argument(
|
||||
"--regret-fallback",
|
||||
choices=("uniform", "argmax_tiebreak"),
|
||||
help="Override regret_matching.all_negative_fallback.",
|
||||
)
|
||||
train.add_argument(
|
||||
"--training-weighting",
|
||||
choices=("none", "lcfr", "dcfr"),
|
||||
help="Override training_weighting.mode.",
|
||||
)
|
||||
train.add_argument("--seed", type=int)
|
||||
train.add_argument(
|
||||
"--set",
|
||||
action="append",
|
||||
default=[],
|
||||
dest="config_overrides",
|
||||
metavar="PATH=VALUE",
|
||||
help="Override any config field using dotted paths, e.g. traversal.sampling_mode=external.",
|
||||
help=(
|
||||
"Override config fields using dotted paths. Repeatable. VALUE is parsed as "
|
||||
"YAML, e.g. --set traversal.num_workers=4 --set run.max_hours=null."
|
||||
),
|
||||
)
|
||||
train.add_argument("--no-save", action="store_true")
|
||||
train.add_argument("--save-latest-only", action="store_true")
|
||||
train.add_argument("--save-iteration-interval", type=int)
|
||||
train.add_argument(
|
||||
"--wandb",
|
||||
action="store_true",
|
||||
|
||||
@@ -74,26 +74,24 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> None:
|
||||
def test_deep_cfr_train_cli_accepts_run_and_traversal_config_overrides() -> None:
|
||||
args = type(
|
||||
"Args",
|
||||
(),
|
||||
{
|
||||
"iterations": 1,
|
||||
"max_hours": None,
|
||||
"max_iterations": None,
|
||||
"seed": None,
|
||||
"traversals_per_iteration": 1,
|
||||
"num_workers": "0",
|
||||
"checkpoint_dir": None,
|
||||
"eval_every": None,
|
||||
"eval_games": None,
|
||||
"regret_fallback": "argmax_tiebreak",
|
||||
"training_weighting": "lcfr",
|
||||
"no_save": True,
|
||||
"save_latest_only": False,
|
||||
"save_iteration_interval": None,
|
||||
"exact_resume": False,
|
||||
"config_overrides": [
|
||||
"run.iterations=1",
|
||||
"run.max_hours=null",
|
||||
"run.max_iterations=null",
|
||||
"traversal.traversals_per_iteration=1",
|
||||
"traversal.traversals_per_player=null",
|
||||
"traversal.num_workers=0",
|
||||
"regret_matching.all_negative_fallback=argmax_tiebreak",
|
||||
"training_weighting.mode=lcfr",
|
||||
"checkpoint.save_latest=false",
|
||||
"checkpoint.save_every_iteration=false",
|
||||
"checkpoint.save_iteration_interval=0",
|
||||
],
|
||||
},
|
||||
)()
|
||||
config = load_config("configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml")
|
||||
@@ -123,21 +121,12 @@ def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None:
|
||||
"Args",
|
||||
(),
|
||||
{
|
||||
"iterations": None,
|
||||
"max_hours": None,
|
||||
"max_iterations": None,
|
||||
"seed": None,
|
||||
"traversals_per_iteration": None,
|
||||
"num_workers": None,
|
||||
"checkpoint_dir": None,
|
||||
"eval_every": None,
|
||||
"eval_games": None,
|
||||
"regret_fallback": None,
|
||||
"training_weighting": None,
|
||||
"no_save": False,
|
||||
"save_latest_only": True,
|
||||
"save_iteration_interval": 1,
|
||||
"exact_resume": False,
|
||||
"config_overrides": [
|
||||
"checkpoint.save_latest=true",
|
||||
"checkpoint.save_latest_only=true",
|
||||
"checkpoint.save_every_iteration=false",
|
||||
"checkpoint.save_iteration_interval=1",
|
||||
],
|
||||
},
|
||||
)()
|
||||
|
||||
@@ -154,21 +143,6 @@ def test_deep_cfr_train_cli_accepts_generic_config_overrides() -> None:
|
||||
"Args",
|
||||
(),
|
||||
{
|
||||
"iterations": None,
|
||||
"max_hours": None,
|
||||
"max_iterations": None,
|
||||
"seed": None,
|
||||
"traversals_per_iteration": None,
|
||||
"num_workers": None,
|
||||
"checkpoint_dir": None,
|
||||
"eval_every": None,
|
||||
"eval_games": None,
|
||||
"regret_fallback": None,
|
||||
"training_weighting": None,
|
||||
"no_save": False,
|
||||
"save_latest_only": False,
|
||||
"save_iteration_interval": None,
|
||||
"exact_resume": False,
|
||||
"config_overrides": [
|
||||
"traversal.sampling_mode=external",
|
||||
"traversal.max_depth=null",
|
||||
|
||||
Reference in New Issue
Block a user