diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py index 029ff6d..7aabfbb 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -197,10 +197,16 @@ class TraversalConfig(StrictModel): if self.scheduler == "interleaved": if self.sampling_mode != "outcome": raise ValueError("scheduler='interleaved' currently supports only outcome sampling") - if self.opponent_policy not in {"network", "average_strategy", "discard_only"}: + if self.opponent_policy not in { + "network", + "average_strategy", + "discard_only", + "heuristic_balanced", + }: raise ValueError( "scheduler='interleaved' currently supports only " - "opponent_policy='network', 'average_strategy', or 'discard_only'" + "opponent_policy='network', 'average_strategy', 'discard_only', " + "or 'heuristic_balanced'" ) if self.cutoff_rollouts != 0 or self.cutoff_value_mode != "score_diff": raise ValueError( diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py b/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py index d464a74..621f871 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py @@ -9,6 +9,7 @@ import numpy as np import torch from coolrl_lost_cities.games.classic.bots.discard_only import DiscardOnlyBot +from coolrl_lost_cities.games.classic.bots.heuristic_py import HeuristicBot from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state from coolrl_lost_cities.games.classic.deep_cfr.memory import TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.traversal_stats import TraversalStats @@ -381,6 +382,9 @@ class InterleavedContext: self._discard_only_bot: DiscardOnlyBot | None = ( DiscardOnlyBot() if cfg.opponent_policy == "discard_only" else None ) + self._heuristic_bot: HeuristicBot | None = ( + HeuristicBot() if cfg.opponent_policy == "heuristic_balanced" else None + ) def advance_until_policy(self, context_index: int) -> None: while not self.done and self.pending is None and self.stack: @@ -482,6 +486,14 @@ class InterleavedContext: self.stack.append(FixedActionFrame(swapped_deck_index=swapped_deck_index)) self.stack.append(EnterFrame(depth + 1)) return + if player != self.traverser and self._heuristic_bot is not None: + phase_local_action = int(self._heuristic_bot.act(self.state)) + action = int(self.state.to_unified_action(phase_local_action)) + swapped_deck_index = self._sample_deck_draw_chance(action) + self.state.push_unified_action(action) + self.stack.append(FixedActionFrame(swapped_deck_index=swapped_deck_index)) + self.stack.append(EnterFrame(depth + 1)) + return info_state = encode_info_state(self.state, player, self.cfg.encoding) legal_mask = np.zeros(self.cfg.action_size, dtype=bool) legal_mask[legal_actions] = True diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 5ec2f89..e7450d1 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -139,7 +139,10 @@ def test_deep_cfr_config_accepts_interleaved_scheduler() -> None: def test_deep_cfr_config_rejects_unsupported_interleaved_options() -> None: with pytest.raises( - ValueError, match="opponent_policy='network', 'average_strategy', or 'discard_only'" + ValueError, + match=( + "opponent_policy='network', 'average_strategy', 'discard_only', or 'heuristic_balanced'" + ), ): _deep_cfr_config( {"traversal": {"scheduler": "interleaved", "opponent_policy": "self_play_league"}} @@ -163,6 +166,13 @@ def test_deep_cfr_config_accepts_discard_only_with_interleaved() -> None: assert config.traversal.opponent_policy == "discard_only" +def test_deep_cfr_config_accepts_heuristic_balanced_with_interleaved() -> None: + config = _deep_cfr_config( + {"traversal": {"scheduler": "interleaved", "opponent_policy": "heuristic_balanced"}} + ) + assert config.traversal.opponent_policy == "heuristic_balanced" + + def test_deep_cfr_config_rejects_discard_only_with_recursive() -> None: with pytest.raises(ValueError, match="discard_only.*scheduler='interleaved'"): _deep_cfr_config( @@ -504,6 +514,42 @@ def test_deep_cfr_trainer_discard_only_opponent_smoke_run(tmp_path) -> None: assert metrics[0].traversal_nodes > 0 +def test_deep_cfr_trainer_heuristic_balanced_opponent_smoke_run(tmp_path) -> None: + trainer = DeepCFRTrainer( + _deep_cfr_config( + { + "run": {"max_iterations": 1, "seed": 27}, + "network": {"hidden_size": 16}, + "traversal": { + "scheduler": "interleaved", + "opponent_policy": "heuristic_balanced", + "traversals_per_player": 2, + "max_depth": 3, + "max_nodes_per_traversal": 64, + "interleave_width": 4, + "interleave_max_batch": 8, + }, + "optimization": { + "advantage_updates_per_iteration": 1, + "strategy_updates_per_iteration": 1, + "advantage_batch_size": 2, + "strategy_batch_size": 2, + }, + "checkpoint": {"save_every": 0, "save_latest": False}, + "evaluation": {"eval_every": 0}, + } + ), + LostCitiesConfig(seed=27), + run_dir=tmp_path / "heuristic_balanced", + ) + + metrics = trainer.train() + + assert len(metrics) == 1 + assert metrics[0].advantage_samples > 0 + assert metrics[0].traversal_nodes > 0 + + def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() -> None: config = _deep_cfr_config( {