diff --git a/docs/performance.md b/docs/performance.md index 48e5fa7..a91a268 100644 --- a/docs/performance.md +++ b/docs/performance.md @@ -535,6 +535,20 @@ single-traversal parity test matches recursive stats and sample target checksums under identical RNG seed. Longer learning-curve A/B is still required before considering a default switch. +Follow-up: `average_strategy` support was added after the initial Phase 3 +network-opponent A/B so the interleaved path can run the actual default opponent +policy. A 10-iteration throughput check with default opponent policy, +evaluation/checkpoint disabled, and the same 8-worker chunk-64 interleaving +settings produced warm-up-excluded means: + +| Mode | iter s | traversal s | batch mean | batch max | +| --- | ---: | ---: | ---: | ---: | +| interleaved, default `average_strategy` | 10.61 | 4.85 | 28.3 | 64 | + +Run: `runs/2026-05-07_230419_option-b-interleaved-average-strategy-10i`. +This prepares the long-run default-config A/B, but it does not replace it: +learning-curve stability still has to be measured before any default switch. + ## Batched Traversal Inference: Design Decision (2026-05-07) Three structural options were considered for Priority #5: diff --git a/docs/plans/option_b_interleaved_traversal.md b/docs/plans/option_b_interleaved_traversal.md index d496343..e85d048 100644 --- a/docs/plans/option_b_interleaved_traversal.md +++ b/docs/plans/option_b_interleaved_traversal.md @@ -1,7 +1,7 @@ # Plan: Option B Per-Worker Interleaved Traversal -**Status:** Phase 3 benchmark passed. Default behavior is unchanged; the -interleaved path remains opt-in. +**Status:** Phase 4 long-run A/B is prepared. Default behavior is unchanged; +the interleaved path remains opt-in. **Owner:** Codex for prototype design and implementation; operator for long-run benchmarks on `home`. **Background:** Option A, the central traversal inference server, was implemented @@ -306,6 +306,81 @@ path with larger worker chunks. The single-process CUDA path makes forward cheap but gives up multiprocessing game-state throughput, so it is slower end-to-end. +### Phase 4: Feature Expansion And Long-Run A/B Prep + +Required feature expansion: + +- Support `opponent_policy: average_strategy`, matching the default config's + opponent branch. +- Keep unsupported branches guarded (`self_play_league`, `safe_heuristic`, + random rollout cutoffs, external sampling). +- Add parity tests for the average-strategy fixed-opponent branch. +- Verify a default-policy interleaved run starts and emits batch metrics. + +Long-run A/B protocol: + +- Baseline: `configs/deep_cfr/default.yaml` unchanged. +- Treatment: exactly one structural scheduler change plus the chunk/batch + settings below. +- Run sequentially, not in parallel on the same GPU. +- Keep default evaluation/checkpoint cadence for the real A/B unless measuring + pure throughput only. +- Stop and report if learning metrics drift materially despite wall-clock + speedup. + +Treatment command: + +```bash +uv run lost-cities-deep-cfr train \ + --config configs/deep_cfr/default.yaml \ + --keep \ + --set run.experiment_name=option-b-interleaved-default-ab \ + --set traversal.scheduler=interleaved \ + --set traversal.num_workers=8 \ + --set traversal.worker_chunk_size=64 \ + --set traversal.interleave_width=64 \ + --set traversal.interleave_max_batch=128 \ + --set traversal.progress_every_traversals=0 +``` + +### Phase 4 Result (2026-05-07) + +Implemented `average_strategy` support in the interleaved scheduler. Opponent +nodes now use the strategy network as a fixed-opponent policy, matching the +Cython recursive path rather than recording CFR regret at those nodes. + +Validation: + +- average-strategy single-traversal parity against `run_cython_traversal_batch`: + PASS, +- trainer smoke with `scheduler=interleaved` and `opponent_policy=average_strategy`: + PASS, +- focused trainer test suite: PASS. + +Default-policy throughput check: + +```bash +uv run lost-cities-deep-cfr train \ + --config configs/deep_cfr/default.yaml \ + --keep \ + --set run.max_iterations=10 \ + --set run.experiment_name=option-b-interleaved-average-strategy-10i \ + --set traversal.scheduler=interleaved \ + --set traversal.num_workers=8 \ + --set traversal.worker_chunk_size=64 \ + --set traversal.interleave_width=64 \ + --set traversal.interleave_max_batch=128 \ + --set traversal.progress_every_traversals=0 \ + --set checkpoint.save_latest=false \ + --set checkpoint.save_every=0 \ + --set evaluation.eval_every=0 +``` + +Result path: `runs/2026-05-07_230419_option-b-interleaved-average-strategy-10i`. +Warm-up-excluded means: `iteration_seconds=10.61s`, +`traversal_seconds=4.85s`, `interleaved/avg_batch_size=28.3`, +`interleaved/max_batch_size=64`. + ## Risks - **State-machine complexity.** Recursive CFR control flow has many local 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 3015cab..85fce3f 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -187,9 +187,10 @@ 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 != "network": + if self.opponent_policy not in {"network", "average_strategy"}: raise ValueError( - "scheduler='interleaved' currently supports only opponent_policy='network'" + "scheduler='interleaved' currently supports only " + "opponent_policy='network' or 'average_strategy'" ) if self.cutoff_rollouts != 0 or self.cutoff_value_mode != "score_diff": raise ValueError( diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py b/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py index 9d0cf75..105f2ee 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/interleaved_traversal.py @@ -76,6 +76,22 @@ def _sampling_policy(policy: np.ndarray, legal_mask: np.ndarray, epsilon: float) return out +def _masked_softmax(logits: np.ndarray, legal_mask: np.ndarray) -> np.ndarray: + policy = np.zeros_like(logits, dtype=np.float32) + legal = np.flatnonzero(legal_mask) + if len(legal) == 0: + return policy + legal_logits = logits[legal].astype(np.float32) + shifted = legal_logits - float(np.max(legal_logits)) + values = np.exp(shifted, dtype=np.float32) + total = float(values.sum()) + if total <= 0.0: + policy[legal] = 1.0 / float(len(legal)) + return policy + policy[legal] = values / total + return policy + + def _record_endpoint(stats: TraversalStats, depth: int, width: int, max_depth: int) -> None: stats.endpoint_depth_sum += depth start = (depth // width) * width @@ -95,6 +111,7 @@ class InterleavedTraversalConfig: strategy_sample_interval: int store_strategy_on_traverser_nodes: bool store_strategy_on_opponent_nodes: bool + opponent_policy: str endpoint_depth_bucket_width: int endpoint_depth_bucket_max: int @@ -109,6 +126,7 @@ class PolicyResult: full_tie: bool player: int = -1 depth: int = -1 + kind: str = "advantage" @dataclass @@ -118,6 +136,7 @@ class PolicyRequest: info_state: np.ndarray legal_mask: np.ndarray depth: int + network_kind: str = "advantage" @dataclass @@ -142,12 +161,17 @@ class AfterChildFrame: swapped_deck_index: int +@dataclass +class FixedActionFrame: + swapped_deck_index: int + + @dataclass class EnterFrame: depth: int -Frame = EnterFrame | AfterChildFrame +Frame = EnterFrame | AfterChildFrame | FixedActionFrame class BatchedPolicy: @@ -157,8 +181,10 @@ class BatchedPolicy: *, device: torch.device, epsilon: float, + strategy_network: torch.nn.Module | None = None, ) -> None: self.networks = networks + self.strategy_network = strategy_network self.device = device self.epsilon = epsilon self.batch_sizes: list[int] = [] @@ -168,24 +194,36 @@ class BatchedPolicy: if not requests: return [] out: list[PolicyResult | None] = [None] * len(requests) - for player in sorted({request.player for request in requests}): - indices = [idx for idx, request in enumerate(requests) if request.player == player] + group_keys = sorted({(request.network_kind, request.player) for request in requests}) + for network_kind, player in group_keys: + indices = [ + idx + for idx, request in enumerate(requests) + if (request.network_kind, request.player) == (network_kind, player) + ] states = np.stack([requests[idx].info_state for idx in indices]).astype(np.float32) x = torch.as_tensor(states, dtype=torch.float32, device=self.device) if self.device.type == "cuda": torch.cuda.synchronize(self.device) start = time.perf_counter() with torch.inference_mode(): - values = self.networks[player](x).detach().cpu().numpy().astype(np.float32) + network = self._network(network_kind, player) + values = network(x).detach().cpu().numpy().astype(np.float32) if self.device.type == "cuda": torch.cuda.synchronize(self.device) self.forward_seconds += time.perf_counter() - start self.batch_sizes.append(len(indices)) for local_idx, request_idx in enumerate(indices): request = requests[request_idx] - policy, fallback, tie_size, full_tie = _regret_matching( - values[local_idx], request.legal_mask, self.epsilon - ) + if request.network_kind == "strategy": + policy = _masked_softmax(values[local_idx], request.legal_mask) + fallback = False + tie_size = 0 + full_tie = False + else: + policy, fallback, tie_size, full_tie = _regret_matching( + values[local_idx], request.legal_mask, self.epsilon + ) out[request_idx] = PolicyResult( info_state=request.info_state, legal_mask=request.legal_mask, @@ -196,6 +234,13 @@ class BatchedPolicy: ) return [result for result in out if result is not None] + def _network(self, network_kind: str, player: int) -> torch.nn.Module: + if network_kind == "advantage": + return self.networks[player] + if network_kind == "strategy" and self.strategy_network is not None: + return self.strategy_network + raise ValueError(f"unsupported policy network request: {network_kind!r}") + class InterleavedContext: def __init__( @@ -225,8 +270,10 @@ class InterleavedContext: frame = self.stack.pop() if isinstance(frame, EnterFrame): self._enter(frame.depth, context_index) - else: + elif isinstance(frame, AfterChildFrame): self._after_child(frame) + else: + self._after_fixed_action(frame) if not self.stack and self.pending is None and not self.done: self.done = True self.value = self.last_value @@ -239,7 +286,6 @@ class InterleavedContext: depth = result.depth if depth < 0: raise RuntimeError("policy result is missing depth") - self._record_strategy(result, player, depth) actions = [int(action) for action in np.flatnonzero(result.legal_mask)] if not actions: self.stats.terminals += 1 @@ -252,6 +298,16 @@ class InterleavedContext: self._return_value(float(self.state.score_diff(self.traverser))) return + if result.kind == "strategy": + self.rng, random_value = _next_double(self.rng) + action = _sample_policy(result.policy, actions, random_value) + 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 + + self._record_strategy(result, player, depth) sampling_policy = _sampling_policy( result.policy, result.legal_mask, self.cfg.outcome_sampling_epsilon ) @@ -289,7 +345,12 @@ class InterleavedContext: info_state = encode_info_state(self.state, player, self.cfg.encoding) legal_mask = np.zeros(self.cfg.action_size, dtype=bool) legal_mask[self.state.unified_legal_actions()] = True - request = PolicyRequest(context_index, player, info_state, legal_mask, depth) + network_kind = ( + "strategy" + if player != self.traverser and self.cfg.opponent_policy == "average_strategy" + else "advantage" + ) + request = PolicyRequest(context_index, player, info_state, legal_mask, depth, network_kind) self.pending = request def _after_child(self, frame: AfterChildFrame) -> None: @@ -320,6 +381,13 @@ class InterleavedContext: self.stats.advantage_samples += 1 self._return_value(node_value) + def _after_fixed_action(self, frame: FixedActionFrame) -> None: + child_value = self.last_value + self.state.pop_action() + if frame.swapped_deck_index >= 0: + self.state.swap_deck_cards(frame.swapped_deck_index, len(self.state.deck) - 1) + self._return_value(child_value) + def _sample_deck_draw_chance(self, unified_action: int) -> int: deck_draw_action = 2 * self.state.config.hand_size deck_len = len(self.state.deck) @@ -496,6 +564,7 @@ class InterleavedTraversalScheduler: for context_index, request, result in zip( request_contexts, requests, results, strict=True ): + result.kind = request.network_kind result.player = request.player result.depth = request.depth contexts[context_index].apply_policy(result) @@ -513,6 +582,7 @@ class InterleavedTraversalScheduler: def run_interleaved_traversal_batch( advantage_networks: list[torch.nn.Module], + strategy_network: torch.nn.Module | None, game_config: Any, seeds: list[int], player: int, @@ -529,6 +599,7 @@ def run_interleaved_traversal_batch( max_nodes: int | None, outcome_sampling_epsilon: float, outcome_sampling_value_clip: float | None, + opponent_policy: str, endpoint_depth_bucket_width: int, endpoint_depth_bucket_max: int, seed: int, @@ -546,12 +617,18 @@ def run_interleaved_traversal_batch( strategy_sample_interval=strategy_sample_interval, store_strategy_on_traverser_nodes=store_strategy_on_traverser_nodes, store_strategy_on_opponent_nodes=store_strategy_on_opponent_nodes, + opponent_policy=opponent_policy, endpoint_depth_bucket_width=endpoint_depth_bucket_width, endpoint_depth_bucket_max=endpoint_depth_bucket_max, ) states = [GameState.new_game(game_config, seed=game_seed) for game_seed in seeds] rng_seeds = [int(seed) + index * 1_000_003 for index in range(len(seeds))] - policy = BatchedPolicy(advantage_networks, device=device, epsilon=cfg.epsilon) + policy = BatchedPolicy( + advantage_networks, + device=device, + epsilon=cfg.epsilon, + strategy_network=strategy_network, + ) scheduler = InterleavedTraversalScheduler(cfg, policy) _values, _rng_out, stats_rows, sample_rows, batch_sizes = scheduler.run( states, 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 254d766..c3e3259 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -401,6 +401,11 @@ class DeepCFRTrainer: stats, advantage_samples, strategy_samples, runtime_metrics = ( run_interleaved_traversal_batch( self.advantage_networks, + ( + self.strategy_network + if self.config.traversal.opponent_policy == "average_strategy" + else None + ), self.game_config, seeds, player, @@ -422,6 +427,7 @@ class DeepCFRTrainer: outcome_sampling_value_clip=( self.config.traversal.outcome_sampling_value_clip ), + opponent_policy=self.config.traversal.opponent_policy, endpoint_depth_bucket_width=( self.config.traversal.endpoint_depth_bucket_width ), 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 da6fb89..b9c8878 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -150,6 +150,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe total_stats, advantage_samples, strategy_samples, runtime_metrics = ( run_interleaved_traversal_batch( networks, + strategy_network, game_config, batch.seeds, batch.player, @@ -169,6 +170,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe max_nodes=cfg.traversal.max_nodes_per_traversal, outcome_sampling_epsilon=cfg.traversal.outcome_sampling_epsilon, outcome_sampling_value_clip=cfg.traversal.outcome_sampling_value_clip, + opponent_policy=cfg.traversal.opponent_policy, endpoint_depth_bucket_width=cfg.traversal.endpoint_depth_bucket_width, endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max, seed=batch.worker_seed, diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 65e9e71..8db4ae8 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -128,15 +128,15 @@ def test_deep_cfr_config_accepts_external_sampling_mode() -> None: def test_deep_cfr_config_accepts_interleaved_scheduler() -> None: config = _deep_cfr_config( - {"traversal": {"scheduler": "interleaved", "opponent_policy": "network"}} + {"traversal": {"scheduler": "interleaved", "opponent_policy": "average_strategy"}} ) assert config.traversal.scheduler == "interleaved" - assert config.traversal.opponent_policy == "network" + assert config.traversal.opponent_policy == "average_strategy" def test_deep_cfr_config_rejects_unsupported_interleaved_options() -> None: - with pytest.raises(ValueError, match="opponent_policy='network'"): + with pytest.raises(ValueError, match="opponent_policy='network' or 'average_strategy'"): _deep_cfr_config( {"traversal": {"scheduler": "interleaved", "opponent_policy": "self_play_league"}} ) @@ -516,11 +516,110 @@ def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() -> interleaved_stats, interleaved_advantage, interleaved_strategy, _runtime = ( run_interleaved_traversal_batch( networks, + None, game_config, [260], 0, 1, **common, + opponent_policy=config.traversal.opponent_policy, + interleave_width=4, + interleave_max_batch=8, + ) + ) + + assert interleaved_stats.to_dict() == recursive_stats.to_dict() + assert len(interleaved_advantage) == len(recursive_advantage) + assert len(interleaved_strategy) == len(recursive_strategy) + assert np.allclose( + [sample.target.sum() for sample in interleaved_advantage], + [sample.target.sum() for sample in recursive_advantage], + atol=1.0e-5, + ) + assert np.allclose( + [sample.target.sum() for sample in interleaved_strategy], + [sample.target.sum() for sample in recursive_strategy], + atol=1.0e-6, + ) + + +def test_deep_cfr_interleaved_scheduler_matches_average_strategy_opponent() -> None: + config = _deep_cfr_config( + { + "run": {"seed": 27}, + "network": {"hidden_size": 16}, + "traversal": { + "opponent_policy": "average_strategy", + "store_strategy_on_opponent_nodes": False, + "max_depth": 4, + "max_nodes_per_traversal": 64, + }, + } + ) + game_config = LostCitiesConfig(seed=27) + probe = GameState.new_game(game_config, seed=27) + action_size = game_config.action_size + torch.manual_seed(27) + networks = [ + DeepCFRMLP.from_config(input_dim(probe, config.encoding), action_size, config.network) + for _ in range(2) + ] + strategy_network = DeepCFRMLP.from_config( + input_dim(probe, config.encoding), action_size, config.network + ) + for network in [*networks, strategy_network]: + network.eval() + common = { + "device": torch.device("cpu"), + "action_size": action_size, + "encoding": config.encoding, + "epsilon": config.traversal.regret_matching_epsilon, + "strategy_sample_interval": config.traversal.strategy_sample_interval, + "store_strategy_on_traverser_nodes": config.traversal.store_strategy_on_traverser_nodes, + "store_strategy_on_opponent_nodes": config.traversal.store_strategy_on_opponent_nodes, + "max_depth": config.traversal.max_depth, + "max_nodes": config.traversal.max_nodes_per_traversal, + "outcome_sampling_epsilon": config.traversal.outcome_sampling_epsilon, + "outcome_sampling_value_clip": config.traversal.outcome_sampling_value_clip, + "endpoint_depth_bucket_width": config.traversal.endpoint_depth_bucket_width, + "endpoint_depth_bucket_max": config.traversal.endpoint_depth_bucket_max, + "seed": 2701, + } + + recursive_stats, recursive_advantage, recursive_strategy = run_cython_traversal_batch( + networks, + game_config, + [270], + 0, + 1, + **common, + strategy_network=strategy_network, + sampling_mode=config.traversal.sampling_mode, + outcome_unsampled_regret=config.traversal.outcome_unsampled_regret, + cutoff_value_mode=config.traversal.cutoff_value_mode, + cutoff_rollouts=config.traversal.cutoff_rollouts, + cutoff_rollout_policy=config.traversal.cutoff_rollout_policy, + cutoff_rollout_max_steps=config.traversal.cutoff_rollout_max_steps, + opponent_policy=config.traversal.opponent_policy, + all_negative_fallback=config.regret_matching.all_negative_fallback, + league_advantage_networks=[], + self_play_anchor_probability=config.self_play.anchor_probability, + self_play_current_weight=config.self_play.current_weight, + self_play_recent_weight=config.self_play.recent_weight, + self_play_older_weight=config.self_play.older_weight, + self_play_anchor_weight=config.self_play.anchor_weight, + self_play_recent_window=config.self_play.recent_window, + ) + interleaved_stats, interleaved_advantage, interleaved_strategy, _runtime = ( + run_interleaved_traversal_batch( + networks, + strategy_network, + game_config, + [270], + 0, + 1, + **common, + opponent_policy=config.traversal.opponent_policy, interleave_width=4, interleave_max_batch=8, )