Deep CFR outcome sampling과 rollout cutoff 추가
This commit is contained in:
@@ -10,6 +10,13 @@ class DeepCFRConfig:
|
|||||||
max_traversal_depth: int | None = 8
|
max_traversal_depth: int | None = 8
|
||||||
max_nodes_per_traversal: int | None = 10_000
|
max_nodes_per_traversal: int | None = 10_000
|
||||||
regret_matching_epsilon: float = 1.0e-8
|
regret_matching_epsilon: float = 1.0e-8
|
||||||
|
outcome_sampling_epsilon: float = 0.0
|
||||||
|
outcome_sampling_value_clip: float | None = None
|
||||||
|
outcome_unsampled_regret: str = "negative_node_value"
|
||||||
|
cutoff_value_mode: str = "score_diff"
|
||||||
|
cutoff_rollouts: int = 0
|
||||||
|
cutoff_rollout_policy: str = "random"
|
||||||
|
cutoff_rollout_max_steps: int = 10_000
|
||||||
strategy_sample_interval: int = 1
|
strategy_sample_interval: int = 1
|
||||||
store_strategy_on_traverser_nodes: bool = True
|
store_strategy_on_traverser_nodes: bool = True
|
||||||
store_strategy_on_opponent_nodes: bool = True
|
store_strategy_on_opponent_nodes: bool = True
|
||||||
|
|||||||
@@ -77,6 +77,13 @@ class DeepCFRTrainer:
|
|||||||
store_strategy_on_opponent_nodes=self.config.store_strategy_on_opponent_nodes,
|
store_strategy_on_opponent_nodes=self.config.store_strategy_on_opponent_nodes,
|
||||||
max_depth=self.config.max_traversal_depth,
|
max_depth=self.config.max_traversal_depth,
|
||||||
max_nodes=self.config.max_nodes_per_traversal,
|
max_nodes=self.config.max_nodes_per_traversal,
|
||||||
|
outcome_sampling_epsilon=self.config.outcome_sampling_epsilon,
|
||||||
|
outcome_sampling_value_clip=self.config.outcome_sampling_value_clip,
|
||||||
|
outcome_unsampled_regret=self.config.outcome_unsampled_regret,
|
||||||
|
cutoff_value_mode=self.config.cutoff_value_mode,
|
||||||
|
cutoff_rollouts=self.config.cutoff_rollouts,
|
||||||
|
cutoff_rollout_policy=self.config.cutoff_rollout_policy,
|
||||||
|
cutoff_rollout_max_steps=self.config.cutoff_rollout_max_steps,
|
||||||
rng=self.rng,
|
rng=self.rng,
|
||||||
)
|
)
|
||||||
for network in self.advantage_networks:
|
for network in self.advantage_networks:
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from dataclasses import dataclass, field
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from coolrl_lost_cities.games.classic.bots import SafeHeuristicBot
|
||||||
from coolrl_lost_cities.games.classic.deep_cfr.cfr_math import regret_matching
|
from coolrl_lost_cities.games.classic.deep_cfr.cfr_math import regret_matching
|
||||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
|
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
|
||||||
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
|
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
|
||||||
@@ -21,6 +22,9 @@ class TraversalStats:
|
|||||||
advantage_samples: int = 0
|
advantage_samples: int = 0
|
||||||
strategy_samples: int = 0
|
strategy_samples: int = 0
|
||||||
sampled_actions: int = 0
|
sampled_actions: int = 0
|
||||||
|
cutoff_rollouts: int = 0
|
||||||
|
cutoff_rollout_steps: int = 0
|
||||||
|
cutoff_rollout_timeouts: int = 0
|
||||||
endpoint_depth_sum: int = 0
|
endpoint_depth_sum: int = 0
|
||||||
endpoint_depth_buckets: dict[str, int] = field(default_factory=dict)
|
endpoint_depth_buckets: dict[str, int] = field(default_factory=dict)
|
||||||
|
|
||||||
@@ -33,6 +37,9 @@ class TraversalStats:
|
|||||||
self.advantage_samples += other.advantage_samples
|
self.advantage_samples += other.advantage_samples
|
||||||
self.strategy_samples += other.strategy_samples
|
self.strategy_samples += other.strategy_samples
|
||||||
self.sampled_actions += other.sampled_actions
|
self.sampled_actions += other.sampled_actions
|
||||||
|
self.cutoff_rollouts += other.cutoff_rollouts
|
||||||
|
self.cutoff_rollout_steps += other.cutoff_rollout_steps
|
||||||
|
self.cutoff_rollout_timeouts += other.cutoff_rollout_timeouts
|
||||||
self.endpoint_depth_sum += other.endpoint_depth_sum
|
self.endpoint_depth_sum += other.endpoint_depth_sum
|
||||||
for key, value in other.endpoint_depth_buckets.items():
|
for key, value in other.endpoint_depth_buckets.items():
|
||||||
self.endpoint_depth_buckets[key] = self.endpoint_depth_buckets.get(key, 0) + value
|
self.endpoint_depth_buckets[key] = self.endpoint_depth_buckets.get(key, 0) + value
|
||||||
@@ -55,6 +62,9 @@ class TraversalStats:
|
|||||||
"traversal_advantage_samples": self.advantage_samples,
|
"traversal_advantage_samples": self.advantage_samples,
|
||||||
"traversal_strategy_samples": self.strategy_samples,
|
"traversal_strategy_samples": self.strategy_samples,
|
||||||
"traversal_sampled_actions": self.sampled_actions,
|
"traversal_sampled_actions": self.sampled_actions,
|
||||||
|
"traversal_cutoff_rollouts": self.cutoff_rollouts,
|
||||||
|
"traversal_cutoff_rollout_steps": self.cutoff_rollout_steps,
|
||||||
|
"traversal_cutoff_rollout_timeouts": self.cutoff_rollout_timeouts,
|
||||||
"traversal_avg_endpoint_depth": self.avg_endpoint_depth,
|
"traversal_avg_endpoint_depth": self.avg_endpoint_depth,
|
||||||
**{
|
**{
|
||||||
f"traversal_endpoint_depth_bucket_{key}": value
|
f"traversal_endpoint_depth_bucket_{key}": value
|
||||||
@@ -78,6 +88,13 @@ class DeepCFRTraverser:
|
|||||||
store_strategy_on_opponent_nodes: bool = True,
|
store_strategy_on_opponent_nodes: bool = True,
|
||||||
max_depth: int | None = None,
|
max_depth: int | None = None,
|
||||||
max_nodes: int | None = None,
|
max_nodes: int | None = None,
|
||||||
|
outcome_sampling_epsilon: float = 0.0,
|
||||||
|
outcome_sampling_value_clip: float | None = None,
|
||||||
|
outcome_unsampled_regret: str = "negative_node_value",
|
||||||
|
cutoff_value_mode: str = "score_diff",
|
||||||
|
cutoff_rollouts: int = 0,
|
||||||
|
cutoff_rollout_policy: str = "random",
|
||||||
|
cutoff_rollout_max_steps: int = 10_000,
|
||||||
rng: np.random.Generator | None = None,
|
rng: np.random.Generator | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.advantage_networks = advantage_networks
|
self.advantage_networks = advantage_networks
|
||||||
@@ -91,7 +108,27 @@ class DeepCFRTraverser:
|
|||||||
self.store_strategy_on_opponent_nodes = store_strategy_on_opponent_nodes
|
self.store_strategy_on_opponent_nodes = store_strategy_on_opponent_nodes
|
||||||
self.max_depth = max_depth
|
self.max_depth = max_depth
|
||||||
self.max_nodes = max_nodes
|
self.max_nodes = max_nodes
|
||||||
|
self.outcome_sampling_epsilon = min(1.0, max(0.0, float(outcome_sampling_epsilon)))
|
||||||
|
self.outcome_sampling_value_clip = (
|
||||||
|
None
|
||||||
|
if outcome_sampling_value_clip is None
|
||||||
|
else max(1.0e-9, float(outcome_sampling_value_clip))
|
||||||
|
)
|
||||||
|
self.outcome_unsampled_regret = outcome_unsampled_regret
|
||||||
|
if self.outcome_unsampled_regret not in {"negative_node_value", "zero"}:
|
||||||
|
raise ValueError("outcome_unsampled_regret must be 'negative_node_value' or 'zero'")
|
||||||
|
self.cutoff_value_mode = cutoff_value_mode
|
||||||
|
if self.cutoff_value_mode not in {"score_diff", "random_rollout"}:
|
||||||
|
raise ValueError("cutoff_value_mode must be 'score_diff' or 'random_rollout'")
|
||||||
|
self.cutoff_rollouts = max(0, int(cutoff_rollouts))
|
||||||
|
self.cutoff_rollout_policy = cutoff_rollout_policy
|
||||||
|
if self.cutoff_rollout_policy not in {"random", "safe_heuristic"}:
|
||||||
|
raise ValueError("cutoff_rollout_policy must be 'random' or 'safe_heuristic'")
|
||||||
|
self.cutoff_rollout_max_steps = max(1, int(cutoff_rollout_max_steps))
|
||||||
self.rng = rng or np.random.default_rng()
|
self.rng = rng or np.random.default_rng()
|
||||||
|
self._safe_heuristic_rollout_bot = (
|
||||||
|
SafeHeuristicBot() if self.cutoff_rollout_policy == "safe_heuristic" else None
|
||||||
|
)
|
||||||
|
|
||||||
def traverse(
|
def traverse(
|
||||||
self, state: GameState, traverser: int, iteration: int
|
self, state: GameState, traverser: int, iteration: int
|
||||||
@@ -115,7 +152,7 @@ class DeepCFRTraverser:
|
|||||||
if self.max_nodes is not None and stats.nodes >= self.max_nodes:
|
if self.max_nodes is not None and stats.nodes >= self.max_nodes:
|
||||||
stats.node_limit_cutoffs += 1
|
stats.node_limit_cutoffs += 1
|
||||||
self._record_endpoint(stats, depth)
|
self._record_endpoint(stats, depth)
|
||||||
return float(state.score_diff(traverser))
|
return self._cutoff_value(state, traverser, stats)
|
||||||
if state.terminal:
|
if state.terminal:
|
||||||
stats.terminals += 1
|
stats.terminals += 1
|
||||||
self._record_endpoint(stats, depth)
|
self._record_endpoint(stats, depth)
|
||||||
@@ -123,7 +160,7 @@ class DeepCFRTraverser:
|
|||||||
if self.max_depth is not None and depth >= self.max_depth:
|
if self.max_depth is not None and depth >= self.max_depth:
|
||||||
stats.depth_cutoffs += 1
|
stats.depth_cutoffs += 1
|
||||||
self._record_endpoint(stats, depth)
|
self._record_endpoint(stats, depth)
|
||||||
return float(state.score_diff(traverser))
|
return self._cutoff_value(state, traverser, stats)
|
||||||
|
|
||||||
player = state.current_player
|
player = state.current_player
|
||||||
info_state, legal, policy = self._policy(state, player)
|
info_state, legal, policy = self._policy(state, player)
|
||||||
@@ -135,8 +172,10 @@ class DeepCFRTraverser:
|
|||||||
self._record_endpoint(stats, depth)
|
self._record_endpoint(stats, depth)
|
||||||
return float(state.score_diff(traverser))
|
return float(state.score_diff(traverser))
|
||||||
|
|
||||||
action = self._sample_action(policy, legal_actions)
|
sampling_policy = self._sampling_policy(policy, legal)
|
||||||
|
action = self._sample_action(sampling_policy, legal_actions)
|
||||||
local_action = state.from_unified_action(int(action))
|
local_action = state.from_unified_action(int(action))
|
||||||
|
swapped_deck_index = self._sample_deck_draw_chance(state, int(action))
|
||||||
state.push_action(local_action)
|
state.push_action(local_action)
|
||||||
try:
|
try:
|
||||||
child_value = self._traverse(
|
child_value = self._traverse(
|
||||||
@@ -148,14 +187,27 @@ class DeepCFRTraverser:
|
|||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
state.pop_action()
|
state.pop_action()
|
||||||
|
if swapped_deck_index is not None:
|
||||||
|
state.swap_deck_cards(swapped_deck_index, len(state.deck) - 1)
|
||||||
|
|
||||||
stats.sampled_actions += 1
|
stats.sampled_actions += 1
|
||||||
action_prob = max(float(policy[action]), self.epsilon)
|
action_prob = max(float(sampling_policy[action]), self.epsilon)
|
||||||
sampled_action_value = child_value / action_prob
|
sampled_action_value = child_value / action_prob
|
||||||
|
if self.outcome_sampling_value_clip is not None:
|
||||||
|
sampled_action_value = float(
|
||||||
|
np.clip(
|
||||||
|
sampled_action_value,
|
||||||
|
-self.outcome_sampling_value_clip,
|
||||||
|
self.outcome_sampling_value_clip,
|
||||||
|
)
|
||||||
|
)
|
||||||
node_value = float(policy[action]) * sampled_action_value
|
node_value = float(policy[action]) * sampled_action_value
|
||||||
|
|
||||||
if player == traverser:
|
if player == traverser:
|
||||||
regrets = np.where(legal, -node_value, 0.0).astype(np.float32)
|
if self.outcome_unsampled_regret == "zero":
|
||||||
|
regrets = np.zeros_like(policy, dtype=np.float32)
|
||||||
|
else:
|
||||||
|
regrets = np.where(legal, -node_value, 0.0).astype(np.float32)
|
||||||
regrets[action] = np.float32(sampled_action_value - node_value)
|
regrets[action] = np.float32(sampled_action_value - node_value)
|
||||||
self.advantage_memory.add(
|
self.advantage_memory.add(
|
||||||
TrainingSample(
|
TrainingSample(
|
||||||
@@ -195,6 +247,62 @@ class DeepCFRTraverser:
|
|||||||
probs /= total
|
probs /= total
|
||||||
return int(self.rng.choice(legal_actions, p=probs))
|
return int(self.rng.choice(legal_actions, p=probs))
|
||||||
|
|
||||||
|
def _sampling_policy(self, policy: np.ndarray, legal: np.ndarray) -> np.ndarray:
|
||||||
|
legal_count = int(np.count_nonzero(legal))
|
||||||
|
if legal_count <= 0:
|
||||||
|
return np.zeros_like(policy, dtype=np.float32)
|
||||||
|
if self.outcome_sampling_epsilon <= 0.0:
|
||||||
|
return policy.astype(np.float32)
|
||||||
|
uniform = legal.astype(np.float32) / float(legal_count)
|
||||||
|
return (
|
||||||
|
(1.0 - self.outcome_sampling_epsilon) * policy + self.outcome_sampling_epsilon * uniform
|
||||||
|
).astype(np.float32)
|
||||||
|
|
||||||
|
def _cutoff_value(self, state: GameState, traverser: int, stats: TraversalStats) -> float:
|
||||||
|
if self.cutoff_value_mode == "score_diff" or self.cutoff_rollouts <= 0:
|
||||||
|
return float(state.score_diff(traverser))
|
||||||
|
total = 0.0
|
||||||
|
for _ in range(self.cutoff_rollouts):
|
||||||
|
total += self._rollout_value(state, traverser, stats)
|
||||||
|
return total / float(self.cutoff_rollouts)
|
||||||
|
|
||||||
|
def _rollout_value(self, state: GameState, traverser: int, stats: TraversalStats) -> float:
|
||||||
|
rollout_state = state.clone()
|
||||||
|
steps = 0
|
||||||
|
while not rollout_state.terminal and steps < self.cutoff_rollout_max_steps:
|
||||||
|
action = self._rollout_action(rollout_state)
|
||||||
|
if action is None:
|
||||||
|
break
|
||||||
|
unified_action = rollout_state.to_unified_action(action)
|
||||||
|
self._sample_deck_draw_chance(rollout_state, unified_action)
|
||||||
|
rollout_state.apply_action(action)
|
||||||
|
steps += 1
|
||||||
|
stats.cutoff_rollouts += 1
|
||||||
|
stats.cutoff_rollout_steps += steps
|
||||||
|
if not rollout_state.terminal:
|
||||||
|
stats.cutoff_rollout_timeouts += 1
|
||||||
|
return float(rollout_state.score_diff(traverser))
|
||||||
|
|
||||||
|
def _rollout_action(self, state: GameState) -> int | None:
|
||||||
|
if self.cutoff_rollout_policy == "safe_heuristic":
|
||||||
|
if self._safe_heuristic_rollout_bot is None:
|
||||||
|
self._safe_heuristic_rollout_bot = SafeHeuristicBot()
|
||||||
|
return self._safe_heuristic_rollout_bot.act(state)
|
||||||
|
legal_actions = state.unified_legal_actions()
|
||||||
|
if not legal_actions:
|
||||||
|
return None
|
||||||
|
return state.from_unified_action(int(self.rng.choice(legal_actions)))
|
||||||
|
|
||||||
|
def _sample_deck_draw_chance(self, state: GameState, unified_action: int) -> int | None:
|
||||||
|
deck_draw_action = 2 * state.config.hand_size
|
||||||
|
if state.phase != "draw" or unified_action != deck_draw_action or len(state.deck) <= 1:
|
||||||
|
return None
|
||||||
|
sampled_index = int(self.rng.integers(0, len(state.deck)))
|
||||||
|
if sampled_index == len(state.deck) - 1:
|
||||||
|
return None
|
||||||
|
state.swap_deck_cards(sampled_index, len(state.deck) - 1)
|
||||||
|
return sampled_index
|
||||||
|
|
||||||
def _record_strategy(
|
def _record_strategy(
|
||||||
self,
|
self,
|
||||||
info_state: np.ndarray,
|
info_state: np.ndarray,
|
||||||
|
|||||||
@@ -74,3 +74,55 @@ def test_deep_cfr_recursive_traverser_restores_state_and_collects_samples() -> N
|
|||||||
sample = trainer.advantage_memory.all()[0]
|
sample = trainer.advantage_memory.all()[0]
|
||||||
assert sample.legal_mask.dtype == bool
|
assert sample.legal_mask.dtype == bool
|
||||||
assert sample.target.shape == sample.legal_mask.shape
|
assert sample.target.shape == sample.legal_mask.shape
|
||||||
|
|
||||||
|
|
||||||
|
def test_deep_cfr_traverser_supports_outcome_sampling_and_rollout_cutoffs() -> None:
|
||||||
|
trainer = DeepCFRTrainer(
|
||||||
|
DeepCFRConfig(
|
||||||
|
iterations=1,
|
||||||
|
traversals_per_iteration=1,
|
||||||
|
max_traversal_depth=1,
|
||||||
|
max_nodes_per_traversal=32,
|
||||||
|
outcome_sampling_epsilon=0.25,
|
||||||
|
outcome_sampling_value_clip=100.0,
|
||||||
|
outcome_unsampled_regret="zero",
|
||||||
|
cutoff_value_mode="random_rollout",
|
||||||
|
cutoff_rollouts=2,
|
||||||
|
cutoff_rollout_policy="random",
|
||||||
|
cutoff_rollout_max_steps=16,
|
||||||
|
batch_size=2,
|
||||||
|
hidden_size=16,
|
||||||
|
seed=31,
|
||||||
|
),
|
||||||
|
LostCitiesConfig(seed=31),
|
||||||
|
)
|
||||||
|
state = GameState.new_game(LostCitiesConfig(seed=31), seed=31)
|
||||||
|
before = state.to_snapshot()
|
||||||
|
traverser = DeepCFRTraverser(
|
||||||
|
trainer.advantage_networks,
|
||||||
|
trainer.advantage_memory,
|
||||||
|
trainer.strategy_memory,
|
||||||
|
device=trainer.device,
|
||||||
|
action_size=trainer.action_size,
|
||||||
|
max_depth=1,
|
||||||
|
max_nodes=32,
|
||||||
|
outcome_sampling_epsilon=0.25,
|
||||||
|
outcome_sampling_value_clip=100.0,
|
||||||
|
outcome_unsampled_regret="zero",
|
||||||
|
cutoff_value_mode="random_rollout",
|
||||||
|
cutoff_rollouts=2,
|
||||||
|
cutoff_rollout_policy="random",
|
||||||
|
cutoff_rollout_max_steps=16,
|
||||||
|
rng=np.random.default_rng(31),
|
||||||
|
)
|
||||||
|
|
||||||
|
_, stats = traverser.traverse(state, traverser=0, iteration=1)
|
||||||
|
|
||||||
|
assert state.to_snapshot() == before
|
||||||
|
assert stats.depth_cutoffs > 0
|
||||||
|
assert stats.cutoff_rollouts == stats.depth_cutoffs * 2
|
||||||
|
assert stats.cutoff_rollout_steps > 0
|
||||||
|
sample = trainer.advantage_memory.all()[0]
|
||||||
|
unsampled_legal = sample.legal_mask.copy()
|
||||||
|
unsampled_legal[np.nonzero(sample.target)[0]] = False
|
||||||
|
assert np.all(sample.target[unsampled_legal] == 0.0)
|
||||||
|
|||||||
Reference in New Issue
Block a user