Add opponent_policy=average_strategy support

Strategy network (학습 중인 average policy)를 traversal opponent로 사용하는
새 옵션 추가. Deep CFR 이론적 수렴이 average strategy에 대한 보장이라는
점에 착안 — opponent_policy=network의 발산 문제를 완화할 수 있는지 실증.

구현:
- config: opponent_policy validator에 average_strategy 추가
- traversal.pyx: opponent_policy_id=3, strategy_network 인자, softmax 기반
  policy 도출 (_policy_from_strategy_network)
- workers.py: TraversalWorkerBatch에 strategy_network state_dict 추가
- trainer.py: 직렬/병렬 traversal call에 strategy_network 전달
- 1000-iter 실험 config 추가 (opponent_policy=network와 동일 hyperparam)

Co-Authored-By: Claude Haiku 4.5 <noreply@anthropic.com>
This commit is contained in:
2026-05-07 14:07:28 +09:00
co-authored by Claude Haiku 4.5
parent 6b8c36ebb2
commit f398c9fc4c
5 changed files with 224 additions and 3 deletions
@@ -0,0 +1,109 @@
run:
experiment_name: lost_cities_deep_cfr_opponent_average_strategy_512x3_1000iter
iterations: null
seed: 79
max_iterations: 1000
max_hours: null
device: cuda
use_amp: false
rules:
n_colors: 5
n_ranks: 9
min_rank: 2
n_handshakes: 3
hand_size: 8
expedition_penalty: -20
bonus_threshold: 8
bonus_amount: 20
encoding:
derived_playability: true
slot_aware_playability: true
network:
hidden_size: 512
num_layers: 3
activation: relu
traversal:
traversals_per_iteration: 2
traversals_per_player: 70
max_depth: null
max_nodes: 10000
max_nodes_per_traversal: 1000
regret_matching_epsilon: 0.0001
outcome_sampling_epsilon: 0.2
outcome_sampling_value_clip: 500.0
outcome_unsampled_regret: zero
cutoff_value_mode: score_diff
cutoff_rollouts: 0
cutoff_rollout_policy: random
cutoff_rollout_max_steps: 300
opponent_policy: average_strategy
strategy_sample_interval: 1
store_strategy_on_traverser_nodes: true
store_strategy_on_opponent_nodes: false
num_workers: 8
worker_chunk_size: 4
traversal_worker_chunk_size: 8
progress_every_traversals: 10
endpoint_depth_bucket_width: 100
endpoint_depth_bucket_max: 1000
regret_matching:
all_negative_fallback: argmax_tiebreak
training_weighting:
mode: none
self_play:
snapshot_every: 1
max_snapshots: 0
anchor_probability: 0.0
current_weight: 1.0
recent_weight: 0.0
older_weight: 0.0
anchor_weight: 0.0
recent_window: 5
optimization:
advantage_train_steps: 1
strategy_train_steps: 1
batch_size: 32
advantage_batch_size: 1024
strategy_batch_size: 1024
advantage_updates_per_iteration: 512
strategy_updates_per_iteration: 512
learning_rate: 0.00003
weight_decay: 0.0001
grad_clip: 1.0
memory:
advantage_capacity: 2000000
strategy_capacity: 2000000
checkpoint:
directory: runs/deep_cfr/deep_cfr_opponent_average_strategy_512x3_1000iter
save_latest: true
save_every_iteration: false
save_iteration_interval: 100
save_latest_only: false
progress_interval_seconds: 20.0
exact_resume: false
evaluation:
eval_every: 5
games: 100
opponents:
- random
- passive_discard
- safe_heuristic
- safe_heuristic_loose
- safe_heuristic_strict
- noisy_safe
max_steps: 10000
on_max_steps: score_diff
batch_size: 64
device: trainer
num_workers: 4
@@ -143,8 +143,10 @@ class TraversalConfig(StrictModel):
@field_validator("opponent_policy")
@classmethod
def _validate_opponent_policy(cls, value: str) -> str:
if value not in {"network", "safe_heuristic", "self_play_league"}:
raise ValueError("must be 'network', 'safe_heuristic', or 'self_play_league'")
if value not in {"network", "safe_heuristic", "self_play_league", "average_strategy"}:
raise ValueError(
"must be 'network', 'safe_heuristic', 'self_play_league', or 'average_strategy'"
)
return value
def resolved_num_workers(self, batches: int | None = None) -> int:
@@ -337,6 +337,11 @@ class DeepCFRTrainer:
player,
iteration,
device=self.device,
strategy_network=(
self.strategy_network
if self.config.traversal.opponent_policy == "average_strategy"
else None
),
action_size=self.action_size,
encoding=self.config.encoding,
epsilon=self.config.traversal.regret_matching_epsilon,
@@ -464,6 +469,12 @@ class DeepCFRTrainer:
{name: value.detach().cpu() for name, value in network.state_dict().items()}
for network in self.advantage_networks
]
strategy_payload: dict | None = None
if self.config.traversal.opponent_policy == "average_strategy":
strategy_payload = {
name: value.detach().cpu()
for name, value in self.strategy_network.state_dict().items()
}
chunk_size = self.config.traversal.resolved_worker_chunk_size()
batch_index = 0
for player in range(2):
@@ -485,6 +496,7 @@ class DeepCFRTrainer:
advantage_networks=network_payloads,
league_advantage_networks=self._league_payloads(),
worker_seed=self.config.run.seed + iteration * 1_000_003 + batch_index,
strategy_network=strategy_payload,
)
)
batch_index += 1
@@ -1,6 +1,7 @@
# cython: language_level=3, boundscheck=False, wraparound=False, cdivision=True, initializedcheck=False
"""Cython traversal engine and rollout primitives for Deep CFR."""
from libc.math cimport exp
from libc.stdlib cimport free, malloc
import numpy as np
@@ -63,6 +64,7 @@ cdef int _depth_bucket_start(int depth, int width, int max_depth) noexcept:
cdef class CythonDeepCFRTraverser:
cdef object advantage_networks
cdef object strategy_network
cdef object advantage_samples
cdef object strategy_samples
cdef object device
@@ -111,6 +113,7 @@ cdef class CythonDeepCFRTraverser:
object advantage_networks,
*,
object device,
object strategy_network=None,
int action_size,
object encoding=None,
float epsilon=1.0e-8,
@@ -140,6 +143,7 @@ cdef class CythonDeepCFRTraverser:
unsigned int seed=1,
):
self.advantage_networks = advantage_networks
self.strategy_network = strategy_network
self.advantage_samples = []
self.strategy_samples = []
self.device = device
@@ -185,8 +189,16 @@ cdef class CythonDeepCFRTraverser:
self.opponent_policy_id = 1
elif opponent_policy == "self_play_league":
self.opponent_policy_id = 2
elif opponent_policy == "average_strategy":
self.opponent_policy_id = 3
if strategy_network is None:
raise ValueError(
"opponent_policy='average_strategy' requires strategy_network"
)
else:
raise ValueError("opponent_policy must be 'network', 'safe_heuristic', or 'self_play_league'")
raise ValueError(
"opponent_policy must be 'network', 'safe_heuristic', 'self_play_league', or 'average_strategy'"
)
if all_negative_fallback == "uniform":
self.all_negative_fallback_id = 0
elif all_negative_fallback == "argmax_tiebreak":
@@ -447,6 +459,69 @@ cdef class CythonDeepCFRTraverser:
policy[selected] = 1.0
return info_state
cdef void _policy_from_strategy_network(
self,
GameState state,
int player,
unsigned char* legal,
float* policy,
):
cdef float[::1] info_view
cdef float[::1] logits_view
cdef int actions[MAX_ACTIONS]
cdef int action_count
cdef int i
cdef float max_legal_logit
cdef float total
cdef float value
cdef bint any_legal
info_state = np.empty(self.input_dim, dtype=np.float32)
info_view = info_state
_encode_info_state_with_flags_c(
state,
player,
&info_view[0],
self.derived_playability,
self.slot_aware_playability,
)
for i in range(self.action_size):
legal[i] = 0
action_count = state._unified_legal_actions_c(actions)
for i in range(action_count):
legal[actions[i]] = 1
with torch.inference_mode():
x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0)
logits = self.strategy_network(x).squeeze(0).detach().cpu().numpy().astype(np.float32)
logits_view = logits
max_legal_logit = 0.0
any_legal = False
for i in range(self.action_size):
if legal[i] != 0:
if not any_legal or logits_view[i] > max_legal_logit:
max_legal_logit = logits_view[i]
any_legal = True
if not any_legal:
for i in range(self.action_size):
policy[i] = 0.0
return
total = 0.0
for i in range(self.action_size):
if legal[i] != 0:
value = logits_view[i] - max_legal_logit
policy[i] = exp(value)
total += policy[i]
else:
policy[i] = 0.0
if total <= 0.0:
value = 1.0 / <float>action_count
for i in range(self.action_size):
policy[i] = value if legal[i] != 0 else 0.0
return
for i in range(self.action_size):
if legal[i] != 0:
policy[i] = policy[i] / total
cdef void _sampling_policy(
self,
const float* policy,
@@ -499,6 +574,18 @@ cdef class CythonDeepCFRTraverser:
if self.safe_heuristic_opponent_bot is None:
self.safe_heuristic_opponent_bot = SafeHeuristicBot()
return int(self.safe_heuristic_opponent_bot.act(state))
if self.opponent_policy_id == 3:
self._policy_from_strategy_network(state, player, legal, policy)
for i in range(self.action_size):
if legal[i] != 0:
actions[count] = i
count += 1
if count <= 0:
return -1
unified_action = _sample_policy_from_actions_c(
policy, actions, count, _next_double(&self.rng)
)
return self._from_unified_action_c(state, unified_action)
bucket = self.active_self_play_bucket
if bucket == 0:
return -1
@@ -840,6 +927,7 @@ def run_cython_traversal_batch(
*,
object device,
int action_size,
object strategy_network=None,
object encoding=None,
float epsilon=1.0e-8,
int strategy_sample_interval=1,
@@ -875,6 +963,7 @@ def run_cython_traversal_batch(
traverser = CythonDeepCFRTraverser(
advantage_networks,
device=device,
strategy_network=strategy_network,
action_size=action_size,
encoding=encoding,
epsilon=epsilon,
@@ -43,6 +43,7 @@ class TraversalWorkerBatch:
advantage_networks: list[dict[str, Any]]
league_advantage_networks: list[list[dict[str, Any]]]
worker_seed: int
strategy_network: dict[str, Any] | None = None
@dataclass(frozen=True)
@@ -76,6 +77,13 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
network.load_state_dict(state_dict)
network.eval()
league_networks.append(snapshot_networks)
strategy_network: torch.nn.Module | None = None
if batch.strategy_network is not None:
strategy_network = DeepCFRMLP.from_config(
batch.input_dim, batch.action_size, cfg.network
).to(device)
strategy_network.load_state_dict(batch.strategy_network)
strategy_network.eval()
game_config = LostCitiesConfig(**batch.game_config)
total_stats, advantage_samples, strategy_samples = run_cython_traversal_batch(
networks,
@@ -84,6 +92,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
batch.player,
batch.iteration,
device=device,
strategy_network=strategy_network,
action_size=batch.action_size,
encoding=cfg.encoding,
epsilon=cfg.traversal.regret_matching_epsilon,