Support average strategy in interleaved traversal

Co-Authored-By: Codex <codex@openai.com>
This commit is contained in:
2026-05-07 23:07:47 +09:00
co-authored by Codex
parent 2ce70ea804
commit 102f8cc91d
7 changed files with 292 additions and 18 deletions
+14
View File
@@ -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:
+77 -2
View File
@@ -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,24 +194,36 @@ 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]
policy, fallback, tie_size, full_tie = _regret_matching( if request.network_kind == "strategy":
values[local_idx], request.legal_mask, self.epsilon 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( out[request_idx] = PolicyResult(
info_state=request.info_state, info_state=request.info_state,
legal_mask=request.legal_mask, legal_mask=request.legal_mask,
@@ -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,
+102 -3
View File
@@ -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,
) )