Add generic Deep CFR config overrides
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user