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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,277 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import json
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.cli import _with_overrides
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load_config
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.tracking import FileRunTracker
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
|
||||
|
||||
METRIC_KEYS = (
|
||||
"time/iteration_seconds",
|
||||
"time/traversal_seconds",
|
||||
"time/advantage_train_seconds",
|
||||
"time/strategy_train_seconds",
|
||||
"time/memory_add_seconds",
|
||||
"time/batch_tensor_seconds",
|
||||
)
|
||||
|
||||
TABLE_COLUMNS = (
|
||||
("iter", "time/iteration_seconds"),
|
||||
("traversal", "time/traversal_seconds"),
|
||||
("adv_train", "time/advantage_train_seconds"),
|
||||
("strat_train", "time/strategy_train_seconds"),
|
||||
("mem_add", "time/memory_add_seconds"),
|
||||
("batch_tensor", "time/batch_tensor_seconds"),
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Benchmark Deep CFR traversal inference backends.")
|
||||
parser.add_argument(
|
||||
"--config-local",
|
||||
default="configs/deep_cfr/default.yaml",
|
||||
help="Config for the local traversal inference backend.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config-server",
|
||||
default="configs/deep_cfr/default_server.yaml",
|
||||
help="Config for the server traversal inference backend.",
|
||||
)
|
||||
parser.add_argument("--iterations", type=int, default=10)
|
||||
parser.add_argument("--warmup", type=int, default=2)
|
||||
parser.add_argument("--device", help="Override run.device for both backend configs.")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def _benchmark_overrides(
|
||||
*,
|
||||
backend: str,
|
||||
iterations: int,
|
||||
seed: int,
|
||||
device: str | None,
|
||||
) -> dict[str, Any]:
|
||||
overrides: dict[str, Any] = {
|
||||
"run": {
|
||||
"max_iterations": iterations,
|
||||
"max_minutes": None,
|
||||
"seed": seed,
|
||||
},
|
||||
"traversal": {"inference_backend": backend},
|
||||
"checkpoint": {
|
||||
"save_latest": False,
|
||||
"save_every": 0,
|
||||
},
|
||||
"evaluation": {"eval_every": 0},
|
||||
}
|
||||
if device is not None:
|
||||
overrides["run"]["device"] = device
|
||||
return overrides
|
||||
|
||||
|
||||
def _load_benchmark_config(
|
||||
path: str,
|
||||
*,
|
||||
backend: str,
|
||||
iterations: int,
|
||||
seed: int,
|
||||
device: str | None,
|
||||
) -> DeepCFRConfig:
|
||||
config = load_config(path)
|
||||
return _with_overrides(
|
||||
config,
|
||||
_benchmark_overrides(
|
||||
backend=backend,
|
||||
iterations=iterations,
|
||||
seed=seed,
|
||||
device=device,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _read_metrics(path: Path) -> list[dict[str, float | int]]:
|
||||
rows: list[dict[str, float | int]] = []
|
||||
with path.open(encoding="utf-8") as handle:
|
||||
for line in handle:
|
||||
line = line.strip()
|
||||
if line:
|
||||
rows.append(json.loads(line))
|
||||
return rows
|
||||
|
||||
|
||||
def _mean(rows: list[dict[str, float | int]], key: str) -> float:
|
||||
if not rows:
|
||||
return 0.0
|
||||
return sum(float(row.get(key, 0.0) or 0.0) for row in rows) / len(rows)
|
||||
|
||||
|
||||
def _summarize(rows: list[dict[str, float | int]], *, warmup: int) -> dict[str, float]:
|
||||
measured = rows[warmup:]
|
||||
return {key: _mean(measured, key) for key in METRIC_KEYS}
|
||||
|
||||
|
||||
def _clear_gpu_memory() -> None:
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def _run_backend(
|
||||
*,
|
||||
name: str,
|
||||
config: DeepCFRConfig,
|
||||
run_dir: Path,
|
||||
warmup: int,
|
||||
) -> dict[str, Any]:
|
||||
_clear_gpu_memory()
|
||||
tracker = FileRunTracker(
|
||||
log_path=run_dir / "train.log",
|
||||
metrics_path=run_dir / "metrics.jsonl",
|
||||
progress_path=run_dir / "runtime_progress.json",
|
||||
)
|
||||
trainer: DeepCFRTrainer | None = DeepCFRTrainer(
|
||||
config,
|
||||
config.rules.to_lost_cities_config(seed=config.run.seed),
|
||||
run_dir=run_dir,
|
||||
device=config.run.device,
|
||||
tracker=tracker,
|
||||
)
|
||||
try:
|
||||
trainer.train()
|
||||
rows = _read_metrics(run_dir / "metrics.jsonl")
|
||||
finally:
|
||||
trainer = None
|
||||
_clear_gpu_memory()
|
||||
return {
|
||||
"backend": name,
|
||||
"run_dir": str(run_dir),
|
||||
"raw": rows,
|
||||
"mean": _summarize(rows, warmup=warmup),
|
||||
}
|
||||
|
||||
|
||||
def _speedups(local: dict[str, float], server: dict[str, float]) -> dict[str, float | None]:
|
||||
result: dict[str, float | None] = {}
|
||||
for key in METRIC_KEYS:
|
||||
server_value = server.get(key, 0.0)
|
||||
result[key] = None if server_value <= 0.0 else local.get(key, 0.0) / server_value
|
||||
return result
|
||||
|
||||
|
||||
def _format_seconds(value: float) -> str:
|
||||
return f"{value:.2f}s"
|
||||
|
||||
|
||||
def _format_multiplier(value: float | None) -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
return f"{value:.2f}x"
|
||||
|
||||
|
||||
def _print_backend_table(results: dict[str, Any]) -> None:
|
||||
local = results["backends"]["local"]["mean"]
|
||||
server = results["backends"]["server"]["mean"]
|
||||
speedup = results["speedup"]
|
||||
rows = [
|
||||
("local", [_format_seconds(local[key]) for _label, key in TABLE_COLUMNS]),
|
||||
("server", [_format_seconds(server[key]) for _label, key in TABLE_COLUMNS]),
|
||||
("speedup", [_format_multiplier(speedup[key]) for _label, key in TABLE_COLUMNS]),
|
||||
]
|
||||
headers = ["Backend", *[label for label, _key in TABLE_COLUMNS]]
|
||||
widths = [max(len(headers[0]), *(len(row[0]) for row in rows))]
|
||||
for index, header in enumerate(headers[1:]):
|
||||
widths.append(max(len(header), *(len(row[1][index]) for row in rows)))
|
||||
|
||||
print(" ".join(header.ljust(widths[index]) for index, header in enumerate(headers)))
|
||||
for backend, values in rows:
|
||||
cells = [backend.ljust(widths[0])]
|
||||
cells.extend(value.rjust(widths[index + 1]) for index, value in enumerate(values))
|
||||
print(" ".join(cells))
|
||||
|
||||
|
||||
def _print_projection(results: dict[str, Any]) -> None:
|
||||
print()
|
||||
print("1000-iteration projection")
|
||||
print("Backend hours")
|
||||
for backend in ("local", "server"):
|
||||
mean_iter = results["backends"][backend]["mean"]["time/iteration_seconds"]
|
||||
hours = mean_iter * 1000.0 / 3600.0
|
||||
print(f"{backend:<7} {hours:.2f}h")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.iterations <= 0:
|
||||
raise SystemExit("--iterations must be positive")
|
||||
if args.warmup < 0:
|
||||
raise SystemExit("--warmup must be non-negative")
|
||||
if args.warmup >= args.iterations:
|
||||
raise SystemExit("--warmup must be less than --iterations")
|
||||
|
||||
local_base = load_config(args.config_local)
|
||||
seed = local_base.run.seed
|
||||
timestamp = datetime.now().strftime("%Y-%m-%d_%H%M%S")
|
||||
bench_dir = Path("runs/bench") / f"{timestamp}_inference_backend"
|
||||
bench_dir.mkdir(parents=True, exist_ok=False)
|
||||
|
||||
local_config = _load_benchmark_config(
|
||||
args.config_local,
|
||||
backend="local",
|
||||
iterations=args.iterations,
|
||||
seed=seed,
|
||||
device=args.device,
|
||||
)
|
||||
server_config = _load_benchmark_config(
|
||||
args.config_server,
|
||||
backend="server",
|
||||
iterations=args.iterations,
|
||||
seed=seed,
|
||||
device=args.device,
|
||||
)
|
||||
|
||||
local_result = _run_backend(
|
||||
name="local",
|
||||
config=local_config,
|
||||
run_dir=bench_dir / "local",
|
||||
warmup=args.warmup,
|
||||
)
|
||||
server_result = _run_backend(
|
||||
name="server",
|
||||
config=server_config,
|
||||
run_dir=bench_dir / "server",
|
||||
warmup=args.warmup,
|
||||
)
|
||||
|
||||
results = {
|
||||
"config": {
|
||||
"config_local": args.config_local,
|
||||
"config_server": args.config_server,
|
||||
"iterations": args.iterations,
|
||||
"warmup": args.warmup,
|
||||
"seed": seed,
|
||||
"device_override": args.device,
|
||||
},
|
||||
"backends": {
|
||||
"local": local_result,
|
||||
"server": server_result,
|
||||
},
|
||||
"speedup": _speedups(local_result["mean"], server_result["mean"]),
|
||||
}
|
||||
results_path = bench_dir / "results.json"
|
||||
results_path.write_text(json.dumps(results, indent=2, sort_keys=True), encoding="utf-8")
|
||||
|
||||
_print_backend_table(results)
|
||||
_print_projection(results)
|
||||
print()
|
||||
print(f"Results written to: {results_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -111,6 +111,7 @@ class TraversalConfig(StrictModel):
|
||||
progress_every_traversals: int = 0
|
||||
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,15 +533,28 @@ 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] = []
|
||||
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
|
||||
]
|
||||
strategy_payload: dict | None = None
|
||||
league_payloads = self._league_payloads()
|
||||
strategy_payload = None
|
||||
if self.config.traversal.opponent_policy == "average_strategy":
|
||||
strategy_payload = {
|
||||
name: value.detach().cpu()
|
||||
@@ -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,6 +76,45 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
||||
|
||||
cfg = config_from_dict(batch.config)
|
||||
device = torch.device("cpu")
|
||||
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)
|
||||
@@ -70,7 +125,9 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
||||
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)
|
||||
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):
|
||||
@@ -122,6 +179,9 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
||||
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,
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user