diff --git a/src/coolrl_lost_cities/games/classic/bots/registry.py b/src/coolrl_lost_cities/games/classic/bots/registry.py index 9945b18..3bc113c 100644 --- a/src/coolrl_lost_cities/games/classic/bots/registry.py +++ b/src/coolrl_lost_cities/games/classic/bots/registry.py @@ -2,8 +2,8 @@ from __future__ import annotations from collections.abc import Callable -from ..policy import LostCitiesPolicy -from .heuristic import SafeHeuristicBot +from ..policy import LostCitiesPolicy, PolicyInput +from .heuristic import SafeHeuristicBot, SafeHeuristicParams from .passive import PassiveDiscardBot from .random import RandomBot @@ -11,20 +11,67 @@ BotName = str DEFAULT_BOT: BotName = "random" PolicyFactory = Callable[[int | None], LostCitiesPolicy] + +class NoisyPolicy(LostCitiesPolicy): + def __init__( + self, + base: LostCitiesPolicy, + random_policy: RandomBot, + *, + epsilon: float = 0.15, + ): + self.base = base + self.random_policy = random_policy + self.epsilon = epsilon + + def act(self, obs_or_state: PolicyInput) -> int: + if self.random_policy.rng.random() < self.epsilon: + return self.random_policy.act(obs_or_state) + return self.base.act(obs_or_state) + + +LOOSE_SAFE_HEURISTIC_PARAMS = SafeHeuristicParams( + open_target_ratio=0.42, + open_min_card_ratio=0.30, + handshake_target_multiplier=1.00, + handshake_min_card_ratio=0.25, + late_open_block_ratio=0.12, +) + +STRICT_SAFE_HEURISTIC_PARAMS = SafeHeuristicParams( + open_target_ratio=0.62, + open_min_card_ratio=0.50, + handshake_target_multiplier=1.35, + handshake_min_card_ratio=0.45, + late_open_block_ratio=0.30, +) + + BOT_REGISTRY: dict[BotName, PolicyFactory] = { DEFAULT_BOT: RandomBot, "passive-discard": lambda seed: PassiveDiscardBot(), "safe-heuristic": lambda seed: SafeHeuristicBot(), + "safe-heuristic-loose": lambda seed: SafeHeuristicBot(LOOSE_SAFE_HEURISTIC_PARAMS), + "safe-heuristic-strict": lambda seed: SafeHeuristicBot(STRICT_SAFE_HEURISTIC_PARAMS), + "noisy-safe": lambda seed: NoisyPolicy( + SafeHeuristicBot(), + RandomBot(seed), + ), } +def canonical_bot_name(name: BotName) -> BotName: + return name.strip().lower().replace("_", "-") + + def available_bot_names() -> list[BotName]: return sorted(BOT_REGISTRY) def build_bot(name: BotName, *, seed: int | None = None) -> LostCitiesPolicy: + canonical = canonical_bot_name(name) try: - policy_factory = BOT_REGISTRY[name] + policy_factory = BOT_REGISTRY[canonical] except KeyError as exc: raise ValueError(f"unknown Lost Cities bot: {name}") from exc return policy_factory(seed) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py index b38c670..44de22d 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -45,17 +45,26 @@ def _with_overrides(config: DeepCFRConfig, overrides: dict[str, Any]) -> DeepCFR return DeepCFRConfig.model_validate(data) -def train_command(args: argparse.Namespace) -> None: - config = _load_config(args.config) +def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]: overrides: dict[str, Any] = {} + run_overrides = overrides.setdefault("run", {}) if args.iterations is not None: - overrides.setdefault("run", {})["iterations"] = args.iterations + run_overrides["iterations"] = args.iterations + if args.max_hours is None and args.max_iterations is None: + run_overrides["max_hours"] = None + run_overrides["max_iterations"] = None + if args.max_hours is not None: + run_overrides["max_hours"] = args.max_hours + if args.max_iterations is not None: + run_overrides["max_iterations"] = args.max_iterations if args.seed is not None: - overrides.setdefault("run", {})["seed"] = args.seed + run_overrides["seed"] = args.seed if args.traversals_per_iteration is not None: - overrides.setdefault("traversal", {})["traversals_per_iteration"] = ( - args.traversals_per_iteration - ) + traversal_overrides = overrides.setdefault("traversal", {}) + traversal_overrides["traversals_per_iteration"] = args.traversals_per_iteration + traversal_overrides["traversals_per_player"] = None + if args.num_workers is not None: + overrides.setdefault("traversal", {})["num_workers"] = args.num_workers if args.checkpoint_dir is not None: overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir if args.eval_every is not None: @@ -64,6 +73,12 @@ def train_command(args: argparse.Namespace) -> None: overrides.setdefault("evaluation", {})["games"] = args.eval_games if args.no_save: overrides.setdefault("checkpoint", {})["save_every_iteration"] = False + return overrides + + +def train_command(args: argparse.Namespace) -> None: + config = _load_config(args.config) + overrides = _train_overrides_from_args(args) config = _with_overrides(config, overrides) trainer = DeepCFRTrainer( config, @@ -165,7 +180,10 @@ def main(argv: list[str] | None = None) -> None: train = subparsers.add_parser("train") train.add_argument("--config") train.add_argument("--iterations", type=int) + train.add_argument("--max-hours", type=float) + train.add_argument("--max-iterations", type=int) train.add_argument("--traversals-per-iteration", type=int) + train.add_argument("--num-workers") train.add_argument("--checkpoint-dir") train.add_argument("--resume") train.add_argument("--device") diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 748db20..91b7aee 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -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) diff --git a/tests/games/classic/test_public_api.py b/tests/games/classic/test_public_api.py index 846310a..5af3166 100644 --- a/tests/games/classic/test_public_api.py +++ b/tests/games/classic/test_public_api.py @@ -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)