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:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user