diff --git a/configs/deep_cfr/deep_cfr_opponent_average_strategy_512x3_1000iter.yaml b/configs/deep_cfr/deep_cfr_opponent_average_strategy_512x3_1000iter.yaml new file mode 100644 index 0000000..d8d42a5 --- /dev/null +++ b/configs/deep_cfr/deep_cfr_opponent_average_strategy_512x3_1000iter.yaml @@ -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 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 42b283e..55b084c 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -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: diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py index d4c4026..bd17b67 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -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 diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/traversal.pyx b/src/coolrl_lost_cities/games/classic/deep_cfr/traversal.pyx index 97e1bfc..6ef482c 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/traversal.pyx +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/traversal.pyx @@ -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 / 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, diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py index bad6294..fa358de 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -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,