Deep CFR 재현 config 실행 보강
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user