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