Deep CFR 재현 config 실행 보강

This commit is contained in:
2026-05-07 01:12:09 +09:00
parent bbf8950c3a
commit df8f979ce4
4 changed files with 123 additions and 10 deletions
@@ -8,6 +8,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
benchmark_traversal,
benchmark_traversal_modes,
)
from coolrl_lost_cities.games.classic.deep_cfr.cli import (
_train_overrides_from_args,
_with_overrides,
)
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load_config
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
@@ -58,6 +62,38 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
)
def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> 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,
"no_save": True,
},
)()
config = load_config(
"configs/deep_cfr/pure_self_play_zero_pit_poc_full_depth_slot_aware_playability.yaml"
)
overridden = _with_overrides(config, _train_overrides_from_args(args))
assert overridden.run.iterations == 1
assert overridden.run.max_hours is None
assert overridden.run.max_iterations is None
assert overridden.traversal.traversals_per_player is None
assert overridden.traversal.resolved_traversals_per_player() == 1
assert overridden.traversal.resolved_num_workers() == 0
assert overridden.checkpoint.save_every_iteration is False
def test_deep_cfr_playability_encoding_extends_input_shape() -> None:
state = GameState.new_game(LostCitiesConfig(seed=61), seed=61)
base_dim = input_dim(state)
+12
View File
@@ -38,3 +38,15 @@ def test_classic_package_exports_snapshot_alias() -> None:
def test_classic_package_exports_bot_registry_helpers() -> None:
assert "random" in classic.available_bot_names()
assert isinstance(classic.build_bot("random", seed=1), classic.LostCitiesPolicy)
def test_classic_bot_registry_accepts_reproduction_opponent_names() -> None:
for name in [
"random",
"passive_discard",
"safe_heuristic",
"safe_heuristic_loose",
"safe_heuristic_strict",
"noisy_safe",
]:
assert isinstance(classic.build_bot(name, seed=1), classic.LostCitiesPolicy)