Add non-default interleaved traversal scheduler
Co-Authored-By: Codex <codex@openai.com>
This commit is contained in:
@@ -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,
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user