diff --git a/configs/deep_cfr/default.yaml b/configs/deep_cfr/default.yaml index 0a9ea83..246176e 100644 --- a/configs/deep_cfr/default.yaml +++ b/configs/deep_cfr/default.yaml @@ -46,6 +46,7 @@ traversal: progress_every_traversals: 10 endpoint_depth_bucket_width: 100 endpoint_depth_bucket_max: 1000 + inference_backend: local regret_matching: all_negative_fallback: argmax_tiebreak @@ -97,3 +98,11 @@ evaluation: 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 diff --git a/configs/deep_cfr/default_server.yaml b/configs/deep_cfr/default_server.yaml new file mode 100644 index 0000000..eed2349 --- /dev/null +++ b/configs/deep_cfr/default_server.yaml @@ -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 diff --git a/scripts/bench_inference_backend.py b/scripts/bench_inference_backend.py new file mode 100644 index 0000000..ef7efd3 --- /dev/null +++ b/scripts/bench_inference_backend.py @@ -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() diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py index e73401c..d141c3a 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -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") diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/inference_buffers.py b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_buffers.py new file mode 100644 index 0000000..0909ad4 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_buffers.py @@ -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() diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/inference_client.py b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_client.py new file mode 100644 index 0000000..b4d6773 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_client.py @@ -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) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/inference_server.py b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_server.py new file mode 100644 index 0000000..3100b3a --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/inference_server.py @@ -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() diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py index 40a0a66..9274422 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -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 diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py index 634b1b0..4eac6db 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -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, diff --git a/tests/games/classic/test_inference_server.py b/tests/games/classic/test_inference_server.py new file mode 100644 index 0000000..c32f92f --- /dev/null +++ b/tests/games/classic/test_inference_server.py @@ -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