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
@@ -111,6 +111,7 @@ class TraversalConfig(StrictModel):
progress_every_traversals: int = 0
endpoint_depth_bucket_width: int = 100
endpoint_depth_bucket_max: int = 1000
inference_backend: str = "local"
@field_validator("sampling_mode")
@classmethod
@@ -150,6 +151,14 @@ class TraversalConfig(StrictModel):
)
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:
if isinstance(self.num_workers, str):
token = self.num_workers.strip().lower()
@@ -256,6 +265,23 @@ class EvaluationConfig(StrictModel):
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):
run: RunConfig = Field(default_factory=RunConfig)
rules: RulesConfig = Field(default_factory=RulesConfig)
@@ -269,6 +295,7 @@ class DeepCFRConfig(StrictModel):
memory: MemoryConfig = Field(default_factory=MemoryConfig)
checkpoint: CheckpointConfig = Field(default_factory=CheckpointConfig)
evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig)
inference_server: InferenceServerConfig = Field(default_factory=InferenceServerConfig)
def to_dict(self) -> dict[str, Any]:
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.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.networks import DeepCFRMLP
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.workers import (
TraversalWorkerBatch,
initialize_traversal_worker,
run_traversal_worker_batch,
)
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
@@ -251,6 +256,8 @@ class DeepCFRTrainer:
self.tracker = CompositeRunTracker(trackers)
self.self_play_league_snapshots: list[list[dict]] = []
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:
return {
@@ -308,8 +315,14 @@ class DeepCFRTrainer:
def run_iteration(self, iteration: int) -> IterationMetrics:
self.iteration = iteration
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()
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)
else:
total_stats = self._run_traversals_single_process(iteration)
@@ -460,10 +473,17 @@ class DeepCFRTrainer:
progress_nodes = 0
progress_traversals = 0
progress_started = time.perf_counter()
with ProcessPoolExecutor(
max_workers=max_workers,
mp_context=mp.get_context("spawn"),
) as executor:
executor_kwargs: dict[str, object] = {
"max_workers": max_workers,
"mp_context": mp.get_context("spawn"),
}
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)
in_flight_limit = min(total_batches, max(1, max_workers * 2))
batch_iter = iter(batches)
@@ -473,7 +493,15 @@ class DeepCFRTrainer:
}
completed_batches = 0
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:
result = future.result()
completed_batches += 1
@@ -505,20 +533,33 @@ class DeepCFRTrainer:
next_batch = next(batch_iter, None)
if next_batch is not None:
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
def _worker_batches(self, iteration: int) -> list[TraversalWorkerBatch]:
batches: list[TraversalWorkerBatch] = []
network_payloads = [
{name: value.detach().cpu() for name, value in network.state_dict().items()}
for network in self.advantage_networks
]
strategy_payload: dict | None = 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()
}
if self.config.traversal.inference_backend == "server":
network_payloads: list[dict] = []
league_payloads = [[{}, {}] for _snapshot in self.self_play_league_snapshots]
strategy_payload: dict | None = (
{} if self.config.traversal.opponent_policy == "average_strategy" else None
)
else:
network_payloads = [
{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
batch_index = 0
for player in range(2):
@@ -538,9 +579,10 @@ class DeepCFRTrainer:
input_dim=self.input_dim,
action_size=self.action_size,
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,
strategy_network=strategy_payload,
inference_handles=None,
)
)
batch_index += 1
@@ -582,6 +624,90 @@ class DeepCFRTrainer:
def _league_payloads(self) -> list[list[dict]]:
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]:
self._start_run_logging()
metrics: list[IterationMetrics] = []
@@ -606,6 +732,7 @@ class DeepCFRTrainer:
break
iteration += 1
finally:
self._shutdown_inference_server()
self.tracker.close()
return metrics
@@ -7,6 +7,14 @@ from typing import Any
import torch
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.networks import DeepCFRMLP
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
_TORCH_THREADS_CONFIGURED = False
_INFERENCE_HANDLES: InferenceClientHandles | None = None
def _configure_worker_torch_threads() -> None:
@@ -31,6 +40,12 @@ def _configure_worker_torch_threads() -> None:
_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)
class TraversalWorkerBatch:
player: int
@@ -44,6 +59,7 @@ class TraversalWorkerBatch:
league_advantage_networks: list[list[dict[str, Any]]]
worker_seed: int
strategy_network: dict[str, Any] | None = None
inference_handles: InferenceClientHandles | None = None
@dataclass(frozen=True)
@@ -60,68 +76,112 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
cfg = config_from_dict(batch.config)
device = torch.device("cpu")
networks = [
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device)
for _ in range(2)
]
for network, state_dict in zip(networks, batch.advantage_networks, strict=True):
network.load_state_dict(state_dict)
network.eval()
league_networks: list[list[torch.nn.Module]] = []
for snapshot in batch.league_advantage_networks:
snapshot_networks = [
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device)
for _ in range(2)
]
for network, state_dict in zip(snapshot_networks, snapshot, strict=True):
network.load_state_dict(state_dict)
network.eval()
league_networks.append(snapshot_networks)
strategy_network: torch.nn.Module | None = None
if batch.strategy_network is not None:
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,
)
client: InferenceClient | None = None
try:
if cfg.traversal.inference_backend == "server":
inference_handles = batch.inference_handles or _INFERENCE_HANDLES
if inference_handles is None:
raise ValueError("server inference backend requires inference handles")
client = InferenceClient(inference_handles)
networks = [
NetworkProxy(
client,
network_kind=NETWORK_KIND_ADVANTAGE,
player=player,
network_index=player,
)
for player in range(2)
]
league_networks = [
[
NetworkProxy(
client,
network_kind=NETWORK_KIND_LEAGUE,
player=player,
network_index=snapshot_index,
)
for player in range(2)
]
for snapshot_index, _snapshot in enumerate(batch.league_advantage_networks)
]
strategy_network = (
NetworkProxy(
client,
network_kind=NETWORK_KIND_STRATEGY,
player=-1,
network_index=0,
)
if batch.strategy_network is not None
else None
)
else:
networks = [
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(device)
for _ in range(2)
]
for network, state_dict in zip(networks, batch.advantage_networks, strict=True):
network.load_state_dict(state_dict)
network.eval()
league_networks: list[list[torch.nn.Module]] = []
for snapshot in batch.league_advantage_networks:
snapshot_networks = [
DeepCFRMLP.from_config(batch.input_dim, batch.action_size, cfg.network).to(
device
)
for _ in range(2)
]
for network, state_dict in zip(snapshot_networks, snapshot, strict=True):
network.load_state_dict(state_dict)
network.eval()
league_networks.append(snapshot_networks)
strategy_network: torch.nn.Module | None = None
if batch.strategy_network is not None:
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(
player=batch.player,
stats=total_stats,