Add batched traversal inference server (Option A) behind opt-in flag

Implements the central inference server pattern: a dedicated GPU
process owns advantage/strategy/league networks, batches policy
requests across traversal workers via shared-memory tensor pool, and
returns logits. Workers route forward calls through InferenceClient /
NetworkProxy when traversal.inference_backend == "server".

Default remains traversal.inference_backend: local. The server
backend regresses iter time ~3.8× on the inspected default config
(small-model dispatch + sync-blocking traversal capping realized
batch at ~num_workers=8 instead of the bs=64-256 needed to amortize
IPC overhead). Keeping the implementation behind the flag lets us
re-enable when (a) model size grows, (b) per-worker interleaved
traversal lands, or (c) eval becomes dominant — see
docs/performance.md "Option A Bench Result and Structural Ceiling"
for the full diagnosis.

Plumbing included:
- inference_buffers.py: shared-memory tensor pool with slot
  management.
- inference_client.py: per-worker client + NetworkProxy adapter for
  the existing traversal.pyx call sites.
- inference_server.py: spawn-context server process with
  batch-window aggregation, weight sync, shutdown sentinel.
- bench_inference_backend.py: A/B between local and server backends
  with eval/checkpoint disabled.
- test_inference_server.py: round-trip and integration tests.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-05-07 20:05:16 +09:00
co-authored by Claude Opus 4.7
parent 05de0e2a81
commit a7ab94e096
10 changed files with 1375 additions and 79 deletions
+9
View File
@@ -46,6 +46,7 @@ traversal:
progress_every_traversals: 10 progress_every_traversals: 10
endpoint_depth_bucket_width: 100 endpoint_depth_bucket_width: 100
endpoint_depth_bucket_max: 1000 endpoint_depth_bucket_max: 1000
inference_backend: local
regret_matching: regret_matching:
all_negative_fallback: argmax_tiebreak all_negative_fallback: argmax_tiebreak
@@ -97,3 +98,11 @@ evaluation:
batch_size: 64 batch_size: 64
device: trainer device: trainer
num_workers: 4 num_workers: 4
inference_server:
device: cuda
num_slots: null
max_batch: 256
batch_window_us: 200
weight_sync_every: 1
use_amp: false
+108
View File
@@ -0,0 +1,108 @@
run:
experiment_name: deep-cfr-default-server
seed: 79
max_iterations: 1000
max_minutes: null
device: cuda
use_amp: false
rules:
n_colors: 5
n_ranks: 9
min_rank: 2
n_handshakes: 3
hand_size: 8
expedition_penalty: -20
bonus_threshold: 8
bonus_amount: 20
encoding:
derived_playability: true
slot_aware_playability: true
network:
hidden_size: 512
num_layers: 3
activation: relu
traversal:
traversals_per_player: 280
max_depth: null
max_nodes_per_traversal: 1000
regret_matching_epsilon: 0.0001
outcome_sampling_epsilon: 0.2
outcome_sampling_value_clip: 500.0
outcome_unsampled_regret: zero
cutoff_value_mode: score_diff
cutoff_rollouts: 0
cutoff_rollout_policy: random
cutoff_rollout_max_steps: 300
opponent_policy: average_strategy
strategy_sample_interval: 1
store_strategy_on_traverser_nodes: true
store_strategy_on_opponent_nodes: false
num_workers: 8
worker_chunk_size: 8
progress_every_traversals: 10
endpoint_depth_bucket_width: 100
endpoint_depth_bucket_max: 1000
inference_backend: server
regret_matching:
all_negative_fallback: argmax_tiebreak
training_weighting:
mode: lcfr
self_play:
snapshot_every: 1
max_snapshots: 0
anchor_probability: 0.0
current_weight: 1.0
recent_weight: 0.0
older_weight: 0.0
anchor_weight: 0.0
recent_window: 5
optimization:
advantage_updates_per_iteration: 512
strategy_updates_per_iteration: 512
advantage_batch_size: 1024
strategy_batch_size: 1024
learning_rate: 1.0e-4
weight_decay: 0.0001
grad_clip: 1.0
memory:
advantage_capacity: 2000000
strategy_capacity: 2000000
checkpoint:
save_latest: true
save_every: 100
progress_interval_seconds: 20.0
exact_resume: false
evaluation:
eval_every: 25
games: 100
opponents:
- random
- passive_discard
- safe_heuristic
- safe_heuristic_loose
- safe_heuristic_strict
- noisy_safe
max_steps: 10000
on_max_steps: score_diff
batch_size: 64
device: trainer
num_workers: 4
inference_server:
device: cuda
num_slots: null
max_batch: 256
batch_window_us: 200
weight_sync_every: 1
use_amp: false
+277
View File
@@ -0,0 +1,277 @@
from __future__ import annotations
import argparse
import gc
import json
from datetime import datetime
from pathlib import Path
from typing import Any
import torch
from coolrl_lost_cities.games.classic.deep_cfr.cli import _with_overrides
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load_config
from coolrl_lost_cities.games.classic.deep_cfr.tracking import FileRunTracker
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
METRIC_KEYS = (
"time/iteration_seconds",
"time/traversal_seconds",
"time/advantage_train_seconds",
"time/strategy_train_seconds",
"time/memory_add_seconds",
"time/batch_tensor_seconds",
)
TABLE_COLUMNS = (
("iter", "time/iteration_seconds"),
("traversal", "time/traversal_seconds"),
("adv_train", "time/advantage_train_seconds"),
("strat_train", "time/strategy_train_seconds"),
("mem_add", "time/memory_add_seconds"),
("batch_tensor", "time/batch_tensor_seconds"),
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Benchmark Deep CFR traversal inference backends.")
parser.add_argument(
"--config-local",
default="configs/deep_cfr/default.yaml",
help="Config for the local traversal inference backend.",
)
parser.add_argument(
"--config-server",
default="configs/deep_cfr/default_server.yaml",
help="Config for the server traversal inference backend.",
)
parser.add_argument("--iterations", type=int, default=10)
parser.add_argument("--warmup", type=int, default=2)
parser.add_argument("--device", help="Override run.device for both backend configs.")
return parser.parse_args()
def _benchmark_overrides(
*,
backend: str,
iterations: int,
seed: int,
device: str | None,
) -> dict[str, Any]:
overrides: dict[str, Any] = {
"run": {
"max_iterations": iterations,
"max_minutes": None,
"seed": seed,
},
"traversal": {"inference_backend": backend},
"checkpoint": {
"save_latest": False,
"save_every": 0,
},
"evaluation": {"eval_every": 0},
}
if device is not None:
overrides["run"]["device"] = device
return overrides
def _load_benchmark_config(
path: str,
*,
backend: str,
iterations: int,
seed: int,
device: str | None,
) -> DeepCFRConfig:
config = load_config(path)
return _with_overrides(
config,
_benchmark_overrides(
backend=backend,
iterations=iterations,
seed=seed,
device=device,
),
)
def _read_metrics(path: Path) -> list[dict[str, float | int]]:
rows: list[dict[str, float | int]] = []
with path.open(encoding="utf-8") as handle:
for line in handle:
line = line.strip()
if line:
rows.append(json.loads(line))
return rows
def _mean(rows: list[dict[str, float | int]], key: str) -> float:
if not rows:
return 0.0
return sum(float(row.get(key, 0.0) or 0.0) for row in rows) / len(rows)
def _summarize(rows: list[dict[str, float | int]], *, warmup: int) -> dict[str, float]:
measured = rows[warmup:]
return {key: _mean(measured, key) for key in METRIC_KEYS}
def _clear_gpu_memory() -> None:
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _run_backend(
*,
name: str,
config: DeepCFRConfig,
run_dir: Path,
warmup: int,
) -> dict[str, Any]:
_clear_gpu_memory()
tracker = FileRunTracker(
log_path=run_dir / "train.log",
metrics_path=run_dir / "metrics.jsonl",
progress_path=run_dir / "runtime_progress.json",
)
trainer: DeepCFRTrainer | None = DeepCFRTrainer(
config,
config.rules.to_lost_cities_config(seed=config.run.seed),
run_dir=run_dir,
device=config.run.device,
tracker=tracker,
)
try:
trainer.train()
rows = _read_metrics(run_dir / "metrics.jsonl")
finally:
trainer = None
_clear_gpu_memory()
return {
"backend": name,
"run_dir": str(run_dir),
"raw": rows,
"mean": _summarize(rows, warmup=warmup),
}
def _speedups(local: dict[str, float], server: dict[str, float]) -> dict[str, float | None]:
result: dict[str, float | None] = {}
for key in METRIC_KEYS:
server_value = server.get(key, 0.0)
result[key] = None if server_value <= 0.0 else local.get(key, 0.0) / server_value
return result
def _format_seconds(value: float) -> str:
return f"{value:.2f}s"
def _format_multiplier(value: float | None) -> str:
if value is None:
return "n/a"
return f"{value:.2f}x"
def _print_backend_table(results: dict[str, Any]) -> None:
local = results["backends"]["local"]["mean"]
server = results["backends"]["server"]["mean"]
speedup = results["speedup"]
rows = [
("local", [_format_seconds(local[key]) for _label, key in TABLE_COLUMNS]),
("server", [_format_seconds(server[key]) for _label, key in TABLE_COLUMNS]),
("speedup", [_format_multiplier(speedup[key]) for _label, key in TABLE_COLUMNS]),
]
headers = ["Backend", *[label for label, _key in TABLE_COLUMNS]]
widths = [max(len(headers[0]), *(len(row[0]) for row in rows))]
for index, header in enumerate(headers[1:]):
widths.append(max(len(header), *(len(row[1][index]) for row in rows)))
print(" ".join(header.ljust(widths[index]) for index, header in enumerate(headers)))
for backend, values in rows:
cells = [backend.ljust(widths[0])]
cells.extend(value.rjust(widths[index + 1]) for index, value in enumerate(values))
print(" ".join(cells))
def _print_projection(results: dict[str, Any]) -> None:
print()
print("1000-iteration projection")
print("Backend hours")
for backend in ("local", "server"):
mean_iter = results["backends"][backend]["mean"]["time/iteration_seconds"]
hours = mean_iter * 1000.0 / 3600.0
print(f"{backend:<7} {hours:.2f}h")
def main() -> None:
args = parse_args()
if args.iterations <= 0:
raise SystemExit("--iterations must be positive")
if args.warmup < 0:
raise SystemExit("--warmup must be non-negative")
if args.warmup >= args.iterations:
raise SystemExit("--warmup must be less than --iterations")
local_base = load_config(args.config_local)
seed = local_base.run.seed
timestamp = datetime.now().strftime("%Y-%m-%d_%H%M%S")
bench_dir = Path("runs/bench") / f"{timestamp}_inference_backend"
bench_dir.mkdir(parents=True, exist_ok=False)
local_config = _load_benchmark_config(
args.config_local,
backend="local",
iterations=args.iterations,
seed=seed,
device=args.device,
)
server_config = _load_benchmark_config(
args.config_server,
backend="server",
iterations=args.iterations,
seed=seed,
device=args.device,
)
local_result = _run_backend(
name="local",
config=local_config,
run_dir=bench_dir / "local",
warmup=args.warmup,
)
server_result = _run_backend(
name="server",
config=server_config,
run_dir=bench_dir / "server",
warmup=args.warmup,
)
results = {
"config": {
"config_local": args.config_local,
"config_server": args.config_server,
"iterations": args.iterations,
"warmup": args.warmup,
"seed": seed,
"device_override": args.device,
},
"backends": {
"local": local_result,
"server": server_result,
},
"speedup": _speedups(local_result["mean"], server_result["mean"]),
}
results_path = bench_dir / "results.json"
results_path.write_text(json.dumps(results, indent=2, sort_keys=True), encoding="utf-8")
_print_backend_table(results)
_print_projection(results)
print()
print(f"Results written to: {results_path}")
if __name__ == "__main__":
main()
@@ -111,6 +111,7 @@ class TraversalConfig(StrictModel):
progress_every_traversals: int = 0 progress_every_traversals: int = 0
endpoint_depth_bucket_width: int = 100 endpoint_depth_bucket_width: int = 100
endpoint_depth_bucket_max: int = 1000 endpoint_depth_bucket_max: int = 1000
inference_backend: str = "local"
@field_validator("sampling_mode") @field_validator("sampling_mode")
@classmethod @classmethod
@@ -150,6 +151,14 @@ class TraversalConfig(StrictModel):
) )
return value return value
@field_validator("inference_backend")
@classmethod
def _validate_inference_backend(cls, value: str) -> str:
token = value.strip().lower()
if token not in {"local", "server"}:
raise ValueError("must be 'local' or 'server'")
return token
def resolved_num_workers(self, batches: int | None = None) -> int: def resolved_num_workers(self, batches: int | None = None) -> int:
if isinstance(self.num_workers, str): if isinstance(self.num_workers, str):
token = self.num_workers.strip().lower() token = self.num_workers.strip().lower()
@@ -256,6 +265,23 @@ class EvaluationConfig(StrictModel):
return workers return workers
class InferenceServerConfig(StrictModel):
device: str = "cuda"
num_slots: int | None = None
max_batch: int = 256
batch_window_us: int = 200
weight_sync_every: int = 1
use_amp: bool = False
@field_validator("device")
@classmethod
def _validate_device(cls, value: str) -> str:
token = value.strip().lower()
if token in {"auto", "cpu", "cuda"}:
return token
return token
class DeepCFRConfig(StrictModel): class DeepCFRConfig(StrictModel):
run: RunConfig = Field(default_factory=RunConfig) run: RunConfig = Field(default_factory=RunConfig)
rules: RulesConfig = Field(default_factory=RulesConfig) rules: RulesConfig = Field(default_factory=RulesConfig)
@@ -269,6 +295,7 @@ class DeepCFRConfig(StrictModel):
memory: MemoryConfig = Field(default_factory=MemoryConfig) memory: MemoryConfig = Field(default_factory=MemoryConfig)
checkpoint: CheckpointConfig = Field(default_factory=CheckpointConfig) checkpoint: CheckpointConfig = Field(default_factory=CheckpointConfig)
evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig) evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig)
inference_server: InferenceServerConfig = Field(default_factory=InferenceServerConfig)
def to_dict(self) -> dict[str, Any]: def to_dict(self) -> dict[str, Any]:
return self.model_dump(mode="json") return self.model_dump(mode="json")
@@ -0,0 +1,118 @@
from __future__ import annotations
import multiprocessing as mp
from dataclasses import dataclass
from multiprocessing.shared_memory import SharedMemory
from typing import Any
import numpy as np
@dataclass(frozen=True)
class InferenceClientHandles:
request_shm_name: str
response_shm_name: str
num_slots: int
input_dim: int
action_size: int
request_queue: Any
weight_queue: Any
free_slots: Any
ready_events: list[Any]
weight_sync_event: Any
stats_queue: Any
class InferenceBuffers:
def __init__(
self,
*,
num_slots: int,
input_dim: int,
action_size: int,
mp_context: mp.context.BaseContext | None = None,
) -> None:
self.num_slots = int(num_slots)
self.input_dim = int(input_dim)
self.action_size = int(action_size)
if self.num_slots <= 0:
raise ValueError("num_slots must be positive")
if self.input_dim <= 0:
raise ValueError("input_dim must be positive")
if self.action_size <= 0:
raise ValueError("action_size must be positive")
self._ctx = mp_context or mp.get_context("spawn")
request_nbytes = self.num_slots * self.input_dim * np.dtype(np.float32).itemsize
response_nbytes = self.num_slots * self.action_size * np.dtype(np.float32).itemsize
self._request_shm = SharedMemory(create=True, size=request_nbytes)
self._response_shm = SharedMemory(create=True, size=response_nbytes)
self.requests = np.ndarray(
(self.num_slots, self.input_dim),
dtype=np.float32,
buffer=self._request_shm.buf,
)
self.responses = np.ndarray(
(self.num_slots, self.action_size),
dtype=np.float32,
buffer=self._response_shm.buf,
)
self.requests.fill(0.0)
self.responses.fill(0.0)
self.request_queue = self._ctx.Queue()
self.weight_queue = self._ctx.Queue()
self.free_slots = self._ctx.Queue()
self.ready_events = [self._ctx.Event() for _ in range(self.num_slots)]
self.weight_sync_event = self._ctx.Event()
self.stats_queue = self._ctx.Queue()
for slot in range(self.num_slots):
self.free_slots.put(slot)
def handles(self) -> InferenceClientHandles:
return InferenceClientHandles(
request_shm_name=self._request_shm.name,
response_shm_name=self._response_shm.name,
num_slots=self.num_slots,
input_dim=self.input_dim,
action_size=self.action_size,
request_queue=self.request_queue,
weight_queue=self.weight_queue,
free_slots=self.free_slots,
ready_events=self.ready_events,
weight_sync_event=self.weight_sync_event,
stats_queue=self.stats_queue,
)
@staticmethod
def attach_requests(handles: InferenceClientHandles) -> tuple[SharedMemory, np.ndarray]:
shm = SharedMemory(name=handles.request_shm_name)
array = np.ndarray(
(handles.num_slots, handles.input_dim),
dtype=np.float32,
buffer=shm.buf,
)
return shm, array
@staticmethod
def attach_responses(handles: InferenceClientHandles) -> tuple[SharedMemory, np.ndarray]:
shm = SharedMemory(name=handles.response_shm_name)
array = np.ndarray(
(handles.num_slots, handles.action_size),
dtype=np.float32,
buffer=shm.buf,
)
return shm, array
def close(self) -> None:
self._request_shm.close()
self._response_shm.close()
def unlink(self) -> None:
self._request_shm.unlink()
self._response_shm.unlink()
def release(self) -> None:
self.close()
self.unlink()
@@ -0,0 +1,100 @@
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.inference_buffers import (
InferenceBuffers,
InferenceClientHandles,
)
NETWORK_KIND_ADVANTAGE = "advantage"
NETWORK_KIND_STRATEGY = "strategy"
NETWORK_KIND_LEAGUE = "league"
@dataclass(frozen=True)
class RequestMessage:
slot_id: int
network_kind: str
player: int
network_index: int
class InferenceClient:
def __init__(self, handles: InferenceClientHandles) -> None:
self.handles = handles
self._request_shm, self._requests = InferenceBuffers.attach_requests(handles)
self._response_shm, self._responses = InferenceBuffers.attach_responses(handles)
self._slot_id = int(self.handles.free_slots.get())
self._closed = False
def close(self) -> None:
if not self._closed:
self.handles.ready_events[self._slot_id].clear()
self.handles.free_slots.put(self._slot_id)
self._closed = True
self._request_shm.close()
self._response_shm.close()
def forward(
self,
*,
network_kind: str,
player: int,
network_index: int,
state: np.ndarray,
) -> np.ndarray:
state = np.asarray(state, dtype=np.float32)
if state.shape != (self.handles.input_dim,):
raise ValueError(
f"state must have shape {(self.handles.input_dim,)}, got {state.shape}"
)
if self._closed:
raise RuntimeError("InferenceClient is closed")
slot_id = self._slot_id
ready = self.handles.ready_events[slot_id]
ready.clear()
self._requests[slot_id, :] = state
self.handles.request_queue.put(
RequestMessage(
slot_id=slot_id,
network_kind=network_kind,
player=player,
network_index=network_index,
)
)
ready.wait()
ready.clear()
return self._responses[slot_id, :].copy()
class NetworkProxy:
def __init__(
self,
client: InferenceClient,
*,
network_kind: str,
player: int,
network_index: int,
) -> None:
self.client = client
self.network_kind = network_kind
self.player = player
self.network_index = network_index
def __call__(self, x: torch.Tensor) -> torch.Tensor:
if x.ndim != 2 or x.shape[0] != 1:
raise ValueError(
f"NetworkProxy expects a single-row tensor, got shape {tuple(x.shape)}"
)
state = x.squeeze(0).detach().cpu().numpy().astype(np.float32, copy=False)
result = self.client.forward(
network_kind=self.network_kind,
player=self.player,
network_index=self.network_index,
state=state,
)
return torch.from_numpy(result).unsqueeze(0)
@@ -0,0 +1,324 @@
from __future__ import annotations
import multiprocessing as mp
import queue
import time
from dataclasses import dataclass
from typing import Any
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.config import InferenceServerConfig, NetworkConfig
from coolrl_lost_cities.games.classic.deep_cfr.inference_buffers import (
InferenceBuffers,
InferenceClientHandles,
)
from coolrl_lost_cities.games.classic.deep_cfr.inference_client import (
NETWORK_KIND_ADVANTAGE,
NETWORK_KIND_LEAGUE,
NETWORK_KIND_STRATEGY,
RequestMessage,
)
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
@dataclass(frozen=True)
class ShutdownMessage:
pass
@dataclass(frozen=True)
class WeightUpdateMessage:
advantage_networks: list[dict[str, torch.Tensor]]
strategy_network: dict[str, torch.Tensor] | None
league_advantage_networks: list[list[dict[str, torch.Tensor]]]
@dataclass(frozen=True)
class BatchStatsMessage:
batch_size: int
group_count: int
def _resolve_server_device(device: str) -> torch.device:
token = device.strip().lower()
if token == "auto":
token = "cuda" if torch.cuda.is_available() else "cpu"
return torch.device(token)
def _new_network(
*,
input_dim: int,
action_size: int,
network_config: NetworkConfig,
device: torch.device,
) -> torch.nn.Module:
return DeepCFRMLP.from_config(input_dim, action_size, network_config).to(device).eval()
def _load_state_dict_on_device(
network: torch.nn.Module,
state_dict: dict[str, torch.Tensor],
device: torch.device,
) -> None:
network.load_state_dict({name: value.to(device) for name, value in state_dict.items()})
network.eval()
def run_inference_server(
handles: InferenceClientHandles,
*,
network_config_data: dict[str, Any],
server_config_data: dict[str, Any],
) -> None:
network_config = NetworkConfig.model_validate(network_config_data)
server_config = InferenceServerConfig.model_validate(server_config_data)
device = _resolve_server_device(server_config.device)
request_shm, requests = InferenceBuffers.attach_requests(handles)
response_shm, responses = InferenceBuffers.attach_responses(handles)
advantage_networks = [
_new_network(
input_dim=handles.input_dim,
action_size=handles.action_size,
network_config=network_config,
device=device,
)
for _ in range(2)
]
strategy_network = _new_network(
input_dim=handles.input_dim,
action_size=handles.action_size,
network_config=network_config,
device=device,
)
league_advantage_networks: list[list[torch.nn.Module]] = []
try:
while True:
if _apply_pending_weight_updates(
handles,
advantage_networks,
strategy_network,
league_advantage_networks,
input_dim=handles.input_dim,
action_size=handles.action_size,
network_config=network_config,
device=device,
):
continue
try:
first = handles.request_queue.get(timeout=0.01)
except queue.Empty:
continue
if isinstance(first, ShutdownMessage):
return
if not isinstance(first, RequestMessage):
continue
batch = [first]
deadline = time.perf_counter() + max(0, server_config.batch_window_us) / 1_000_000.0
while len(batch) < server_config.max_batch:
remaining = deadline - time.perf_counter()
if remaining <= 0.0:
break
try:
item = handles.request_queue.get(timeout=remaining)
except queue.Empty:
break
if isinstance(item, ShutdownMessage):
return
if isinstance(item, RequestMessage):
batch.append(item)
_serve_request_batch(
batch,
requests,
responses,
handles,
advantage_networks,
strategy_network,
league_advantage_networks,
device=device,
use_amp=server_config.use_amp,
)
finally:
request_shm.close()
response_shm.close()
def _apply_pending_weight_updates(
handles: InferenceClientHandles,
advantage_networks: list[torch.nn.Module],
strategy_network: torch.nn.Module,
league_advantage_networks: list[list[torch.nn.Module]],
*,
input_dim: int,
action_size: int,
network_config: NetworkConfig,
device: torch.device,
) -> bool:
applied = False
while True:
try:
item = handles.weight_queue.get_nowait()
except queue.Empty:
break
if isinstance(item, ShutdownMessage):
handles.request_queue.put(item)
return True
if not isinstance(item, WeightUpdateMessage):
continue
for network, state_dict in zip(advantage_networks, item.advantage_networks, strict=True):
_load_state_dict_on_device(network, state_dict, device)
if item.strategy_network is not None:
_load_state_dict_on_device(strategy_network, item.strategy_network, device)
league_advantage_networks[:] = []
for snapshot in item.league_advantage_networks:
snapshot_networks = [
_new_network(
input_dim=input_dim,
action_size=action_size,
network_config=network_config,
device=device,
)
for _ in range(2)
]
for network, state_dict in zip(snapshot_networks, snapshot, strict=True):
_load_state_dict_on_device(network, state_dict, device)
league_advantage_networks.append(snapshot_networks)
handles.weight_sync_event.set()
applied = True
return applied
def _network_for_request(
request: RequestMessage,
advantage_networks: list[torch.nn.Module],
strategy_network: torch.nn.Module,
league_advantage_networks: list[list[torch.nn.Module]],
) -> torch.nn.Module:
if request.network_kind == NETWORK_KIND_ADVANTAGE:
return advantage_networks[request.network_index]
if request.network_kind == NETWORK_KIND_STRATEGY:
return strategy_network
if request.network_kind == NETWORK_KIND_LEAGUE:
return league_advantage_networks[request.network_index][request.player]
raise ValueError(f"unknown network kind: {request.network_kind!r}")
def _serve_request_batch(
batch: list[RequestMessage],
requests: np.ndarray,
responses: np.ndarray,
handles: InferenceClientHandles,
advantage_networks: list[torch.nn.Module],
strategy_network: torch.nn.Module,
league_advantage_networks: list[list[torch.nn.Module]],
*,
device: torch.device,
use_amp: bool,
) -> None:
groups: dict[tuple[str, int, int], list[RequestMessage]] = {}
for request in batch:
key = (request.network_kind, request.player, request.network_index)
groups.setdefault(key, []).append(request)
handles.stats_queue.put(BatchStatsMessage(batch_size=len(batch), group_count=len(groups)))
with torch.inference_mode():
for requests_for_network in groups.values():
network = _network_for_request(
requests_for_network[0],
advantage_networks,
strategy_network,
league_advantage_networks,
)
slots = [request.slot_id for request in requests_for_network]
x = torch.as_tensor(requests[slots, :], dtype=torch.float32, device=device)
if use_amp and device.type == "cuda":
with torch.autocast(device_type="cuda"):
output = network(x)
else:
output = network(x)
values = output.detach().to("cpu", dtype=torch.float32).numpy()
for row, slot in enumerate(slots):
responses[slot, :] = values[row]
handles.ready_events[slot].set()
class InferenceServerController:
def __init__(
self,
*,
input_dim: int,
action_size: int,
num_slots: int,
network_config: NetworkConfig,
server_config: InferenceServerConfig,
) -> None:
self._ctx = mp.get_context("spawn")
self._buffers = InferenceBuffers(
num_slots=num_slots,
input_dim=input_dim,
action_size=action_size,
mp_context=self._ctx,
)
self.handles = self._buffers.handles()
self._process = self._ctx.Process(
target=run_inference_server,
kwargs={
"handles": self.handles,
"network_config_data": network_config.model_dump(mode="json"),
"server_config_data": server_config.model_dump(mode="json"),
},
daemon=True,
)
self._process.start()
@property
def is_alive(self) -> bool:
return self._process.is_alive()
def push_weights(
self,
*,
advantage_networks: list[dict[str, torch.Tensor]],
strategy_network: dict[str, torch.Tensor] | None,
league_advantage_networks: list[list[dict[str, torch.Tensor]]],
timeout: float = 60.0,
) -> None:
if not self.is_alive:
raise RuntimeError("inference server process is not alive")
self.handles.weight_sync_event.clear()
self.handles.weight_queue.put(
WeightUpdateMessage(
advantage_networks=advantage_networks,
strategy_network=strategy_network,
league_advantage_networks=league_advantage_networks,
)
)
if not self.handles.weight_sync_event.wait(timeout=timeout):
raise TimeoutError("timed out waiting for inference server weight sync")
def drain_batch_stats(self) -> list[BatchStatsMessage]:
stats: list[BatchStatsMessage] = []
while True:
try:
item = self.handles.stats_queue.get_nowait()
except queue.Empty:
break
if isinstance(item, BatchStatsMessage):
stats.append(item)
return stats
def shutdown(self, timeout: float = 10.0) -> None:
try:
self.handles.request_queue.put(ShutdownMessage())
self.handles.weight_queue.put(ShutdownMessage())
self._process.join(timeout=timeout)
if self._process.is_alive():
self._process.terminate()
self._process.join(timeout=timeout)
finally:
self._buffers.release()
@@ -23,6 +23,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.config import (
) )
from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim
from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy_network from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy_network
from coolrl_lost_cities.games.classic.deep_cfr.inference_server import (
BatchStatsMessage,
InferenceServerController,
)
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample 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.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.deep_cfr.tracking import ( from coolrl_lost_cities.games.classic.deep_cfr.tracking import (
@@ -35,6 +39,7 @@ from coolrl_lost_cities.games.classic.deep_cfr.traversal import run_cython_trave
from coolrl_lost_cities.games.classic.deep_cfr.traversal_stats import TraversalStats from coolrl_lost_cities.games.classic.deep_cfr.traversal_stats import TraversalStats
from coolrl_lost_cities.games.classic.deep_cfr.workers import ( from coolrl_lost_cities.games.classic.deep_cfr.workers import (
TraversalWorkerBatch, TraversalWorkerBatch,
initialize_traversal_worker,
run_traversal_worker_batch, run_traversal_worker_batch,
) )
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
@@ -251,6 +256,8 @@ class DeepCFRTrainer:
self.tracker = CompositeRunTracker(trackers) self.tracker = CompositeRunTracker(trackers)
self.self_play_league_snapshots: list[list[dict]] = [] self.self_play_league_snapshots: list[list[dict]] = []
self._runtime_metrics: dict[str, float | int] = {} self._runtime_metrics: dict[str, float | int] = {}
self._inference_server: InferenceServerController | None = None
self._last_inference_weight_sync_iteration: int | None = None
def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict: def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict:
return { return {
@@ -308,8 +315,14 @@ class DeepCFRTrainer:
def run_iteration(self, iteration: int) -> IterationMetrics: def run_iteration(self, iteration: int) -> IterationMetrics:
self.iteration = iteration self.iteration = iteration
self._runtime_metrics = {} self._runtime_metrics = {}
if self.config.traversal.inference_backend == "server":
self._ensure_inference_server()
self._maybe_sync_inference_server(iteration)
traversal_started = time.perf_counter() traversal_started = time.perf_counter()
if self.config.traversal.resolved_num_workers() > 1: if (
self.config.traversal.resolved_num_workers() > 1
or self.config.traversal.inference_backend == "server"
):
total_stats = self._run_traversals_parallel(iteration) total_stats = self._run_traversals_parallel(iteration)
else: else:
total_stats = self._run_traversals_single_process(iteration) total_stats = self._run_traversals_single_process(iteration)
@@ -460,10 +473,17 @@ class DeepCFRTrainer:
progress_nodes = 0 progress_nodes = 0
progress_traversals = 0 progress_traversals = 0
progress_started = time.perf_counter() progress_started = time.perf_counter()
with ProcessPoolExecutor( executor_kwargs: dict[str, object] = {
max_workers=max_workers, "max_workers": max_workers,
mp_context=mp.get_context("spawn"), "mp_context": mp.get_context("spawn"),
) as executor: }
if self.config.traversal.inference_backend == "server":
if self._inference_server is None:
raise RuntimeError("inference server is not initialized")
self._record_inference_batch_stats(self._inference_server.drain_batch_stats())
executor_kwargs["initializer"] = initialize_traversal_worker
executor_kwargs["initargs"] = (self._inference_server.handles,)
with ProcessPoolExecutor(**executor_kwargs) as executor:
total_batches = len(batches) total_batches = len(batches)
in_flight_limit = min(total_batches, max(1, max_workers * 2)) in_flight_limit = min(total_batches, max(1, max_workers * 2))
batch_iter = iter(batches) batch_iter = iter(batches)
@@ -473,7 +493,15 @@ class DeepCFRTrainer:
} }
completed_batches = 0 completed_batches = 0
while futures: while futures:
done, futures = wait(futures, return_when=FIRST_COMPLETED) done, futures = wait(futures, timeout=5.0, return_when=FIRST_COMPLETED)
if not done:
if (
self.config.traversal.inference_backend == "server"
and self._inference_server is not None
and not self._inference_server.is_alive
):
raise RuntimeError("inference server process exited during traversal")
continue
for future in done: for future in done:
result = future.result() result = future.result()
completed_batches += 1 completed_batches += 1
@@ -505,20 +533,33 @@ class DeepCFRTrainer:
next_batch = next(batch_iter, None) next_batch = next(batch_iter, None)
if next_batch is not None: if next_batch is not None:
futures.add(executor.submit(run_traversal_worker_batch, next_batch)) futures.add(executor.submit(run_traversal_worker_batch, next_batch))
if (
self.config.traversal.inference_backend == "server"
and self._inference_server is not None
):
self._record_inference_batch_stats(self._inference_server.drain_batch_stats())
return total_stats return total_stats
def _worker_batches(self, iteration: int) -> list[TraversalWorkerBatch]: def _worker_batches(self, iteration: int) -> list[TraversalWorkerBatch]:
batches: list[TraversalWorkerBatch] = [] batches: list[TraversalWorkerBatch] = []
network_payloads = [ if self.config.traversal.inference_backend == "server":
{name: value.detach().cpu() for name, value in network.state_dict().items()} network_payloads: list[dict] = []
for network in self.advantage_networks league_payloads = [[{}, {}] for _snapshot in self.self_play_league_snapshots]
] strategy_payload: dict | None = (
strategy_payload: dict | None = None {} if self.config.traversal.opponent_policy == "average_strategy" else None
if self.config.traversal.opponent_policy == "average_strategy": )
strategy_payload = { else:
name: value.detach().cpu() network_payloads = [
for name, value in self.strategy_network.state_dict().items() {name: value.detach().cpu() for name, value in network.state_dict().items()}
} for network in self.advantage_networks
]
league_payloads = self._league_payloads()
strategy_payload = None
if self.config.traversal.opponent_policy == "average_strategy":
strategy_payload = {
name: value.detach().cpu()
for name, value in self.strategy_network.state_dict().items()
}
chunk_size = self.config.traversal.worker_chunk_size chunk_size = self.config.traversal.worker_chunk_size
batch_index = 0 batch_index = 0
for player in range(2): for player in range(2):
@@ -538,9 +579,10 @@ class DeepCFRTrainer:
input_dim=self.input_dim, input_dim=self.input_dim,
action_size=self.action_size, action_size=self.action_size,
advantage_networks=network_payloads, advantage_networks=network_payloads,
league_advantage_networks=self._league_payloads(), league_advantage_networks=league_payloads,
worker_seed=self.config.run.seed + iteration * 1_000_003 + batch_index, worker_seed=self.config.run.seed + iteration * 1_000_003 + batch_index,
strategy_network=strategy_payload, strategy_network=strategy_payload,
inference_handles=None,
) )
) )
batch_index += 1 batch_index += 1
@@ -582,6 +624,90 @@ class DeepCFRTrainer:
def _league_payloads(self) -> list[list[dict]]: def _league_payloads(self) -> list[list[dict]]:
return self.self_play_league_snapshots return self.self_play_league_snapshots
def _record_inference_batch_stats(self, stats: list[BatchStatsMessage]) -> None:
if not stats:
return
batch_sizes = [int(item.batch_size) for item in stats]
group_counts = [int(item.group_count) for item in stats]
request_count = sum(batch_sizes)
batch_count = len(batch_sizes)
self._runtime_metrics["inference_server/batches"] = (
int(self._runtime_metrics.get("inference_server/batches", 0)) + batch_count
)
self._runtime_metrics["inference_server/requests"] = (
int(self._runtime_metrics.get("inference_server/requests", 0)) + request_count
)
self._runtime_metrics["inference_server/groups"] = int(
self._runtime_metrics.get("inference_server/groups", 0)
) + sum(group_counts)
total_batches = int(self._runtime_metrics["inference_server/batches"])
total_requests = int(self._runtime_metrics["inference_server/requests"])
total_groups = int(self._runtime_metrics["inference_server/groups"])
self._runtime_metrics["inference_server/avg_batch_size"] = total_requests / max(
total_batches, 1
)
self._runtime_metrics["inference_server/avg_groups_per_batch"] = total_groups / max(
total_batches, 1
)
self._runtime_metrics["inference_server/max_batch_size"] = max(
int(self._runtime_metrics.get("inference_server/max_batch_size", 0)),
max(batch_sizes),
)
current_min = self._runtime_metrics.get("inference_server/min_batch_size")
self._runtime_metrics["inference_server/min_batch_size"] = (
min(int(current_min), min(batch_sizes)) if current_min is not None else min(batch_sizes)
)
def _inference_num_slots(self) -> int:
configured = self.config.inference_server.num_slots
if configured is not None:
return int(configured)
workers = max(1, self.config.traversal.resolved_num_workers())
return max(64, 4 * workers * int(self.config.traversal.worker_chunk_size))
def _ensure_inference_server(self) -> None:
if self.config.traversal.resolved_num_workers() < 1:
raise ValueError(
"traversal.inference_backend='server' requires traversal.num_workers >= 1"
)
if self._inference_server is not None:
if not self._inference_server.is_alive:
raise RuntimeError("inference server process is not alive")
return
self._inference_server = InferenceServerController(
input_dim=self.input_dim,
action_size=self.action_size,
num_slots=self._inference_num_slots(),
network_config=self.config.network,
server_config=self.config.inference_server,
)
def _maybe_sync_inference_server(self, iteration: int) -> None:
if self._inference_server is None:
return
interval = max(1, int(self.config.inference_server.weight_sync_every))
if self._last_inference_weight_sync_iteration is not None and iteration % interval != 0:
return
advantage_payloads = [
{name: value.detach().cpu() for name, value in network.state_dict().items()}
for network in self.advantage_networks
]
strategy_payload = {
name: value.detach().cpu() for name, value in self.strategy_network.state_dict().items()
}
self._inference_server.push_weights(
advantage_networks=advantage_payloads,
strategy_network=strategy_payload,
league_advantage_networks=self._league_payloads(),
)
self._last_inference_weight_sync_iteration = iteration
def _shutdown_inference_server(self) -> None:
if self._inference_server is None:
return
self._inference_server.shutdown()
self._inference_server = None
def train(self) -> list[IterationMetrics]: def train(self) -> list[IterationMetrics]:
self._start_run_logging() self._start_run_logging()
metrics: list[IterationMetrics] = [] metrics: list[IterationMetrics] = []
@@ -606,6 +732,7 @@ class DeepCFRTrainer:
break break
iteration += 1 iteration += 1
finally: finally:
self._shutdown_inference_server()
self.tracker.close() self.tracker.close()
return metrics return metrics
@@ -7,6 +7,14 @@ from typing import Any
import torch import torch
from coolrl_lost_cities.games.classic.deep_cfr.config import config_from_dict from coolrl_lost_cities.games.classic.deep_cfr.config import config_from_dict
from coolrl_lost_cities.games.classic.deep_cfr.inference_buffers import InferenceClientHandles
from coolrl_lost_cities.games.classic.deep_cfr.inference_client import (
NETWORK_KIND_ADVANTAGE,
NETWORK_KIND_LEAGUE,
NETWORK_KIND_STRATEGY,
InferenceClient,
NetworkProxy,
)
from coolrl_lost_cities.games.classic.deep_cfr.memory import TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.memory import TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.deep_cfr.traversal import run_cython_traversal_batch from coolrl_lost_cities.games.classic.deep_cfr.traversal import run_cython_traversal_batch
@@ -14,6 +22,7 @@ from coolrl_lost_cities.games.classic.deep_cfr.traversal_stats import TraversalS
from coolrl_lost_cities.games.classic.game import LostCitiesConfig from coolrl_lost_cities.games.classic.game import LostCitiesConfig
_TORCH_THREADS_CONFIGURED = False _TORCH_THREADS_CONFIGURED = False
_INFERENCE_HANDLES: InferenceClientHandles | None = None
def _configure_worker_torch_threads() -> None: def _configure_worker_torch_threads() -> None:
@@ -31,6 +40,12 @@ def _configure_worker_torch_threads() -> None:
_TORCH_THREADS_CONFIGURED = True _TORCH_THREADS_CONFIGURED = True
def initialize_traversal_worker(inference_handles: InferenceClientHandles | None = None) -> None:
global _INFERENCE_HANDLES
_INFERENCE_HANDLES = inference_handles
_configure_worker_torch_threads()
@dataclass(frozen=True) @dataclass(frozen=True)
class TraversalWorkerBatch: class TraversalWorkerBatch:
player: int player: int
@@ -44,6 +59,7 @@ class TraversalWorkerBatch:
league_advantage_networks: list[list[dict[str, Any]]] league_advantage_networks: list[list[dict[str, Any]]]
worker_seed: int worker_seed: int
strategy_network: dict[str, Any] | None = None strategy_network: dict[str, Any] | None = None
inference_handles: InferenceClientHandles | None = None
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -60,68 +76,112 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
cfg = config_from_dict(batch.config) cfg = config_from_dict(batch.config)
device = torch.device("cpu") device = torch.device("cpu")
networks = [ client: InferenceClient | None = None
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device) try:
for _ in range(2) if cfg.traversal.inference_backend == "server":
] inference_handles = batch.inference_handles or _INFERENCE_HANDLES
for network, state_dict in zip(networks, batch.advantage_networks, strict=True): if inference_handles is None:
network.load_state_dict(state_dict) raise ValueError("server inference backend requires inference handles")
network.eval() client = InferenceClient(inference_handles)
league_networks: list[list[torch.nn.Module]] = [] networks = [
for snapshot in batch.league_advantage_networks: NetworkProxy(
snapshot_networks = [ client,
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device) network_kind=NETWORK_KIND_ADVANTAGE,
for _ in range(2) player=player,
] network_index=player,
for network, state_dict in zip(snapshot_networks, snapshot, strict=True): )
network.load_state_dict(state_dict) for player in range(2)
network.eval() ]
league_networks.append(snapshot_networks) league_networks = [
strategy_network: torch.nn.Module | None = None [
if batch.strategy_network is not None: NetworkProxy(
strategy_network = DeepCFRMLP.from_config( client,
batch.input_dim, batch.action_size, cfg.network network_kind=NETWORK_KIND_LEAGUE,
).to(device) player=player,
strategy_network.load_state_dict(batch.strategy_network) network_index=snapshot_index,
strategy_network.eval() )
game_config = LostCitiesConfig(**batch.game_config) for player in range(2)
total_stats, advantage_samples, strategy_samples = run_cython_traversal_batch( ]
networks, for snapshot_index, _snapshot in enumerate(batch.league_advantage_networks)
game_config, ]
batch.seeds, strategy_network = (
batch.player, NetworkProxy(
batch.iteration, client,
device=device, network_kind=NETWORK_KIND_STRATEGY,
strategy_network=strategy_network, player=-1,
action_size=batch.action_size, network_index=0,
encoding=cfg.encoding, )
epsilon=cfg.traversal.regret_matching_epsilon, if batch.strategy_network is not None
strategy_sample_interval=cfg.traversal.strategy_sample_interval, else None
store_strategy_on_traverser_nodes=cfg.traversal.store_strategy_on_traverser_nodes, )
store_strategy_on_opponent_nodes=cfg.traversal.store_strategy_on_opponent_nodes, else:
max_depth=cfg.traversal.max_depth, networks = [
max_nodes=cfg.traversal.max_nodes_per_traversal, DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device)
sampling_mode=cfg.traversal.sampling_mode, for _ in range(2)
outcome_sampling_epsilon=cfg.traversal.outcome_sampling_epsilon, ]
outcome_sampling_value_clip=cfg.traversal.outcome_sampling_value_clip, for network, state_dict in zip(networks, batch.advantage_networks, strict=True):
outcome_unsampled_regret=cfg.traversal.outcome_unsampled_regret, network.load_state_dict(state_dict)
cutoff_value_mode=cfg.traversal.cutoff_value_mode, network.eval()
cutoff_rollouts=cfg.traversal.cutoff_rollouts, league_networks: list[list[torch.nn.Module]] = []
cutoff_rollout_policy=cfg.traversal.cutoff_rollout_policy, for snapshot in batch.league_advantage_networks:
cutoff_rollout_max_steps=cfg.traversal.cutoff_rollout_max_steps, snapshot_networks = [
opponent_policy=cfg.traversal.opponent_policy, DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(
all_negative_fallback=cfg.regret_matching.all_negative_fallback, device
league_advantage_networks=league_networks, )
self_play_anchor_probability=cfg.self_play.anchor_probability, for _ in range(2)
self_play_current_weight=cfg.self_play.current_weight, ]
self_play_recent_weight=cfg.self_play.recent_weight, for network, state_dict in zip(snapshot_networks, snapshot, strict=True):
self_play_older_weight=cfg.self_play.older_weight, network.load_state_dict(state_dict)
self_play_anchor_weight=cfg.self_play.anchor_weight, network.eval()
self_play_recent_window=cfg.self_play.recent_window, league_networks.append(snapshot_networks)
endpoint_depth_bucket_width=cfg.traversal.endpoint_depth_bucket_width, strategy_network: torch.nn.Module | None = None
endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max, if batch.strategy_network is not None:
seed=batch.worker_seed, strategy_network = DeepCFRMLP.from_config(
) batch.input_dim, batch.action_size, cfg.network
).to(device)
strategy_network.load_state_dict(batch.strategy_network)
strategy_network.eval()
game_config = LostCitiesConfig(**batch.game_config)
total_stats, advantage_samples, strategy_samples = run_cython_traversal_batch(
networks,
game_config,
batch.seeds,
batch.player,
batch.iteration,
device=device,
strategy_network=strategy_network,
action_size=batch.action_size,
encoding=cfg.encoding,
epsilon=cfg.traversal.regret_matching_epsilon,
strategy_sample_interval=cfg.traversal.strategy_sample_interval,
store_strategy_on_traverser_nodes=cfg.traversal.store_strategy_on_traverser_nodes,
store_strategy_on_opponent_nodes=cfg.traversal.store_strategy_on_opponent_nodes,
max_depth=cfg.traversal.max_depth,
max_nodes=cfg.traversal.max_nodes_per_traversal,
sampling_mode=cfg.traversal.sampling_mode,
outcome_sampling_epsilon=cfg.traversal.outcome_sampling_epsilon,
outcome_sampling_value_clip=cfg.traversal.outcome_sampling_value_clip,
outcome_unsampled_regret=cfg.traversal.outcome_unsampled_regret,
cutoff_value_mode=cfg.traversal.cutoff_value_mode,
cutoff_rollouts=cfg.traversal.cutoff_rollouts,
cutoff_rollout_policy=cfg.traversal.cutoff_rollout_policy,
cutoff_rollout_max_steps=cfg.traversal.cutoff_rollout_max_steps,
opponent_policy=cfg.traversal.opponent_policy,
all_negative_fallback=cfg.regret_matching.all_negative_fallback,
league_advantage_networks=league_networks,
self_play_anchor_probability=cfg.self_play.anchor_probability,
self_play_current_weight=cfg.self_play.current_weight,
self_play_recent_weight=cfg.self_play.recent_weight,
self_play_older_weight=cfg.self_play.older_weight,
self_play_anchor_weight=cfg.self_play.anchor_weight,
self_play_recent_window=cfg.self_play.recent_window,
endpoint_depth_bucket_width=cfg.traversal.endpoint_depth_bucket_width,
endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max,
seed=batch.worker_seed,
)
finally:
if client is not None:
client.close()
return TraversalWorkerResult( return TraversalWorkerResult(
player=batch.player, player=batch.player,
stats=total_stats, stats=total_stats,
@@ -0,0 +1,146 @@
from __future__ import annotations
import multiprocessing as mp
import threading
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig
from coolrl_lost_cities.games.classic.deep_cfr.inference_buffers import (
InferenceBuffers,
InferenceClientHandles,
)
from coolrl_lost_cities.games.classic.deep_cfr.inference_client import (
NETWORK_KIND_ADVANTAGE,
InferenceClient,
RequestMessage,
)
from coolrl_lost_cities.games.classic.deep_cfr.inference_server import (
InferenceServerController,
)
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
def _child_write_request_row(handles: InferenceClientHandles) -> None:
request_shm, requests = InferenceBuffers.attach_requests(handles)
try:
requests[1, :] = np.array([1.0, 2.0, 3.0], dtype=np.float32)
finally:
request_shm.close()
def test_inference_buffers_shared_memory_round_trip() -> None:
ctx = mp.get_context("spawn")
buffers = InferenceBuffers(num_slots=2, input_dim=3, action_size=4, mp_context=ctx)
try:
process = ctx.Process(target=_child_write_request_row, args=(buffers.handles(),))
process.start()
process.join(timeout=10.0)
assert process.exitcode == 0
np.testing.assert_allclose(buffers.requests[1], np.array([1.0, 2.0, 3.0]))
finally:
buffers.release()
def test_inference_client_forwards_through_request_queue() -> None:
buffers = InferenceBuffers(num_slots=2, input_dim=3, action_size=3)
handles = buffers.handles()
client = InferenceClient(handles)
def serve_one() -> None:
request = handles.request_queue.get(timeout=5.0)
assert isinstance(request, RequestMessage)
buffers.responses[request.slot_id, :] = buffers.requests[request.slot_id, :] * 2.0
handles.ready_events[request.slot_id].set()
thread = threading.Thread(target=serve_one)
thread.start()
try:
result = client.forward(
network_kind=NETWORK_KIND_ADVANTAGE,
player=0,
network_index=0,
state=np.array([2.0, 3.0, 4.0], dtype=np.float32),
)
np.testing.assert_allclose(result, np.array([4.0, 6.0, 8.0], dtype=np.float32))
finally:
thread.join(timeout=5.0)
client.close()
buffers.release()
def test_inference_server_matches_network_forward() -> None:
config = DeepCFRConfig.model_validate(
{
"network": {"hidden_size": 8, "num_layers": 1},
"inference_server": {"device": "cpu", "num_slots": 4, "max_batch": 4},
}
)
input_dim = 3
action_size = 2
network = DeepCFRMLP.from_config(input_dim, action_size, config.network)
state_dict = network.state_dict()
controller = InferenceServerController(
input_dim=input_dim,
action_size=action_size,
num_slots=4,
network_config=config.network,
server_config=config.inference_server,
)
client = InferenceClient(controller.handles)
try:
controller.push_weights(
advantage_networks=[state_dict, state_dict],
strategy_network=state_dict,
league_advantage_networks=[],
)
state = np.array([0.25, -0.5, 1.5], dtype=np.float32)
result = client.forward(
network_kind=NETWORK_KIND_ADVANTAGE,
player=0,
network_index=0,
state=state,
)
with torch.inference_mode():
expected = network(torch.from_numpy(state).unsqueeze(0)).squeeze(0).numpy()
np.testing.assert_allclose(result, expected, rtol=1.0e-6, atol=1.0e-6)
finally:
client.close()
controller.shutdown()
def test_deep_cfr_training_smoke_with_cpu_inference_server(tmp_path) -> None:
config = DeepCFRConfig.model_validate(
{
"run": {"max_iterations": 1, "seed": 11, "device": "cpu"},
"network": {"hidden_size": 8, "num_layers": 1},
"traversal": {
"traversals_per_player": 1,
"max_depth": 2,
"max_nodes_per_traversal": 64,
"num_workers": 2,
"worker_chunk_size": 1,
"opponent_policy": "average_strategy",
"inference_backend": "server",
},
"optimization": {
"advantage_batch_size": 2,
"strategy_batch_size": 2,
"advantage_updates_per_iteration": 1,
"strategy_updates_per_iteration": 1,
},
"checkpoint": {"save_every": 0, "save_latest": False},
"evaluation": {"eval_every": 0},
"inference_server": {"device": "cpu", "num_slots": 8, "max_batch": 4},
}
)
trainer = DeepCFRTrainer(config=config, run_dir=tmp_path, device="cpu")
metrics = trainer.train()
assert len(metrics) == 1
assert metrics[0].traversal_nodes > 0
assert metrics[0].advantage_samples > 0