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 16db414..95139a7 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -5,6 +5,8 @@ import json from pathlib import Path from typing import Any +import yaml + from coolrl_lost_cities.games.classic.deep_cfr.analyze import analyze_run from coolrl_lost_cities.games.classic.deep_cfr.benchmark import ( benchmark_traversal, @@ -47,6 +49,23 @@ def _with_overrides(config: DeepCFRConfig, overrides: dict[str, Any]) -> DeepCFR return DeepCFRConfig.model_validate(data) +def _set_path_override(overrides: dict[str, Any], assignment: str) -> None: + if "=" not in assignment: + raise ValueError(f"config override must be PATH=VALUE: {assignment}") + path, raw_value = assignment.split("=", 1) + keys = path.split(".") + if any(not key for key in keys): + raise ValueError(f"config override path must use non-empty dotted keys: {path}") + value = yaml.safe_load(raw_value) + cursor = overrides + for key in keys[:-1]: + next_cursor = cursor.setdefault(key, {}) + if not isinstance(next_cursor, dict): + raise ValueError(f"config override path conflicts with scalar value: {path}") + cursor = next_cursor + cursor[keys[-1]] = value + + def _resolve_resume_path(config: DeepCFRConfig, resume: str | None) -> str | None: if resume != _RESUME_LATEST: return resume @@ -106,6 +125,8 @@ def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]: ) 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 @@ -236,6 +257,14 @@ def main(argv: list[str] | None = None) -> None: 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.", + ) 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) diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 795f06c..c30fe48 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -149,6 +149,43 @@ def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None: assert overridden.checkpoint.save_iteration_interval == 1 +def test_deep_cfr_train_cli_accepts_generic_config_overrides() -> None: + args = type( + "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", + "optimization.batch_size=64", + "checkpoint.save_latest_only=true", + ], + }, + )() + + overridden = _with_overrides(DeepCFRConfig(), _train_overrides_from_args(args)) + + assert overridden.traversal.sampling_mode == "external" + assert overridden.traversal.max_depth is None + assert overridden.optimization.batch_size == 64 + assert overridden.checkpoint.save_latest_only is True + + def test_deep_cfr_iteration_weights_use_sample_age() -> None: trainer = DeepCFRTrainer( _deep_cfr_config(