Support average strategy in interleaved traversal
Co-Authored-By: Codex <codex@openai.com>
This commit is contained in:
@@ -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
|
checksums under identical RNG seed. Longer learning-curve A/B is still required
|
||||||
before considering a default switch.
|
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)
|
## Batched Traversal Inference: Design Decision (2026-05-07)
|
||||||
|
|
||||||
Three structural options were considered for Priority #5:
|
Three structural options were considered for Priority #5:
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# Plan: Option B Per-Worker Interleaved Traversal
|
# Plan: Option B Per-Worker Interleaved Traversal
|
||||||
|
|
||||||
**Status:** Phase 3 benchmark passed. Default behavior is unchanged; the
|
**Status:** Phase 4 long-run A/B is prepared. Default behavior is unchanged;
|
||||||
interleaved path remains opt-in.
|
the interleaved path remains opt-in.
|
||||||
**Owner:** Codex for prototype design and implementation; operator for long-run
|
**Owner:** Codex for prototype design and implementation; operator for long-run
|
||||||
benchmarks on `home`.
|
benchmarks on `home`.
|
||||||
**Background:** Option A, the central traversal inference server, was implemented
|
**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
|
cheap but gives up multiprocessing game-state throughput, so it is slower
|
||||||
end-to-end.
|
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
|
## Risks
|
||||||
|
|
||||||
- **State-machine complexity.** Recursive CFR control flow has many local
|
- **State-machine complexity.** Recursive CFR control flow has many local
|
||||||
|
|||||||
@@ -187,9 +187,10 @@ class TraversalConfig(StrictModel):
|
|||||||
if self.scheduler == "interleaved":
|
if self.scheduler == "interleaved":
|
||||||
if self.sampling_mode != "outcome":
|
if self.sampling_mode != "outcome":
|
||||||
raise ValueError("scheduler='interleaved' currently supports only outcome sampling")
|
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(
|
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":
|
if self.cutoff_rollouts != 0 or self.cutoff_value_mode != "score_diff":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -76,6 +76,22 @@ def _sampling_policy(policy: np.ndarray, legal_mask: np.ndarray, epsilon: float)
|
|||||||
return out
|
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:
|
def _record_endpoint(stats: TraversalStats, depth: int, width: int, max_depth: int) -> None:
|
||||||
stats.endpoint_depth_sum += depth
|
stats.endpoint_depth_sum += depth
|
||||||
start = (depth // width) * width
|
start = (depth // width) * width
|
||||||
@@ -95,6 +111,7 @@ class InterleavedTraversalConfig:
|
|||||||
strategy_sample_interval: int
|
strategy_sample_interval: int
|
||||||
store_strategy_on_traverser_nodes: bool
|
store_strategy_on_traverser_nodes: bool
|
||||||
store_strategy_on_opponent_nodes: bool
|
store_strategy_on_opponent_nodes: bool
|
||||||
|
opponent_policy: str
|
||||||
endpoint_depth_bucket_width: int
|
endpoint_depth_bucket_width: int
|
||||||
endpoint_depth_bucket_max: int
|
endpoint_depth_bucket_max: int
|
||||||
|
|
||||||
@@ -109,6 +126,7 @@ class PolicyResult:
|
|||||||
full_tie: bool
|
full_tie: bool
|
||||||
player: int = -1
|
player: int = -1
|
||||||
depth: int = -1
|
depth: int = -1
|
||||||
|
kind: str = "advantage"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -118,6 +136,7 @@ class PolicyRequest:
|
|||||||
info_state: np.ndarray
|
info_state: np.ndarray
|
||||||
legal_mask: np.ndarray
|
legal_mask: np.ndarray
|
||||||
depth: int
|
depth: int
|
||||||
|
network_kind: str = "advantage"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -142,12 +161,17 @@ class AfterChildFrame:
|
|||||||
swapped_deck_index: int
|
swapped_deck_index: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FixedActionFrame:
|
||||||
|
swapped_deck_index: int
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EnterFrame:
|
class EnterFrame:
|
||||||
depth: int
|
depth: int
|
||||||
|
|
||||||
|
|
||||||
Frame = EnterFrame | AfterChildFrame
|
Frame = EnterFrame | AfterChildFrame | FixedActionFrame
|
||||||
|
|
||||||
|
|
||||||
class BatchedPolicy:
|
class BatchedPolicy:
|
||||||
@@ -157,8 +181,10 @@ class BatchedPolicy:
|
|||||||
*,
|
*,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
epsilon: float,
|
epsilon: float,
|
||||||
|
strategy_network: torch.nn.Module | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.networks = networks
|
self.networks = networks
|
||||||
|
self.strategy_network = strategy_network
|
||||||
self.device = device
|
self.device = device
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
self.batch_sizes: list[int] = []
|
self.batch_sizes: list[int] = []
|
||||||
@@ -168,21 +194,33 @@ class BatchedPolicy:
|
|||||||
if not requests:
|
if not requests:
|
||||||
return []
|
return []
|
||||||
out: list[PolicyResult | None] = [None] * len(requests)
|
out: list[PolicyResult | None] = [None] * len(requests)
|
||||||
for player in sorted({request.player for request in requests}):
|
group_keys = sorted({(request.network_kind, request.player) for request in requests})
|
||||||
indices = [idx for idx, request in enumerate(requests) if request.player == player]
|
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)
|
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)
|
x = torch.as_tensor(states, dtype=torch.float32, device=self.device)
|
||||||
if self.device.type == "cuda":
|
if self.device.type == "cuda":
|
||||||
torch.cuda.synchronize(self.device)
|
torch.cuda.synchronize(self.device)
|
||||||
start = time.perf_counter()
|
start = time.perf_counter()
|
||||||
with torch.inference_mode():
|
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":
|
if self.device.type == "cuda":
|
||||||
torch.cuda.synchronize(self.device)
|
torch.cuda.synchronize(self.device)
|
||||||
self.forward_seconds += time.perf_counter() - start
|
self.forward_seconds += time.perf_counter() - start
|
||||||
self.batch_sizes.append(len(indices))
|
self.batch_sizes.append(len(indices))
|
||||||
for local_idx, request_idx in enumerate(indices):
|
for local_idx, request_idx in enumerate(indices):
|
||||||
request = requests[request_idx]
|
request = requests[request_idx]
|
||||||
|
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(
|
policy, fallback, tie_size, full_tie = _regret_matching(
|
||||||
values[local_idx], request.legal_mask, self.epsilon
|
values[local_idx], request.legal_mask, self.epsilon
|
||||||
)
|
)
|
||||||
@@ -196,6 +234,13 @@ class BatchedPolicy:
|
|||||||
)
|
)
|
||||||
return [result for result in out if result is not None]
|
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:
|
class InterleavedContext:
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -225,8 +270,10 @@ class InterleavedContext:
|
|||||||
frame = self.stack.pop()
|
frame = self.stack.pop()
|
||||||
if isinstance(frame, EnterFrame):
|
if isinstance(frame, EnterFrame):
|
||||||
self._enter(frame.depth, context_index)
|
self._enter(frame.depth, context_index)
|
||||||
else:
|
elif isinstance(frame, AfterChildFrame):
|
||||||
self._after_child(frame)
|
self._after_child(frame)
|
||||||
|
else:
|
||||||
|
self._after_fixed_action(frame)
|
||||||
if not self.stack and self.pending is None and not self.done:
|
if not self.stack and self.pending is None and not self.done:
|
||||||
self.done = True
|
self.done = True
|
||||||
self.value = self.last_value
|
self.value = self.last_value
|
||||||
@@ -239,7 +286,6 @@ class InterleavedContext:
|
|||||||
depth = result.depth
|
depth = result.depth
|
||||||
if depth < 0:
|
if depth < 0:
|
||||||
raise RuntimeError("policy result is missing depth")
|
raise RuntimeError("policy result is missing depth")
|
||||||
self._record_strategy(result, player, depth)
|
|
||||||
actions = [int(action) for action in np.flatnonzero(result.legal_mask)]
|
actions = [int(action) for action in np.flatnonzero(result.legal_mask)]
|
||||||
if not actions:
|
if not actions:
|
||||||
self.stats.terminals += 1
|
self.stats.terminals += 1
|
||||||
@@ -252,6 +298,16 @@ class InterleavedContext:
|
|||||||
self._return_value(float(self.state.score_diff(self.traverser)))
|
self._return_value(float(self.state.score_diff(self.traverser)))
|
||||||
return
|
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(
|
sampling_policy = _sampling_policy(
|
||||||
result.policy, result.legal_mask, self.cfg.outcome_sampling_epsilon
|
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)
|
info_state = encode_info_state(self.state, player, self.cfg.encoding)
|
||||||
legal_mask = np.zeros(self.cfg.action_size, dtype=bool)
|
legal_mask = np.zeros(self.cfg.action_size, dtype=bool)
|
||||||
legal_mask[self.state.unified_legal_actions()] = True
|
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
|
self.pending = request
|
||||||
|
|
||||||
def _after_child(self, frame: AfterChildFrame) -> None:
|
def _after_child(self, frame: AfterChildFrame) -> None:
|
||||||
@@ -320,6 +381,13 @@ class InterleavedContext:
|
|||||||
self.stats.advantage_samples += 1
|
self.stats.advantage_samples += 1
|
||||||
self._return_value(node_value)
|
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:
|
def _sample_deck_draw_chance(self, unified_action: int) -> int:
|
||||||
deck_draw_action = 2 * self.state.config.hand_size
|
deck_draw_action = 2 * self.state.config.hand_size
|
||||||
deck_len = len(self.state.deck)
|
deck_len = len(self.state.deck)
|
||||||
@@ -496,6 +564,7 @@ class InterleavedTraversalScheduler:
|
|||||||
for context_index, request, result in zip(
|
for context_index, request, result in zip(
|
||||||
request_contexts, requests, results, strict=True
|
request_contexts, requests, results, strict=True
|
||||||
):
|
):
|
||||||
|
result.kind = request.network_kind
|
||||||
result.player = request.player
|
result.player = request.player
|
||||||
result.depth = request.depth
|
result.depth = request.depth
|
||||||
contexts[context_index].apply_policy(result)
|
contexts[context_index].apply_policy(result)
|
||||||
@@ -513,6 +582,7 @@ class InterleavedTraversalScheduler:
|
|||||||
|
|
||||||
def run_interleaved_traversal_batch(
|
def run_interleaved_traversal_batch(
|
||||||
advantage_networks: list[torch.nn.Module],
|
advantage_networks: list[torch.nn.Module],
|
||||||
|
strategy_network: torch.nn.Module | None,
|
||||||
game_config: Any,
|
game_config: Any,
|
||||||
seeds: list[int],
|
seeds: list[int],
|
||||||
player: int,
|
player: int,
|
||||||
@@ -529,6 +599,7 @@ def run_interleaved_traversal_batch(
|
|||||||
max_nodes: int | None,
|
max_nodes: int | None,
|
||||||
outcome_sampling_epsilon: float,
|
outcome_sampling_epsilon: float,
|
||||||
outcome_sampling_value_clip: float | None,
|
outcome_sampling_value_clip: float | None,
|
||||||
|
opponent_policy: str,
|
||||||
endpoint_depth_bucket_width: int,
|
endpoint_depth_bucket_width: int,
|
||||||
endpoint_depth_bucket_max: int,
|
endpoint_depth_bucket_max: int,
|
||||||
seed: int,
|
seed: int,
|
||||||
@@ -546,12 +617,18 @@ def run_interleaved_traversal_batch(
|
|||||||
strategy_sample_interval=strategy_sample_interval,
|
strategy_sample_interval=strategy_sample_interval,
|
||||||
store_strategy_on_traverser_nodes=store_strategy_on_traverser_nodes,
|
store_strategy_on_traverser_nodes=store_strategy_on_traverser_nodes,
|
||||||
store_strategy_on_opponent_nodes=store_strategy_on_opponent_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_width=endpoint_depth_bucket_width,
|
||||||
endpoint_depth_bucket_max=endpoint_depth_bucket_max,
|
endpoint_depth_bucket_max=endpoint_depth_bucket_max,
|
||||||
)
|
)
|
||||||
states = [GameState.new_game(game_config, seed=game_seed) for game_seed in seeds]
|
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))]
|
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)
|
scheduler = InterleavedTraversalScheduler(cfg, policy)
|
||||||
_values, _rng_out, stats_rows, sample_rows, batch_sizes = scheduler.run(
|
_values, _rng_out, stats_rows, sample_rows, batch_sizes = scheduler.run(
|
||||||
states,
|
states,
|
||||||
|
|||||||
@@ -401,6 +401,11 @@ class DeepCFRTrainer:
|
|||||||
stats, advantage_samples, strategy_samples, runtime_metrics = (
|
stats, advantage_samples, strategy_samples, runtime_metrics = (
|
||||||
run_interleaved_traversal_batch(
|
run_interleaved_traversal_batch(
|
||||||
self.advantage_networks,
|
self.advantage_networks,
|
||||||
|
(
|
||||||
|
self.strategy_network
|
||||||
|
if self.config.traversal.opponent_policy == "average_strategy"
|
||||||
|
else None
|
||||||
|
),
|
||||||
self.game_config,
|
self.game_config,
|
||||||
seeds,
|
seeds,
|
||||||
player,
|
player,
|
||||||
@@ -422,6 +427,7 @@ class DeepCFRTrainer:
|
|||||||
outcome_sampling_value_clip=(
|
outcome_sampling_value_clip=(
|
||||||
self.config.traversal.outcome_sampling_value_clip
|
self.config.traversal.outcome_sampling_value_clip
|
||||||
),
|
),
|
||||||
|
opponent_policy=self.config.traversal.opponent_policy,
|
||||||
endpoint_depth_bucket_width=(
|
endpoint_depth_bucket_width=(
|
||||||
self.config.traversal.endpoint_depth_bucket_width
|
self.config.traversal.endpoint_depth_bucket_width
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -150,6 +150,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
|||||||
total_stats, advantage_samples, strategy_samples, runtime_metrics = (
|
total_stats, advantage_samples, strategy_samples, runtime_metrics = (
|
||||||
run_interleaved_traversal_batch(
|
run_interleaved_traversal_batch(
|
||||||
networks,
|
networks,
|
||||||
|
strategy_network,
|
||||||
game_config,
|
game_config,
|
||||||
batch.seeds,
|
batch.seeds,
|
||||||
batch.player,
|
batch.player,
|
||||||
@@ -169,6 +170,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
|||||||
max_nodes=cfg.traversal.max_nodes_per_traversal,
|
max_nodes=cfg.traversal.max_nodes_per_traversal,
|
||||||
outcome_sampling_epsilon=cfg.traversal.outcome_sampling_epsilon,
|
outcome_sampling_epsilon=cfg.traversal.outcome_sampling_epsilon,
|
||||||
outcome_sampling_value_clip=cfg.traversal.outcome_sampling_value_clip,
|
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_width=cfg.traversal.endpoint_depth_bucket_width,
|
||||||
endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max,
|
endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max,
|
||||||
seed=batch.worker_seed,
|
seed=batch.worker_seed,
|
||||||
|
|||||||
@@ -128,15 +128,15 @@ def test_deep_cfr_config_accepts_external_sampling_mode() -> None:
|
|||||||
|
|
||||||
def test_deep_cfr_config_accepts_interleaved_scheduler() -> None:
|
def test_deep_cfr_config_accepts_interleaved_scheduler() -> None:
|
||||||
config = _deep_cfr_config(
|
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.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:
|
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(
|
_deep_cfr_config(
|
||||||
{"traversal": {"scheduler": "interleaved", "opponent_policy": "self_play_league"}}
|
{"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 = (
|
interleaved_stats, interleaved_advantage, interleaved_strategy, _runtime = (
|
||||||
run_interleaved_traversal_batch(
|
run_interleaved_traversal_batch(
|
||||||
networks,
|
networks,
|
||||||
|
None,
|
||||||
game_config,
|
game_config,
|
||||||
[260],
|
[260],
|
||||||
0,
|
0,
|
||||||
1,
|
1,
|
||||||
**common,
|
**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_width=4,
|
||||||
interleave_max_batch=8,
|
interleave_max_batch=8,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user