Add heuristic_balanced opponent_policy to interleaved scheduler
Mirrors the discard_only plumbing pattern. Uses HeuristicBot() (default balanced params) and converts the bot's phase-local action to unified via state.to_unified_action. Recursive (Cython) path was already supported and is unchanged. Tests: smoke run + accept/reject validators. All 59 tests pass.
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user