Add non-default interleaved traversal scheduler

Co-Authored-By: Codex <codex@openai.com>
This commit is contained in:
2026-05-07 22:49:38 +09:00
co-authored by Codex
parent 09d58159b8
commit 240ef552c6
6 changed files with 1014 additions and 87 deletions
+176 -2
View File
@@ -7,7 +7,10 @@ import numpy as np
import pytest
import torch
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim
from coolrl_lost_cities.games.classic.deep_cfr.traversal import CythonDeepCFRTraverser
from coolrl_lost_cities.games.classic.deep_cfr.traversal import (
CythonDeepCFRTraverser,
run_cython_traversal_batch,
)
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
@@ -23,6 +26,9 @@ from coolrl_lost_cities.games.classic.deep_cfr.cli import (
)
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load_config
from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy_network
from coolrl_lost_cities.games.classic.deep_cfr.interleaved_traversal import (
run_interleaved_traversal_batch,
)
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.deep_cfr.trainer import (
@@ -108,11 +114,44 @@ def test_deep_cfr_train_cli_accepts_run_and_traversal_config_overrides() -> None
def test_deep_cfr_config_accepts_external_sampling_mode() -> None:
config = _deep_cfr_config({"traversal": {"sampling_mode": "external"}})
config = _deep_cfr_config(
{
"traversal": {
"sampling_mode": "external",
"store_strategy_on_traverser_nodes": False,
}
}
)
assert config.traversal.sampling_mode == "external"
def test_deep_cfr_config_accepts_interleaved_scheduler() -> None:
config = _deep_cfr_config(
{"traversal": {"scheduler": "interleaved", "opponent_policy": "network"}}
)
assert config.traversal.scheduler == "interleaved"
assert config.traversal.opponent_policy == "network"
def test_deep_cfr_config_rejects_unsupported_interleaved_options() -> None:
with pytest.raises(ValueError, match="opponent_policy='network'"):
_deep_cfr_config(
{"traversal": {"scheduler": "interleaved", "opponent_policy": "self_play_league"}}
)
with pytest.raises(ValueError, match="requires inference_backend='local'"):
_deep_cfr_config(
{
"traversal": {
"scheduler": "interleaved",
"opponent_policy": "network",
"inference_backend": "server",
}
}
)
def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None:
args = type(
"Args",
@@ -138,6 +177,7 @@ def test_deep_cfr_train_cli_accepts_generic_config_overrides() -> None:
{
"config_overrides": [
"traversal.sampling_mode=external",
"traversal.store_strategy_on_traverser_nodes=false",
"traversal.max_depth=null",
"optimization.advantage_batch_size=64",
"checkpoint.save_latest=true",
@@ -368,6 +408,139 @@ def test_deep_cfr_trainer_smoke_run() -> None:
assert metrics[0].strategy_loss >= 0.0
def test_deep_cfr_trainer_interleaved_scheduler_smoke_run(tmp_path) -> None:
trainer = DeepCFRTrainer(
_deep_cfr_config(
{
"run": {"max_iterations": 1, "seed": 24},
"network": {"hidden_size": 16},
"traversal": {
"scheduler": "interleaved",
"opponent_policy": "network",
"traversals_per_player": 2,
"max_depth": 3,
"max_nodes_per_traversal": 64,
"interleave_width": 4,
"interleave_max_batch": 8,
},
"optimization": {
"advantage_updates_per_iteration": 1,
"strategy_updates_per_iteration": 1,
"advantage_batch_size": 2,
"strategy_batch_size": 2,
},
"checkpoint": {"save_every": 0, "save_latest": False},
"evaluation": {"eval_every": 0},
}
),
LostCitiesConfig(seed=24),
run_dir=tmp_path / "interleaved",
)
metrics = trainer.train()
runtime = metrics[0].runtime_metrics
assert len(metrics) == 1
assert metrics[0].advantage_samples > 0
assert metrics[0].strategy_samples > 0
assert metrics[0].traversal_nodes > 0
assert runtime["interleaved/batches"] > 0
assert runtime["interleaved/requests"] > 0
assert runtime["interleaved/max_batch_size"] >= 1
assert runtime["interleaved/avg_batch_size"] >= 1.0
def test_deep_cfr_interleaved_scheduler_matches_recursive_single_traversal() -> None:
config = _deep_cfr_config(
{
"run": {"seed": 26},
"network": {"hidden_size": 16},
"traversal": {
"opponent_policy": "network",
"max_depth": 3,
"max_nodes_per_traversal": 64,
},
}
)
game_config = LostCitiesConfig(seed=26)
probe = GameState.new_game(game_config, seed=26)
action_size = game_config.action_size
torch.manual_seed(26)
networks = [
DeepCFRMLP.from_config(input_dim(probe, config.encoding), action_size, config.network)
for _ in range(2)
]
for network in networks:
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": 2601,
}
recursive_stats, recursive_advantage, recursive_strategy = run_cython_traversal_batch(
networks,
game_config,
[260],
0,
1,
**common,
strategy_network=None,
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,
game_config,
[260],
0,
1,
**common,
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_trainer_supports_lcfr_and_dcfr_loss_weighting() -> None:
for mode in ("lcfr", "dcfr"):
trainer = DeepCFRTrainer(
@@ -578,6 +751,7 @@ def test_deep_cfr_cython_traverser_supports_external_sampling() -> None:
"traversal": {
"traversals_per_player": 1,
"sampling_mode": "external",
"store_strategy_on_traverser_nodes": False,
"max_depth": 1,
"max_nodes_per_traversal": 64,
},