Add generic Deep CFR config overrides

This commit is contained in:
2026-05-07 15:26:39 +09:00
parent 43507375fc
commit 4c223f76ff
2 changed files with 66 additions and 0 deletions
@@ -5,6 +5,8 @@ import json
from pathlib import Path from pathlib import Path
from typing import Any 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.analyze import analyze_run
from coolrl_lost_cities.games.classic.deep_cfr.benchmark import ( from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
benchmark_traversal, benchmark_traversal,
@@ -47,6 +49,23 @@ def _with_overrides(config: DeepCFRConfig, overrides: dict[str, Any]) -> DeepCFR
return DeepCFRConfig.model_validate(data) 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: def _resolve_resume_path(config: DeepCFRConfig, resume: str | None) -> str | None:
if resume != _RESUME_LATEST: if resume != _RESUME_LATEST:
return resume return resume
@@ -106,6 +125,8 @@ def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
) )
if args.exact_resume: if args.exact_resume:
overrides.setdefault("checkpoint", {})["exact_resume"] = True overrides.setdefault("checkpoint", {})["exact_resume"] = True
for assignment in getattr(args, "config_overrides", None) or ():
_set_path_override(overrides, assignment)
return overrides return overrides
@@ -236,6 +257,14 @@ def main(argv: list[str] | None = None) -> None:
help="Override training_weighting.mode.", help="Override training_weighting.mode.",
) )
train.add_argument("--seed", type=int) 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("--no-save", action="store_true")
train.add_argument("--save-latest-only", action="store_true") train.add_argument("--save-latest-only", action="store_true")
train.add_argument("--save-iteration-interval", type=int) train.add_argument("--save-iteration-interval", type=int)
@@ -149,6 +149,43 @@ def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None:
assert overridden.checkpoint.save_iteration_interval == 1 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: def test_deep_cfr_iteration_weights_use_sample_age() -> None:
trainer = DeepCFRTrainer( trainer = DeepCFRTrainer(
_deep_cfr_config( _deep_cfr_config(