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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user