Use generic train config overrides
This commit is contained in:
@@ -66,8 +66,11 @@ Short fixed-iteration run:
|
|||||||
```bash
|
```bash
|
||||||
uv run lost-cities-deep-cfr train \
|
uv run lost-cities-deep-cfr train \
|
||||||
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
||||||
--iterations 100 \
|
--set run.iterations=100 \
|
||||||
--save-latest-only
|
--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
|
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"
|
RUN_DIR="runs/deep_cfr/$(date +%Y-%m-%d_%H%M%S)_deep_cfr_experiment_name"
|
||||||
uv run lost-cities-deep-cfr train \
|
uv run lost-cities-deep-cfr train \
|
||||||
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
--config configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml \
|
||||||
--checkpoint-dir "$RUN_DIR" \
|
--set checkpoint.directory="$RUN_DIR" \
|
||||||
--max-iterations 100
|
--set run.max_iterations=100
|
||||||
```
|
```
|
||||||
|
|
||||||
Date-prefixed examples:
|
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_100iter`
|
||||||
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_unbounded`
|
- `runs/deep_cfr/YYYY-MM-DD_HHMMSS_deep_cfr_unbounded`
|
||||||
|
|
||||||
Useful train overrides:
|
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.
|
||||||
- `--exact-resume`: require checkpoint config compatibility for exact resume.
|
- `--device DEVICE`: override the trainer device for this invocation.
|
||||||
- `--no-save`: disable checkpoint writes.
|
- `--set PATH=VALUE`: override config fields. It is repeatable and parses
|
||||||
- `--save-latest-only`: keep only `latest.pt`.
|
values as YAML, e.g. `--set traversal.num_workers=4` or
|
||||||
- `--save-iteration-interval N`: archive every N iterations.
|
`--set run.max_hours=null`.
|
||||||
- `--set PATH=VALUE`: override arbitrary config fields, e.g.
|
|
||||||
`--set traversal.num_workers=4`.
|
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
|
## Long Runs
|
||||||
|
|
||||||
@@ -147,8 +157,8 @@ checkpoint:
|
|||||||
```
|
```
|
||||||
|
|
||||||
`latest.pt` is updated continuously; archive checkpoints are written every 100
|
`latest.pt` is updated continuously; archive checkpoints are written every 100
|
||||||
iterations. If disk is tight, prefer `--save-latest-only` or increase
|
iterations. If disk is tight, prefer `--set checkpoint.save_latest_only=true` or
|
||||||
`save_iteration_interval`.
|
increase `save_iteration_interval`.
|
||||||
|
|
||||||
## Evaluation And Analysis
|
## 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]:
|
def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
|
||||||
overrides: 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 ():
|
for assignment in getattr(args, "config_overrides", None) or ():
|
||||||
_set_path_override(overrides, assignment)
|
_set_path_override(overrides, assignment)
|
||||||
return overrides
|
return overrides
|
||||||
@@ -249,39 +203,19 @@ 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("--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("--resume", nargs="?", const=_RESUME_LATEST, default=None)
|
||||||
train.add_argument("--exact-resume", action="store_true")
|
|
||||||
train.add_argument("--device")
|
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(
|
train.add_argument(
|
||||||
"--set",
|
"--set",
|
||||||
action="append",
|
action="append",
|
||||||
default=[],
|
default=[],
|
||||||
dest="config_overrides",
|
dest="config_overrides",
|
||||||
metavar="PATH=VALUE",
|
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(
|
train.add_argument(
|
||||||
"--wandb",
|
"--wandb",
|
||||||
action="store_true",
|
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 = type(
|
||||||
"Args",
|
"Args",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"iterations": 1,
|
"config_overrides": [
|
||||||
"max_hours": None,
|
"run.iterations=1",
|
||||||
"max_iterations": None,
|
"run.max_hours=null",
|
||||||
"seed": None,
|
"run.max_iterations=null",
|
||||||
"traversals_per_iteration": 1,
|
"traversal.traversals_per_iteration=1",
|
||||||
"num_workers": "0",
|
"traversal.traversals_per_player=null",
|
||||||
"checkpoint_dir": None,
|
"traversal.num_workers=0",
|
||||||
"eval_every": None,
|
"regret_matching.all_negative_fallback=argmax_tiebreak",
|
||||||
"eval_games": None,
|
"training_weighting.mode=lcfr",
|
||||||
"regret_fallback": "argmax_tiebreak",
|
"checkpoint.save_latest=false",
|
||||||
"training_weighting": "lcfr",
|
"checkpoint.save_every_iteration=false",
|
||||||
"no_save": True,
|
"checkpoint.save_iteration_interval=0",
|
||||||
"save_latest_only": False,
|
],
|
||||||
"save_iteration_interval": None,
|
|
||||||
"exact_resume": False,
|
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
config = load_config("configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml")
|
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",
|
"Args",
|
||||||
(),
|
(),
|
||||||
{
|
{
|
||||||
"iterations": None,
|
"config_overrides": [
|
||||||
"max_hours": None,
|
"checkpoint.save_latest=true",
|
||||||
"max_iterations": None,
|
"checkpoint.save_latest_only=true",
|
||||||
"seed": None,
|
"checkpoint.save_every_iteration=false",
|
||||||
"traversals_per_iteration": None,
|
"checkpoint.save_iteration_interval=1",
|
||||||
"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,
|
|
||||||
},
|
},
|
||||||
)()
|
)()
|
||||||
|
|
||||||
@@ -154,21 +143,6 @@ def test_deep_cfr_train_cli_accepts_generic_config_overrides() -> None:
|
|||||||
"Args",
|
"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": [
|
"config_overrides": [
|
||||||
"traversal.sampling_mode=external",
|
"traversal.sampling_mode=external",
|
||||||
"traversal.max_depth=null",
|
"traversal.max_depth=null",
|
||||||
|
|||||||
Reference in New Issue
Block a user