From 0fea13786f8df8285282a6b1bb486b047af171f1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Thu, 7 May 2026 15:53:43 +0900 Subject: [PATCH] Use generic train config overrides --- AGENTS.md | 36 +++++---- .../games/classic/deep_cfr/cli.py | 74 +------------------ tests/games/classic/test_deep_cfr_trainer.py | 66 +++++------------ 3 files changed, 47 insertions(+), 129 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 3d3df60..254d950 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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 `/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 diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py index 686bd74..5fd73e8 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -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", diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index c33f721..8b1ad57 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -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",