Optimize get load calls (/v1/loads) using shared-memory load snapshots (#26348)

Co-authored-by: cctry <cctry@meta.com>
This commit is contained in:
Lianmin Zheng
2026-05-29 13:40:26 -07:00
committed by GitHub
co-authored by cctry
parent 3cecc77ccb
commit 4ff1296f5e
14 changed files with 1521 additions and 271 deletions
+26 -90
View File
@@ -18,7 +18,6 @@ This module provides the /v1/loads endpoint which returns detailed scheduler
metrics for load balancing, monitoring, and capacity planning. metrics for load balancing, monitoring, and capacity planning.
""" """
import dataclasses
import time import time
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Optional from typing import Optional
@@ -26,26 +25,10 @@ from typing import Optional
from fastapi import APIRouter, Depends, HTTPException from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import Response from fastapi.responses import Response
from sglang.srt.managers.io_struct import (
DisaggregationMetrics,
GetLoadsReqOutput,
LoRAMetrics,
MemoryMetrics,
QueueMetrics,
SpeculativeMetrics,
)
from sglang.version import __version__ from sglang.version import __version__
router = APIRouter() router = APIRouter()
_OPTIONAL_METRIC_SECTIONS = {
"memory": ("memory", MemoryMetrics),
"speculative": ("spec", SpeculativeMetrics),
"lora": ("lora", LoRAMetrics),
"disaggregation": ("disagg", DisaggregationMetrics),
"queues": ("queues", QueueMetrics),
}
def _get_tokenizer_manager(): def _get_tokenizer_manager():
"""Dependency to get tokenizer_manager from global state.""" """Dependency to get tokenizer_manager from global state."""
@@ -54,77 +37,31 @@ def _get_tokenizer_manager():
return get_global_state().tokenizer_manager return get_global_state().tokenizer_manager
def _loads_dict_factory(items): def _format_loads_prometheus(load_results, include=None) -> Response:
"""Factory for dataclasses.asdict() that excludes None values and timestamp.""" """Format load metrics in Prometheus text exposition format."""
return {k: v for k, v in items if v is not None and k != "timestamp"} section_prefixes = {"speculative": "spec", "disaggregation": "disagg"}
metric_samples = {}
for load in load_results:
load_dict = load.to_dict(include)
dp_rank = load_dict.pop("dp_rank")
def _compute_aggregate(load_dicts: list) -> dict: for key, value in load_dict.items():
"""Compute aggregate metrics from load dicts.""" if isinstance(value, dict):
if not load_dicts: prefix = section_prefixes.get(key, key)
return { for sub_key, sub_value in value.items():
"total_running_reqs": 0, if isinstance(sub_value, (int, float)):
"total_waiting_reqs": 0, metric_samples.setdefault(
"total_reqs": 0, f"sglang_{prefix}_{sub_key}", []
"total_used_tokens": 0, ).append((dp_rank, sub_value))
"total_tokens": 0, elif isinstance(value, (int, float)):
"avg_token_usage": 0.0, metric_samples.setdefault(f"sglang_{key}", []).append((dp_rank, value))
"avg_throughput": 0.0,
"avg_utilization": 0.0,
}
n = len(load_dicts)
return {
"total_running_reqs": sum(d["num_running_reqs"] for d in load_dicts),
"total_waiting_reqs": sum(d["num_waiting_reqs"] for d in load_dicts),
"total_reqs": sum(
d["num_running_reqs"] + d["num_waiting_reqs"] for d in load_dicts
),
"total_used_tokens": sum(d["num_used_tokens"] for d in load_dicts),
"total_tokens": sum(d["num_total_tokens"] for d in load_dicts),
"avg_token_usage": round(sum(d["token_usage"] for d in load_dicts) / n, 4),
"avg_throughput": round(sum(d["gen_throughput"] for d in load_dicts) / n, 2),
"avg_utilization": round(sum(d["utilization"] for d in load_dicts) / n, 4),
}
def _format_loads_prometheus(load_results) -> Response:
"""Format load metrics in Prometheus text exposition format.
Metrics are derived from dataclass field metadata, providing a single source of truth.
"""
lines = [] lines = []
for metric_name, samples in metric_samples.items():
for f in dataclasses.fields(GetLoadsReqOutput): lines.append(f"# TYPE {metric_name} gauge")
if "metric" not in f.metadata: for dp_rank, value in samples:
continue lines.append(f'{metric_name}{{dp_rank="{dp_rank}"}} {value}')
metric_type, description = f.metadata["metric"]
metric_name = f"sglang_{f.name}"
lines.append(f"# HELP {metric_name} {description}")
lines.append(f"# TYPE {metric_name} {metric_type}")
for load in load_results:
value = getattr(load, f.name, None)
if value is not None:
lines.append(f'{metric_name}{{dp_rank="{load.dp_rank}"}} {value}')
for attr_name, (prefix, dataclass_type) in _OPTIONAL_METRIC_SECTIONS.items():
if not any(getattr(load, attr_name, None) for load in load_results):
continue
for f in dataclasses.fields(dataclass_type):
if "metric" not in f.metadata:
continue
metric_type, description = f.metadata["metric"]
metric_name = f"sglang_{prefix}_{f.name}"
lines.append(f"# HELP {metric_name} {description}")
lines.append(f"# TYPE {metric_name} {metric_type}")
for load in load_results:
section = getattr(load, attr_name, None)
if section:
value = getattr(section, f.name, None)
if value is not None:
lines.append(
f'{metric_name}{{dp_rank="{load.dp_rank}"}} {value}'
)
return Response( return Response(
content="\n".join(lines) + "\n", content="\n".join(lines) + "\n",
@@ -150,7 +87,7 @@ async def get_loads(
format: Response format - 'json' (default) or 'prometheus' format: Response format - 'json' (default) or 'prometheus'
Returns: Returns:
JSON response with timestamp, version, dp_rank_count, per-DP-rank loads, and aggregates JSON response with timestamp, version, and per-DP-rank loads
""" """
include_list = [s.strip() for s in include.split(",")] if include else None include_list = [s.strip() for s in include.split(",")] if include else None
@@ -169,19 +106,18 @@ async def get_loads(
time.perf_counter() - start time.perf_counter() - start
) )
include_set = set(include_list) if include_list else None
if format == "prometheus": if format == "prometheus":
return _format_loads_prometheus(load_results) return _format_loads_prometheus(load_results, include_set)
loads = [] loads = []
for load in load_results: for load in load_results:
d = dataclasses.asdict(load, dict_factory=_loads_dict_factory) d = load.to_dict(include_set)
d["num_total_reqs"] = d["num_running_reqs"] + d["num_waiting_reqs"]
loads.append(d) loads.append(d)
return { return {
"timestamp": datetime.now(timezone.utc).isoformat(), "timestamp": datetime.now(timezone.utc).isoformat(),
"version": __version__, "version": __version__,
"dp_rank_count": len(loads),
"loads": loads, "loads": loads,
"aggregate": _compute_aggregate(loads),
} }
+4 -1
View File
@@ -254,6 +254,9 @@ class Envs:
SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY = EnvInt(0) SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY = EnvInt(0)
SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE = EnvBool(True) SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE = EnvBool(True)
# Load snapshot backend
SGLANG_LOAD_SNAPSHOT_USE_ZMQ = EnvBool(False)
# Scheduler: new token ratio hyperparameters # Scheduler: new token ratio hyperparameters
SGLANG_INIT_NEW_TOKEN_RATIO = EnvFloat(0.7) SGLANG_INIT_NEW_TOKEN_RATIO = EnvFloat(0.7)
SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14) SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14)
@@ -819,7 +822,7 @@ _warn_deprecated_env_to_cli_flag(
# Import cuda_coredump to trigger auto-injection of CUDA env vars # Import cuda_coredump to trigger auto-injection of CUDA env vars
# when SGLANG_CUDA_COREDUMP=1. Best-effort; for strict guarantees, # when SGLANG_CUDA_COREDUMP=1. Best-effort; for strict guarantees,
# set CUDA_* env vars in the shell before launching Python. # set CUDA_* env vars in the shell before launching Python.
import sglang.srt.debug_utils.cuda_coredump # noqa: F401, E402 import sglang.srt.debug_utils.cuda_coredump # noqa: F401, E402 # isort: skip
def example_with_exit_stack(): def example_with_exit_stack():
@@ -36,8 +36,8 @@ from sglang.srt.managers.io_struct import (
ProfileReq, ProfileReq,
TokenizedEmbeddingReqInput, TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
WatchLoadUpdateReq,
) )
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import run_scheduler_process from sglang.srt.managers.scheduler import run_scheduler_process
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
@@ -91,10 +91,14 @@ class DPBudget:
self.dp_size = dp_size self.dp_size = dp_size
self.total_requests = [0] * dp_size self.total_requests = [0] * dp_size
self.total_tokens = [0] * dp_size self.total_tokens = [0] * dp_size
self.last_timestamp = [0.0] * dp_size
def update_budget(self, load_update: WatchLoadUpdateReq): def update_budget(self, loads):
"""Update the budget.""" """Update budget from shm snapshots, skipping stale reads."""
for load in load_update.loads: for load in loads:
if load.timestamp == self.last_timestamp[load.dp_rank]:
continue
self.last_timestamp[load.dp_rank] = load.timestamp
self.total_requests[load.dp_rank] = ( self.total_requests[load.dp_rank] = (
load.num_running_reqs + load.num_waiting_reqs load.num_running_reqs + load.num_waiting_reqs
) )
@@ -151,9 +155,19 @@ class DataParallelController:
LoadBalanceMethod.TOTAL_TOKENS: self.total_tokens_scheduler, LoadBalanceMethod.TOTAL_TOKENS: self.total_tokens_scheduler,
} }
self.dispatching = dispatch_lookup[self.load_balance_method] self.dispatching = dispatch_lookup[self.load_balance_method]
self.refresh_load_budget_on_dispatch = self.load_balance_method in (
LoadBalanceMethod.TOTAL_REQUESTS,
LoadBalanceMethod.TOTAL_TOKENS,
)
# Load balance budget # Load balance budget
self.dp_budget = DPBudget(server_args.dp_size) self.dp_budget = DPBudget(server_args.dp_size)
self.load_snapshot_reader = create_load_snapshot_reader(
server_args,
port_args,
caller="dp_controller",
)
self._last_refresh_time = 0.0
# To protect changing env vars to set CUDA_VISIBLE_DEVICES. # To protect changing env vars to set CUDA_VISIBLE_DEVICES.
self.env_lock = threading.Lock() self.env_lock = threading.Lock()
@@ -198,13 +212,30 @@ class DataParallelController:
for worker in self.workers[:: self.control_message_step]: for worker in self.workers[:: self.control_message_step]:
worker.send_pyobj(obj) worker.send_pyobj(obj)
def handle_load_update_req(self, obj):
self.dp_budget.update_budget(obj)
def update_active_ranks(self, ranks: ActiveRanksOutput): def update_active_ranks(self, ranks: ActiveRanksOutput):
self.status = ranks.status self.status = ranks.status
def dispatching_with_trace(self, req: Req): def refresh_load_budget(self):
# Throttle to at most once per 20ms. When a burst of requests
# arrives, dispatching_with_trace() calls this before every
# dispatch. Each call reads the latest scheduler snapshot and
# overwrites the speculative +1 increments that DPBudget.dispatch()
# added for previously dispatched requests in this burst. Without
# throttling, the budget resets to the (stale) scheduler-reported
# value on every request, causing the entire burst to land on a
# single DP rank. The 20ms interval lets the burst complete
# using speculative counters, then refreshes from the real
# scheduler load for the next batch.
now = time.perf_counter()
if now - self._last_refresh_time < 0.02:
return
self._last_refresh_time = now
self.dp_budget.update_budget(self.load_snapshot_reader.read_all())
def dispatching_with_trace(self, req: Req, refresh_load_budget: bool = True):
if refresh_load_budget and self.refresh_load_budget_on_dispatch:
self.refresh_load_budget()
req.time_stats = DPControllerReqTimeStats.new_from_obj(req.time_stats) req.time_stats = DPControllerReqTimeStats.new_from_obj(req.time_stats)
req.time_stats.set_dp_dispatch_time() req.time_stats.set_dp_dispatch_time()
@@ -212,12 +243,16 @@ class DataParallelController:
req.time_stats.set_dp_dispatch_finish_time() req.time_stats.set_dp_dispatch_finish_time()
def dispatch_batch_generate(self, batch_req: BatchTokenizedGenerateReqInput): def dispatch_batch_generate(self, batch_req: BatchTokenizedGenerateReqInput):
if self.refresh_load_budget_on_dispatch:
self.refresh_load_budget()
for req in batch_req: for req in batch_req:
self.dispatching_with_trace(req) self.dispatching_with_trace(req, refresh_load_budget=False)
def dispatch_batch_embedding(self, batch_req: BatchTokenizedEmbeddingReqInput): def dispatch_batch_embedding(self, batch_req: BatchTokenizedEmbeddingReqInput):
if self.refresh_load_budget_on_dispatch:
self.refresh_load_budget()
for req in batch_req: for req in batch_req:
self.dispatching_with_trace(req) self.dispatching_with_trace(req, refresh_load_budget=False)
def init_dispatcher(self): def init_dispatcher(self):
self._request_dispatcher = TypeBasedDispatcher( self._request_dispatcher = TypeBasedDispatcher(
@@ -228,7 +263,6 @@ class DataParallelController:
(BatchTokenizedEmbeddingReqInput, self.dispatch_batch_embedding), (BatchTokenizedEmbeddingReqInput, self.dispatch_batch_embedding),
(BlockReqInput, self.send_to_all_workers), (BlockReqInput, self.send_to_all_workers),
(ProfileReq, self.send_to_all_workers), (ProfileReq, self.send_to_all_workers),
(WatchLoadUpdateReq, self.handle_load_update_req),
(ActiveRanksOutput, self.update_active_ranks), (ActiveRanksOutput, self.update_active_ranks),
] ]
) )
@@ -244,6 +278,7 @@ class DataParallelController:
tmp_port_args = PortArgs.init_new(server_args) tmp_port_args = PortArgs.init_new(server_args)
tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name
tmp_port_args.instance_id = port_args.instance_id
# This port is checked free in PortArgs.init_new. # This port is checked free in PortArgs.init_new.
# We hold it first so that the next dp worker gets a different port # We hold it first so that the next dp worker gets a different port
@@ -494,6 +529,7 @@ class DataParallelController:
# Data parallelism reuses the tensor parallelism group, # Data parallelism reuses the tensor parallelism group,
# so all dp ranks should use the same nccl port. # so all dp ranks should use the same nccl port.
rank_port_args.nccl_port = port_args.nccl_port rank_port_args.nccl_port = port_args.nccl_port
rank_port_args.instance_id = port_args.instance_id
reader, writer = mp.Pipe(duplex=False) reader, writer = mp.Pipe(duplex=False)
gpu_id = ( gpu_id = (
@@ -37,11 +37,7 @@ from sglang.srt.managers.io_struct import (
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import ( from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
configure_logger,
freeze_gc,
kill_itself_when_parent_died,
)
from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.hf_transformers_utils import get_tokenizer
from sglang.srt.utils.network import get_zmq_socket from sglang.srt.utils.network import get_zmq_socket
from sglang.srt.utils.patch_tokenizer import decode_without_hf_kwargs from sglang.srt.utils.patch_tokenizer import decode_without_hf_kwargs
@@ -396,7 +392,6 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
placeholder_tokens_val=None, placeholder_tokens_val=None,
retraction_counts=recv_obj.retraction_counts, retraction_counts=recv_obj.retraction_counts,
token_steps=recv_obj.token_steps, token_steps=recv_obj.token_steps,
load=recv_obj.load,
dp_ranks=recv_obj.dp_ranks, dp_ranks=recv_obj.dp_ranks,
time_stats=recv_obj.time_stats, time_stats=recv_obj.time_stats,
) )
+689
View File
@@ -0,0 +1,689 @@
"""Load snapshot: publish scheduler load metrics for DP balancing and /v1/loads.
Architecture
------------
Each scheduler periodically publishes a ``LoadSnapshot`` containing its
current load metrics (running reqs, tokens, throughput, ...). Two
transport backends are supported:
**SHM mode** (single-node, default)::
Scheduler ──ShmLoadSnapshotWriter──▶ /dev/shm mmap file
▲
TokenizerManager ──ShmLoadSnapshotReader───┘ (for /v1/loads)
DataParallelController ──ShmLoadSnapshotReader─┘ (for dispatch)
**ZMQ mode** (multi-node DP attention, or ``SGLANG_LOAD_SNAPSHOT_USE_ZMQ=1``)::
Scheduler (any node) ──ZmqLoadSnapshotWriter (PUSH)──▶ network
│
ZmqShmLoadSnapshotReader (PULL, node 0) ◀─────────────────┘
│ drains zmq, writes to SHM
▼
/dev/shm mmap file (node 0)
▲
TokenizerManager / DataParallelController ──ShmLoadSnapshotReader──┘
Shared memory does not work across nodes, so multi-node DP attention
requires the ZMQ transport. The ``ZmqShmLoadSnapshotReader`` on node 0
receives snapshots from all schedulers via zmq PUSH/PULL and writes them
into the local SHM file. All readers (tokenizer, dp_controller) on
node 0 then read from SHM.
``zmq_reader_owner()`` decides which process on node 0 binds the zmq
PULL socket (only one can bind); the other reads plain SHM.
"""
from __future__ import annotations
import fcntl
import hashlib
import logging
import mmap
import os
import struct
from contextlib import contextmanager
from typing import TYPE_CHECKING, Optional
import msgspec
import msgspec.msgpack
import msgspec.structs
from sglang.srt.environ import envs
if TYPE_CHECKING:
from sglang.srt.managers.io_struct import GetLoadsReqOutput
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
DISAGG_MODE_TO_INT = {"null": 0, "prefill": 1, "decode": 2}
INT_TO_DISAGG_MODE = {v: k for k, v in DISAGG_MODE_TO_INT.items()}
def _native(v):
"""Coerce numpy scalars to Python int/float for msgpack encoding."""
if hasattr(v, "item"):
return v.item()
return v
def should_use_zmq(server_args) -> bool:
"""Whether to use zmq PUSH/PULL instead of shared memory for load snapshots.
Shared memory (mmap) only works within a single node. When schedulers
run on multiple nodes (multi-node DP attention), they cannot write to
the SHM file on node 0, so we fall back to zmq transport. The env var
``SGLANG_LOAD_SNAPSHOT_USE_ZMQ`` forces zmq mode for testing.
"""
return (
server_args.enable_dp_attention and server_args.nnodes > 1
) or envs.SGLANG_LOAD_SNAPSHOT_USE_ZMQ.get()
_LOAD_AWARE_METHODS = frozenset({"total_requests", "total_tokens"})
def zmq_reader_owner(server_args, caller: str) -> bool:
"""Decide which process owns the zmq PULL socket.
Exactly one of ``"dp_controller"`` or ``"tokenizer"`` must return True
when zmq mode is active. The owner polls zmq -> SHM; the other reads SHM.
Rules:
- Non-zero node_rank: no tokenizer, dp_controller only launches
schedulers and waits -> nobody owns it.
- dp_size == 1: no dp_controller exists -> tokenizer owns it.
- dp_size > 1, load-aware method: dp_controller polls on every
dispatch via refresh_load_budget() -> dp_controller owns it.
- dp_size > 1, round-robin / other: dp_controller never reads
load data -> tokenizer owns it (polls on /v1/loads calls).
"""
if not should_use_zmq(server_args):
return False
if server_args.node_rank != 0:
return False
if server_args.dp_size == 1:
return caller == "tokenizer"
if server_args.load_balance_method.lower() in _LOAD_AWARE_METHODS:
return caller == "dp_controller"
return caller == "tokenizer"
# ---------------------------------------------------------------------------
# LoadSnapshot data class
# ---------------------------------------------------------------------------
CORE_METRIC_FIELDS = (
"timestamp",
"dp_rank",
"num_running_reqs",
"num_waiting_reqs",
"num_used_tokens",
"num_total_tokens",
"max_total_num_tokens",
"max_running_requests",
"token_usage",
"gen_throughput",
"cache_hit_rate",
"utilization",
)
SECTION_FIELDS = (
(
"memory",
"memory",
"has_memory",
(
("weight_gb", "memory_weight_gb"),
("kv_cache_gb", "memory_kv_cache_gb"),
("graph_gb", "memory_graph_gb"),
("token_capacity", "memory_token_capacity"),
),
),
(
"spec",
"speculative",
"has_speculative",
(
("accept_length", "speculative_accept_length"),
("accept_rate", "speculative_accept_rate"),
),
),
(
"lora",
"lora",
"has_lora",
(
("slots_used", "lora_slots_used"),
("slots_total", "lora_slots_total"),
("utilization", "lora_utilization"),
),
),
(
"disagg",
"disaggregation",
"has_disaggregation",
(
("mode", "disagg_mode"),
("prefill_bootstrap_queue_reqs", "prefill_bootstrap_queue_reqs"),
("prefill_inflight_queue_reqs", "prefill_inflight_queue_reqs"),
("decode_prealloc_queue_reqs", "decode_prealloc_queue_reqs"),
("decode_transfer_queue_reqs", "decode_transfer_queue_reqs"),
("decode_retracted_queue_reqs", "decode_retracted_queue_reqs"),
("kv_transfer_speed_gb_s", "kv_transfer_speed_gb_s"),
("kv_transfer_latency_ms", "kv_transfer_latency_ms"),
),
),
(
"queues",
"queues",
"has_queues",
(
("waiting", "queue_waiting"),
("grammar", "queue_grammar"),
("paused", "queue_paused"),
("retracted", "queue_retracted"),
),
),
)
class LoadSnapshot(msgspec.Struct, omit_defaults=True):
timestamp: float = 0.0
dp_rank: int = 0
num_running_reqs: int = 0
num_waiting_reqs: int = 0
num_used_tokens: int = 0
num_total_tokens: int = 0
max_total_num_tokens: int = 0
max_running_requests: int = 0
token_usage: float = 0.0
gen_throughput: float = 0.0
cache_hit_rate: float = 0.0
utilization: float = 0.0
has_memory: int = 0
memory_weight_gb: float = 0.0
memory_kv_cache_gb: float = 0.0
memory_graph_gb: float = 0.0
memory_token_capacity: int = 0
has_speculative: int = 0
speculative_accept_length: float = 0.0
speculative_accept_rate: float = 0.0
has_lora: int = 0
lora_slots_used: int = 0
lora_slots_total: int = 0
lora_utilization: float = 0.0
has_disaggregation: int = 0
disagg_mode: int = 0
prefill_bootstrap_queue_reqs: int = 0
prefill_inflight_queue_reqs: int = 0
decode_prealloc_queue_reqs: int = 0
decode_transfer_queue_reqs: int = 0
decode_retracted_queue_reqs: int = 0
kv_transfer_speed_gb_s: float = 0.0
kv_transfer_latency_ms: float = 0.0
has_queues: int = 0
queue_waiting: int = 0
queue_grammar: int = 0
queue_paused: int = 0
queue_retracted: int = 0
@classmethod
def from_get_loads_output(cls, output: GetLoadsReqOutput) -> LoadSnapshot:
snapshot: dict = {}
for name in CORE_METRIC_FIELDS:
value = getattr(output, name)
if name == "dp_rank":
snapshot[name] = int(value) if value is not None else 0
else:
snapshot[name] = _native(value)
for _, section_name, present_attr, attrs in SECTION_FIELDS:
section = getattr(output, section_name, None)
snapshot[present_attr] = int(section is not None)
if section is None:
continue
for section_attr, snapshot_attr in attrs:
value = getattr(section, section_attr)
if snapshot_attr == "disagg_mode":
value = DISAGG_MODE_TO_INT.get(value, 0)
else:
value = _native(value)
snapshot[snapshot_attr] = value
return cls(**snapshot)
VALID_SECTIONS = frozenset(
{"core", "memory", "spec", "lora", "disagg", "queues", "all"}
)
def to_dict(self, include: Optional[set[str]] = None) -> dict:
load = {
"dp_rank": self.dp_rank,
"num_running_reqs": self.num_running_reqs,
"num_waiting_reqs": self.num_waiting_reqs,
"num_used_tokens": self.num_used_tokens,
"num_total_tokens": self.num_total_tokens,
"max_total_num_tokens": self.max_total_num_tokens,
"max_running_requests": self.max_running_requests,
"token_usage": self.token_usage,
"gen_throughput": self.gen_throughput,
"cache_hit_rate": self.cache_hit_rate,
"utilization": self.utilization,
}
if include is None or "all" in include:
include_all = True
else:
if not (include <= self.VALID_SECTIONS):
raise ValueError(
f"Invalid include sections: {include - self.VALID_SECTIONS}. "
f"Valid options: {sorted(self.VALID_SECTIONS)}"
)
if include == {"core"}:
return load
include_all = False
for include_key, section_name, present_attr, attrs in SECTION_FIELDS:
if not getattr(self, present_attr):
continue
if not include_all and include_key not in include:
continue
section = {}
for section_attr, snapshot_attr in attrs:
value = getattr(self, snapshot_attr)
if snapshot_attr == "disagg_mode":
value = INT_TO_DISAGG_MODE.get(value, "null")
section[section_attr] = value
load[section_name] = section
return load
snapshot_encoder = msgspec.msgpack.Encoder()
snapshot_decoder = msgspec.msgpack.Decoder(LoadSnapshot)
# ---------------------------------------------------------------------------
# SHM file layout utilities
# ---------------------------------------------------------------------------
MAGIC = b"SLNS"
VERSION = 2
HEADER_STRUCT = struct.Struct("<4sHHI")
SLOT_LEN_STRUCT = struct.Struct("<I")
SLOT_SIZE = 16 * 1024
@contextmanager
def file_lock(fd: int, lock_type: int):
fcntl.flock(fd, lock_type)
try:
yield
finally:
fcntl.flock(fd, fcntl.LOCK_UN)
def shm_path_for(ipc_name: str) -> str:
name = os.path.basename(ipc_name.rstrip("/")) or "default"
safe_name = "".join(c if c.isalnum() or c in "._-" else "_" for c in name)
digest = hashlib.blake2s(ipc_name.encode(), digest_size=4).hexdigest()
return f"/dev/shm/sglang_loads_{safe_name}_{digest}.shm"
def file_size(dp_size: int, slot_size: int = SLOT_SIZE) -> int:
return HEADER_STRUCT.size + dp_size * slot_size
def slot_offset(dp_rank: int, slot_size: int = SLOT_SIZE) -> int:
return HEADER_STRUCT.size + dp_rank * slot_size
# ---------------------------------------------------------------------------
# Writers
# ---------------------------------------------------------------------------
class ShmLoadSnapshotWriter:
def __init__(
self, path: str, dp_size: int, dp_rank: int, publish_interval: int = 1
):
if dp_rank < 0 or dp_rank >= dp_size:
raise ValueError(f"invalid dp_rank={dp_rank} for dp_size={dp_size}")
self.publish_interval = max(1, publish_interval)
self.publish_counter = 0
self.path = path
self.dp_size = dp_size
self.dp_rank = dp_rank
self.slot_size = SLOT_SIZE
self.fd = -1
size = file_size(dp_size, self.slot_size)
self.fd = os.open(path, os.O_CREAT | os.O_RDWR, 0o600)
try:
with file_lock(self.fd, fcntl.LOCK_EX):
os.ftruncate(self.fd, size)
self.mmap = mmap.mmap(self.fd, size, access=mmap.ACCESS_WRITE)
HEADER_STRUCT.pack_into(
self.mmap, 0, MAGIC, VERSION, dp_size, self.slot_size
)
self._write_payload(LoadSnapshot(dp_rank=dp_rank))
except Exception:
if self.fd >= 0:
os.close(self.fd)
raise
def write(self, snapshot: LoadSnapshot) -> None:
if snapshot.dp_rank != self.dp_rank:
raise ValueError(
f"snapshot dp_rank={snapshot.dp_rank} does not match writer dp_rank={self.dp_rank}"
)
with file_lock(self.fd, fcntl.LOCK_EX):
self._write_payload(snapshot)
def _write_payload(self, snapshot: LoadSnapshot) -> None:
payload = snapshot_encoder.encode(snapshot)
max_payload_size = self.slot_size - SLOT_LEN_STRUCT.size
if len(payload) > max_payload_size:
raise ValueError(
f"load snapshot payload size {len(payload)} exceeds slot payload "
f"capacity {max_payload_size}"
)
offset = slot_offset(self.dp_rank, self.slot_size)
payload_start = offset + SLOT_LEN_STRUCT.size
payload_end = payload_start + len(payload)
slot_end = offset + self.slot_size
SLOT_LEN_STRUCT.pack_into(self.mmap, offset, 0)
self.mmap[payload_start:payload_end] = payload
self.mmap[payload_end:slot_end] = b"\0" * (slot_end - payload_end)
SLOT_LEN_STRUCT.pack_into(self.mmap, offset, len(payload))
def close(self) -> None:
self.mmap.close()
os.close(self.fd)
class ZmqLoadSnapshotWriter:
"""Sends load snapshots via zmq PUSH to a ZmqShmLoadSnapshotReader.
CONFLATE is set so only the latest message is kept in the send
buffer when the reader is slower than the writer.
"""
def __init__(
self, endpoint: str, dp_size: int, dp_rank: int, publish_interval: int = 1
):
import zmq as _zmq
if dp_rank < 0 or dp_rank >= dp_size:
raise ValueError(f"invalid dp_rank={dp_rank} for dp_size={dp_size}")
self.publish_interval = max(1, publish_interval)
self.publish_counter = 0
self.dp_size = dp_size
self.dp_rank = dp_rank
self._zmq = _zmq
self._ctx = _zmq.Context.instance()
self._socket = self._ctx.socket(_zmq.PUSH)
self._socket.setsockopt(_zmq.LINGER, 0)
self._socket.setsockopt(_zmq.CONFLATE, 1)
self._socket.connect(endpoint)
def write(self, snapshot: LoadSnapshot) -> None:
if snapshot.dp_rank != self.dp_rank:
raise ValueError(
f"snapshot dp_rank={snapshot.dp_rank} does not match "
f"writer dp_rank={self.dp_rank}"
)
try:
self._socket.send(snapshot_encoder.encode(snapshot), self._zmq.NOBLOCK)
except self._zmq.Again:
pass
def close(self) -> None:
self._socket.close()
# ---------------------------------------------------------------------------
# Readers
# ---------------------------------------------------------------------------
class ShmLoadSnapshotReader:
def __init__(self, path: str, dp_size: int):
self.path = path
self.dp_size = dp_size
self.mmap: Optional[mmap.mmap] = None
self.fd: Optional[int] = None
self.slot_size = SLOT_SIZE
self._header_warning_logged = False
self._attach()
def _attach(self) -> bool:
if self.mmap is not None:
return True
try:
fd = os.open(self.path, os.O_RDONLY)
except FileNotFoundError:
return False
size = os.fstat(fd).st_size
if size < HEADER_STRUCT.size:
os.close(fd)
return False
try:
with file_lock(fd, fcntl.LOCK_SH):
mapped = mmap.mmap(fd, size, access=mmap.ACCESS_READ)
magic, version, dp_size, slot_size = HEADER_STRUCT.unpack_from(
mapped, 0
)
except (OSError, ValueError):
os.close(fd)
return False
if (
magic != MAGIC
or version != VERSION
or dp_size != self.dp_size
or slot_size < SLOT_LEN_STRUCT.size
or size < file_size(self.dp_size, slot_size)
):
mapped.close()
os.close(fd)
if not self._header_warning_logged:
logger.warning("load shm header mismatch at %s", self.path)
self._header_warning_logged = True
return False
self.mmap = mapped
self.fd = fd
self.slot_size = slot_size
return True
def read(self, dp_rank: int) -> Optional[LoadSnapshot]:
if dp_rank < 0 or dp_rank >= self.dp_size:
return None
if not self._attach():
return None
assert self.fd is not None
with file_lock(self.fd, fcntl.LOCK_SH):
return self._read_slot(dp_rank)
def _read_slot(self, dp_rank: int) -> Optional[LoadSnapshot]:
assert self.mmap is not None
offset = slot_offset(dp_rank, self.slot_size)
(payload_len,) = SLOT_LEN_STRUCT.unpack_from(self.mmap, offset)
max_payload_size = self.slot_size - SLOT_LEN_STRUCT.size
if payload_len == 0 or payload_len > max_payload_size:
return None
payload_start = offset + SLOT_LEN_STRUCT.size
payload_end = payload_start + payload_len
try:
return snapshot_decoder.decode(self.mmap[payload_start:payload_end])
except Exception as e:
logger.debug("load snapshot decode failed for rank %s: %s", dp_rank, e)
return None
def read_all(self) -> list[LoadSnapshot]:
if not self._attach():
return []
assert self.fd is not None
with file_lock(self.fd, fcntl.LOCK_SH):
loads = []
for r in range(self.dp_size):
load = self._read_slot(r)
if load is not None:
loads.append(load)
return loads
def close(self) -> None:
if self.mmap is not None:
self.mmap.close()
self.mmap = None
if self.fd is not None:
os.close(self.fd)
self.fd = None
class ZmqShmLoadSnapshotReader:
"""Receives snapshots via zmq PULL from writers, writes to SHM, reads from SHM.
Transparently wraps a ShmLoadSnapshotReader. Every read() / read_all()
first drains the PULL socket into SHM so callers always see fresh data.
"""
def __init__(self, endpoint: str, shm_path: str, dp_size: int):
import zmq as _zmq
self._zmq = _zmq
self._ctx = _zmq.Context.instance()
self._socket = self._ctx.socket(_zmq.PULL)
self._socket.setsockopt(_zmq.LINGER, 0)
self._socket.setsockopt(_zmq.CONFLATE, 1)
self._socket.bind(endpoint)
self._endpoint = endpoint
self._shm_path = shm_path
self.dp_size = dp_size
self._shm_reader = ShmLoadSnapshotReader(shm_path, dp_size)
self._shm_writers: dict[int, ShmLoadSnapshotWriter] = {}
def _poll(self) -> None:
"""Drain zmq messages and write latest per dp_rank to SHM."""
latest: dict[int, LoadSnapshot] = {}
while True:
try:
data = self._socket.recv(self._zmq.NOBLOCK)
except self._zmq.Again:
break
try:
snapshot = snapshot_decoder.decode(data)
if 0 <= snapshot.dp_rank < self.dp_size:
latest[snapshot.dp_rank] = snapshot
except Exception as e:
logger.warning("load snapshot zmq decode failed: %s", e)
for dp_rank, snapshot in latest.items():
if dp_rank not in self._shm_writers:
self._shm_writers[dp_rank] = ShmLoadSnapshotWriter(
self._shm_path, self.dp_size, dp_rank
)
try:
self._shm_writers[dp_rank].write(snapshot)
except Exception as e:
logger.warning(
"load snapshot shm write failed for rank %d: %s", dp_rank, e
)
def read(self, dp_rank: int) -> Optional[LoadSnapshot]:
self._poll()
return self._shm_reader.read(dp_rank)
def read_all(self) -> list[LoadSnapshot]:
self._poll()
return self._shm_reader.read_all()
def close(self) -> None:
for w in self._shm_writers.values():
w.close()
self._shm_writers.clear()
self._shm_reader.close()
self._socket.close()
if self._endpoint.startswith("ipc://"):
try:
os.unlink(self._endpoint[len("ipc://") :])
except OSError:
pass
# ---------------------------------------------------------------------------
# Factory functions
# ---------------------------------------------------------------------------
def _zmq_addr_for(port_args) -> str:
"""Return the zmq PUSH/PULL address from PortArgs.
For dp_attention (TCP mode), uses the ``load_collector_ipc_name`` field
stored in PortArgs. For single-node IPC (env-var override), derives
a deterministic IPC path from ``instance_id``.
"""
ipc_name = getattr(port_args, "load_collector_ipc_name", "")
if ipc_name:
return ipc_name
safe = "".join(
c if c.isalnum() or c in "._-" else "_" for c in port_args.instance_id
)
digest = hashlib.blake2s(port_args.instance_id.encode(), digest_size=4).hexdigest()
return f"ipc:///tmp/sglang_load_collector_{safe}_{digest}.sock"
def create_load_snapshot_writer(
server_args,
port_args,
dp_size: int,
dp_rank: int,
publish_interval: int = 1,
):
"""Return a SHM or ZMQ writer based on server configuration."""
if should_use_zmq(server_args):
return ZmqLoadSnapshotWriter(
_zmq_addr_for(port_args), dp_size, dp_rank, publish_interval
)
return ShmLoadSnapshotWriter(
shm_path_for(port_args.instance_id), dp_size, dp_rank, publish_interval
)
def create_load_snapshot_reader(server_args, port_args, caller: str):
"""Create a load snapshot reader.
Args:
caller: ``"dp_controller"`` or ``"tokenizer"`` -- determines who
binds the zmq PULL socket when zmq mode is active.
"""
dp_size = server_args.dp_size
if zmq_reader_owner(server_args, caller):
return ZmqShmLoadSnapshotReader(
_zmq_addr_for(port_args), shm_path_for(port_args.instance_id), dp_size
)
return ShmLoadSnapshotReader(shm_path_for(port_args.instance_id), dp_size)
+47 -18
View File
@@ -27,11 +27,11 @@ from functools import partial
from http import HTTPStatus from http import HTTPStatus
from typing import Any, Deque, Dict, List, Optional, Tuple, Union from typing import Any, Deque, Dict, List, Optional, Tuple, Union
from sglang.srt.utils.common import suppress_noisy_warnings from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip
suppress_noisy_warnings() suppress_noisy_warnings()
import psutil import psutil # isort: skip
import setproctitle import setproctitle
import torch import torch
import torch.distributed import torch.distributed
@@ -144,6 +144,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromIPCReqInput, UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot, create_load_snapshot_writer
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.overlap_utils import decide_needs_cpu_seq_lens from sglang.srt.managers.overlap_utils import decide_needs_cpu_seq_lens
from sglang.srt.managers.prefill_delayer import ( from sglang.srt.managers.prefill_delayer import (
@@ -164,26 +165,18 @@ from sglang.srt.managers.schedule_policy import (
from sglang.srt.managers.scheduler_components.batch_result_processor import ( from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor, SchedulerBatchResultProcessor,
) )
from sglang.srt.managers.scheduler_components.dp_attn import ( from sglang.srt.managers.scheduler_components.dp_attn import SchedulerDPAttnAdapter
SchedulerDPAttnAdapter, from sglang.srt.managers.scheduler_components.flush_wrapper import SchedulerFlushWrapper
)
from sglang.srt.managers.scheduler_components.flush_wrapper import (
SchedulerFlushWrapper,
)
from sglang.srt.managers.scheduler_components.idle_sleeper import IdleSleeper from sglang.srt.managers.scheduler_components.idle_sleeper import IdleSleeper
from sglang.srt.managers.scheduler_components.invariant_checker import ( from sglang.srt.managers.scheduler_components.invariant_checker import (
SchedulerInvariantChecker, SchedulerInvariantChecker,
create_scheduler_watchdog, create_scheduler_watchdog,
) )
from sglang.srt.managers.scheduler_components.ipc_channels import ( from sglang.srt.managers.scheduler_components.ipc_channels import SchedulerIpcChannels
SchedulerIpcChannels,
)
from sglang.srt.managers.scheduler_components.kv_events_publisher import ( from sglang.srt.managers.scheduler_components.kv_events_publisher import (
SchedulerKvEventsPublisher, SchedulerKvEventsPublisher,
) )
from sglang.srt.managers.scheduler_components.load_inquirer import ( from sglang.srt.managers.scheduler_components.load_inquirer import SchedulerLoadInquirer
SchedulerLoadInquirer,
)
from sglang.srt.managers.scheduler_components.logprob_result_processor import ( from sglang.srt.managers.scheduler_components.logprob_result_processor import (
SchedulerLogprobResultProcessor, SchedulerLogprobResultProcessor,
) )
@@ -606,6 +599,22 @@ class Scheduler(
), ),
) )
self.load_snapshot_writer = None
if not is_rank_zero:
return
dp_rank = self.ps.dp_rank if self.ps.dp_rank is not None else 0
try:
self.load_snapshot_writer = create_load_snapshot_writer(
self.server_args,
port_args,
self.ps.dp_size,
dp_rank,
publish_interval=self.server_args.load_snapshot_publish_interval,
)
except Exception as e:
logger.warning("load snapshot writer init failed: %s", e)
def init_idle_sleeper(self) -> None: def init_idle_sleeper(self) -> None:
if ( if (
self.ps.pp_rank == 0 self.ps.pp_rank == 0
@@ -622,6 +631,24 @@ class Scheduler(
else: else:
self.idle_sleeper = None self.idle_sleeper = None
def publish_load_snapshot(self, force: bool = False):
writer = self.load_snapshot_writer
if writer is None:
return
if not force:
writer.publish_counter += 1
if writer.publish_counter < writer.publish_interval:
return
writer.publish_counter = 0
try:
result = self.load_inquirer.get_loads(GetLoadsReqInput(include=["all"]))
writer.write(LoadSnapshot.from_get_loads_output(result))
except Exception as e:
logger.warning("load snapshot publish failed: %s", e)
def handle_get_loads_req(self, req: GetLoadsReqInput):
return self.load_inquirer.get_loads(req)
def init_tokenizer(self): def init_tokenizer(self):
server_args = self.server_args server_args = self.server_args
self.is_generation = self.model_config.is_generation self.is_generation = self.model_config.is_generation
@@ -1313,10 +1340,7 @@ class Scheduler(
self.load_lora_adapter_from_tensors, self.load_lora_adapter_from_tensors,
), ),
(UnloadLoRAAdapterReqInput, self.unload_lora_adapter), (UnloadLoRAAdapterReqInput, self.unload_lora_adapter),
( (GetLoadsReqInput, self.handle_get_loads_req),
GetLoadsReqInput,
lambda req: self.load_inquirer.get_loads(req),
),
(PauseGenerationReqInput, self.pause_generation), (PauseGenerationReqInput, self.pause_generation),
(ContinueGenerationReqInput, self.continue_generation), (ContinueGenerationReqInput, self.continue_generation),
(DumperControlReqInput, self.handle_dumper_control), (DumperControlReqInput, self.handle_dumper_control),
@@ -3100,6 +3124,8 @@ class Scheduler(
batch: ScheduleBatch, batch: ScheduleBatch,
result: Union[GenerationBatchResult, EmbeddingBatchResult], result: Union[GenerationBatchResult, EmbeddingBatchResult],
): ):
self.publish_load_snapshot(force=batch.forward_mode.is_extend())
if batch.forward_mode.is_decode(): if batch.forward_mode.is_decode():
self.batch_result_processor.process_batch_result_decode(batch, result) self.batch_result_processor.process_batch_result_decode(batch, result)
elif batch.forward_mode.is_extend(): elif batch.forward_mode.is_extend():
@@ -3205,6 +3231,9 @@ class Scheduler(
# reset device timer window so idle time isn't counted # reset device timer window so idle time isn't counted
self.metrics_reporter.reset_device_timer_window() self.metrics_reporter.reset_device_timer_window()
# Publish the idle state so /get_loads and DP balancing do not see stale load.
self.publish_load_snapshot(force=True)
# sleep until next event # sleep until next event
self.maybe_sleep_on_idle() self.maybe_sleep_on_idle()
@@ -4,14 +4,7 @@ import asyncio
import logging import logging
import time import time
import uuid import uuid
from typing import ( from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
TYPE_CHECKING,
Any,
Dict,
List,
Optional,
Tuple,
)
import fastapi import fastapi
@@ -39,7 +32,6 @@ from sglang.srt.managers.io_struct import (
FlushCacheReqOutput, FlushCacheReqOutput,
GetInternalStateReq, GetInternalStateReq,
GetInternalStateReqOutput, GetInternalStateReqOutput,
GetLoadsReqInput,
GetLoadsReqOutput, GetLoadsReqOutput,
GetWeightsByNameReqInput, GetWeightsByNameReqInput,
GetWeightsByNameReqOutput, GetWeightsByNameReqOutput,
@@ -79,6 +71,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqInput, UpdateWeightsFromTensorReqInput,
UpdateWeightsFromTensorReqOutput, UpdateWeightsFromTensorReqOutput,
) )
from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.server_args import LoRARef, ServerArgs
from sglang.srt.utils import get_bool_env_var from sglang.srt.utils import get_bool_env_var
from sglang.utils import TypeBasedDispatcher from sglang.utils import TypeBasedDispatcher
@@ -809,41 +802,27 @@ class TokenizerControlMixin:
self: TokenizerManager, self: TokenizerManager,
include: Optional[List[str]] = None, include: Optional[List[str]] = None,
dp_rank: Optional[int] = None, dp_rank: Optional[int] = None,
) -> List[GetLoadsReqOutput]: ) -> List[LoadSnapshot]:
""" """
Get comprehensive load metrics for /v1/loads endpoint. Get load snapshots for /v1/loads endpoint.
Args: Args:
include: List of sections to include. Options: core, memory, spec, lora, disagg, queues, all include: List of sections to include. Options: core, memory, spec, lora, disagg, queues, all
dp_rank: Optional filter for specific DP rank dp_rank: Optional filter for specific DP rank
Returns: Returns:
List of GetLoadsReqOutput, one per scheduler (filtered by dp_rank if specified) List of LoadSnapshot, one per scheduler (filtered by dp_rank if specified)
""" """
self.auto_create_handle_loop() self.auto_create_handle_loop()
# Always request all sections from scheduler — watching mode shares if dp_rank is not None and (dp_rank < 0 or dp_rank >= self.server_args.dp_size):
# results across concurrent callers, so we fetch full data and filter here. return []
req = GetLoadsReqInput(include=["all"], dp_rank=None)
results = await self.get_loads_communicator(req)
# Filter by dp_rank if specified reader = self.load_snapshot_reader
if dp_rank is not None: if dp_rank is not None:
results = [r for r in results if r.dp_rank == dp_rank] load = reader.read(dp_rank)
results = [load] if load is not None else []
# Filter optional sections client-side (scheduler always returns all) else:
if include and "all" not in include: results = reader.read_all()
include_set = set(include)
_section_attrs = {
"memory": "memory",
"spec": "speculative",
"lora": "lora",
"disagg": "disaggregation",
"queues": "queues",
}
for r in results:
for key, attr in _section_attrs.items():
if key not in include_set:
setattr(r, attr, None)
return results return results
@@ -73,16 +73,14 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput, TokenizedGenerateReqInput,
UpdateWeightFromDiskReqInput, UpdateWeightFromDiskReqInput,
UpdateWeightFromDiskReqOutput, UpdateWeightFromDiskReqOutput,
WatchLoadUpdateReq,
) )
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import MultimodalDataItem from sglang.srt.managers.schedule_batch import MultimodalDataItem
from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
from sglang.srt.managers.tokenizer_manager_score_mixin import ( from sglang.srt.managers.tokenizer_manager_score_mixin import TokenizerManagerScoreMixin
TokenizerManagerScoreMixin,
)
from sglang.srt.managers.utils import is_health_check_generate_req from sglang.srt.managers.utils import is_health_check_generate_req
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.metrics_collector import ( from sglang.srt.observability.metrics_collector import (
@@ -378,6 +376,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Make sure that each request carries the tokenizer_ipc_name for response routing # Make sure that each request carries the tokenizer_ipc_name for response routing
self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler) self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler)
self.load_snapshot_reader = create_load_snapshot_reader(
self.server_args,
port_args,
caller="tokenizer",
)
def init_running_status(self): def init_running_status(self):
# Request states # Request states
self.rid_to_state: Dict[str, ReqState] = {} self.rid_to_state: Dict[str, ReqState] = {}
@@ -1947,16 +1951,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
for s in pending_notify.values(): for s in pending_notify.values():
s.event.set() s.event.set()
# When skip_tokenizer_init is enabled, tokensizer_manager receives
# BatchTokenIDOutput.
if (
self.server_args.dp_size > 1
and isinstance(recv_obj, (BatchStrOutput, BatchTokenIDOutput))
and recv_obj.load is not None
):
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
self.send_to_scheduler.send_pyobj(load_update_req)
def add_logprob_to_meta_info( def add_logprob_to_meta_info(
self, self,
meta_info: dict, meta_info: dict,
@@ -73,6 +73,7 @@ class RayDataParallelController(DataParallelController):
tmp_port_args = PortArgs.init_new(server_args) tmp_port_args = PortArgs.init_new(server_args)
tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name
tmp_port_args.instance_id = port_args.instance_id
# Hold NCCL port so the next DP rank gets a different one # Hold NCCL port so the next DP rank gets a different one
sockets.append(bind_port(tmp_port_args.nccl_port)) sockets.append(bind_port(tmp_port_args.nccl_port))
@@ -159,6 +160,7 @@ class RayDataParallelController(DataParallelController):
) )
# All DP ranks share the same NCCL port (reuse TP group) # All DP ranks share the same NCCL port (reuse TP group)
rank_port_args.nccl_port = port_args.nccl_port rank_port_args.nccl_port = port_args.nccl_port
rank_port_args.instance_id = port_args.instance_id
# The detokenizer and tokenizer bind using the # The detokenizer and tokenizer bind using the
# original port_args addresses (127.0.0.1 when # original port_args addresses (127.0.0.1 when
# dist_init_addr is unset). Scheduler actors must # dist_init_addr is unset). Scheduler actors must
+26
View File
@@ -25,6 +25,7 @@ import logging
import os import os
import random import random
import tempfile import tempfile
import uuid
from functools import cached_property from functools import cached_property
from typing import Any, Callable, Dict, List, Literal, Optional, Union from typing import Any, Callable, Dict, List, Literal, Optional, Union
@@ -451,6 +452,7 @@ class ServerArgs:
base_gpu_id: int = 0 base_gpu_id: int = 0
gpu_id_step: int = 1 gpu_id_step: int = 1
sleep_on_idle: bool = False sleep_on_idle: bool = False
load_snapshot_publish_interval: int = 15
use_ray: bool = False use_ray: bool = False
custom_sigquit_handler: Optional[Callable] = None custom_sigquit_handler: Optional[Callable] = None
@@ -4956,6 +4958,12 @@ class ServerArgs:
action="store_true", action="store_true",
help="Reduce CPU usage when sglang is idle.", help="Reduce CPU usage when sglang is idle.",
) )
parser.add_argument(
"--load-snapshot-publish-interval",
type=int,
default=ServerArgs.load_snapshot_publish_interval,
help="Publish load snapshot to shared memory every N decode iterations. Prefill and idle always publish immediately.",
)
parser.add_argument( parser.add_argument(
"--use-ray", "--use-ray",
action="store_true", action="store_true",
@@ -7759,6 +7767,14 @@ class PortArgs:
# The ipc filename for Tokenizer and worker tokenizer # The ipc filename for Tokenizer and worker tokenizer
tokenizer_worker_ipc_name: Optional[str] tokenizer_worker_ipc_name: Optional[str]
# zmq address for load snapshot PUSH/PULL (dp-attention TCP mode only;
# empty when IPC mode derives the address from instance_id).
load_collector_ipc_name: str = ""
# Stable token shared by all processes in one server instance, used to
# derive the /dev/shm path for load snapshots.
instance_id: str = ""
@staticmethod @staticmethod
def init_new( def init_new(
server_args: ServerArgs, server_args: ServerArgs,
@@ -7777,6 +7793,8 @@ class PortArgs:
f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}" f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
) )
instance_id = uuid.uuid4().hex[:12]
if not server_args.enable_dp_attention: if not server_args.enable_dp_attention:
# Normal case, use IPC within a single node # Normal case, use IPC within a single node
return PortArgs( return PortArgs(
@@ -7787,6 +7805,7 @@ class PortArgs:
rpc_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", rpc_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}",
metrics_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}", metrics_ipc_name=f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}",
tokenizer_worker_ipc_name=tokenizer_worker_ipc_name, tokenizer_worker_ipc_name=tokenizer_worker_ipc_name,
instance_id=instance_id,
) )
else: else:
# DP attention. Use TCP + port to handle both single-node and multi-node. # DP attention. Use TCP + port to handle both single-node and multi-node.
@@ -7801,6 +7820,7 @@ class PortArgs:
detokenizer_port = port_base + 1 detokenizer_port = port_base + 1
rpc_port = port_base + 2 rpc_port = port_base + 2
metrics_port = port_base + 3 metrics_port = port_base + 3
load_collector_port = port_base + 5
if dp_rank is None: if dp_rank is None:
# TokenizerManager to DataParallelController # TokenizerManager to DataParallelController
scheduler_input_port = port_base + 4 scheduler_input_port = port_base + 4
@@ -7816,6 +7836,8 @@ class PortArgs:
wait_port_available(nccl_port, "nccl_port") wait_port_available(nccl_port, "nccl_port")
wait_port_available(rpc_port, "rpc_port") wait_port_available(rpc_port, "rpc_port")
wait_port_available(metrics_port, "metrics_port") wait_port_available(metrics_port, "metrics_port")
if server_args.nnodes > 1:
wait_port_available(load_collector_port, "load_collector_port")
# Check scheduler_input_port only for dp. # Check scheduler_input_port only for dp.
# Skip check when using worker_ports since the port is already bound by our ZMQ socket # Skip check when using worker_ports since the port is already bound by our ZMQ socket
if dp_rank is None or worker_ports is None: if dp_rank is None or worker_ports is None:
@@ -7838,6 +7860,10 @@ class PortArgs:
rpc_ipc_name=NetworkAddress(dist_init_host, rpc_port).to_tcp(), rpc_ipc_name=NetworkAddress(dist_init_host, rpc_port).to_tcp(),
metrics_ipc_name=NetworkAddress(dist_init_host, metrics_port).to_tcp(), metrics_ipc_name=NetworkAddress(dist_init_host, metrics_port).to_tcp(),
tokenizer_worker_ipc_name=tokenizer_worker_ipc_name, tokenizer_worker_ipc_name=tokenizer_worker_ipc_name,
load_collector_ipc_name=NetworkAddress(
dist_init_host, load_collector_port
).to_tcp(),
instance_id=instance_id,
) )
@@ -0,0 +1,96 @@
"""Integration tests for load snapshot with real servers.
Tests [no dp, normal dp] x [zmq, shm] by launching real servers
and querying /v1/loads.
"""
import json
import time
import unittest
import urllib.request
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=300, stage="base-b", runner_config="2-gpu-large")
def _query_loads(base_url, retries=5, interval=2.0):
url = f"{base_url}/v1/loads"
for attempt in range(retries):
try:
resp = urllib.request.urlopen(url, timeout=5)
data = json.loads(resp.read())
if data.get("loads"):
return data
except Exception:
pass
if attempt < retries - 1:
time.sleep(interval)
try:
resp = urllib.request.urlopen(url, timeout=5)
return json.loads(resp.read())
except Exception:
return {"loads": []}
def _launch_and_check(test_case, other_args=None, env=None, expected_dp_size=1):
process = popen_launch_server(
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_URL_FOR_TEST,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args or [],
env=env,
)
try:
time.sleep(5)
data = _query_loads(DEFAULT_URL_FOR_TEST)
loads = data.get("loads", [])
test_case.assertGreater(len(loads), 0, f"Expected non-empty loads, got: {data}")
test_case.assertEqual(len(loads), expected_dp_size)
dp_ranks = sorted(l["dp_rank"] for l in loads)
test_case.assertEqual(dp_ranks, list(range(expected_dp_size)))
for load in loads:
test_case.assertGreater(load["max_total_num_tokens"], 0)
finally:
kill_process_tree(process.pid)
class TestLoadSnapshotNoDP(CustomTestCase):
def test_shm_backend(self):
_launch_and_check(self, expected_dp_size=1)
def test_zmq_backend(self):
_launch_and_check(
self,
env={"SGLANG_LOAD_SNAPSHOT_USE_ZMQ": "1"},
expected_dp_size=1,
)
class TestLoadSnapshotNormalDP(CustomTestCase):
def test_shm_backend(self):
_launch_and_check(
self,
other_args=["--dp", "2"],
expected_dp_size=2,
)
def test_zmq_backend(self):
_launch_and_check(
self,
other_args=["--dp", "2"],
env={"SGLANG_LOAD_SNAPSHOT_USE_ZMQ": "1"},
expected_dp_size=2,
)
if __name__ == "__main__":
unittest.main()
@@ -1,75 +1,185 @@
"""Unit tests for /v1/loads _compute_aggregate. """Unit tests for /v1/loads load snapshot response behavior."""
Narrow scope: lock in the semantic of new aggregate keys added by this PR
(total_used_tokens vs total_tokens). Trivial helpers (dict filtering,
zero-init branch) are not covered — they would just restate Python.
"""
import asyncio
import os
import tempfile
import unittest import unittest
from types import SimpleNamespace
from sglang.srt.entrypoints.v1_loads import _compute_aggregate import msgspec.msgpack
from sglang.srt.entrypoints.v1_loads import get_loads
from sglang.srt.managers.load_snapshot import (
HEADER_STRUCT,
MAGIC,
SLOT_LEN_STRUCT,
SLOT_SIZE,
VERSION,
LoadSnapshot,
ShmLoadSnapshotReader,
ShmLoadSnapshotWriter,
slot_offset,
)
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_cpu_ci(est_time=10, suite="base-a-test-cpu")
def _load( def _temp_path() -> str:
*, fd, path = tempfile.mkstemp()
os.close(fd)
os.unlink(path)
return path
class _FakeTokenizerManager(TokenizerControlMixin):
def __init__(self, reader, dp_size: int):
self.load_snapshot_reader = reader
self.server_args = SimpleNamespace(
dp_size=dp_size,
enable_dp_attention=False,
nnodes=1,
)
def auto_create_handle_loop(self):
pass
class _FakeHttpTokenizerManager:
metrics_collector = None
def __init__(self, loads):
self.loads = loads
async def get_loads(self, include=None, dp_rank=None):
results = []
for load in self.loads:
if dp_rank is not None and load.dp_rank != dp_rank:
continue
results.append(load)
return results
class TestLoadsResponse(CustomTestCase):
def test_response_omits_server_side_aggregate_and_redundant_fields(self):
manager = _FakeHttpTokenizerManager(
[
LoadSnapshot(
dp_rank=0, dp_rank=0,
running=0, num_running_reqs=3,
waiting=0, num_waiting_reqs=2,
used=0, num_total_tokens=256,
total=0, )
token_usage=0.0,
throughput=0.0,
utilization=0.0,
):
return {
"dp_rank": dp_rank,
"num_running_reqs": running,
"num_waiting_reqs": waiting,
"num_used_tokens": used,
"num_total_tokens": total,
"token_usage": token_usage,
"gen_throughput": throughput,
"utilization": utilization,
}
class TestComputeAggregate(CustomTestCase):
def test_multi_dp_rank_sums(self):
agg = _compute_aggregate(
[
_load(dp_rank=0, running=3, waiting=1, used=50, total=70),
_load(dp_rank=1, running=5, waiting=2, used=80, total=100),
_load(dp_rank=2, running=0, waiting=4, used=0, total=40),
] ]
) )
self.assertEqual(agg["total_running_reqs"], 8)
self.assertEqual(agg["total_waiting_reqs"], 7)
self.assertEqual(agg["total_reqs"], 15)
self.assertEqual(agg["total_used_tokens"], 130)
self.assertEqual(agg["total_tokens"], 210)
def test_averages_over_dp_count(self): response = asyncio.run(get_loads(tokenizer_manager=manager))
agg = _compute_aggregate(
[ self.assertNotIn("dp_rank_count", response)
_load(token_usage=0.6, throughput=100.0, utilization=0.5), self.assertNotIn("aggregate", response)
_load(token_usage=0.8, throughput=200.0, utilization=0.7), self.assertEqual(len(response["loads"]), 1)
] self.assertNotIn("num_total_reqs", response["loads"][0])
self.assertEqual(response["loads"][0]["num_running_reqs"], 3)
self.assertEqual(response["loads"][0]["num_waiting_reqs"], 2)
class TestGetLoads(CustomTestCase):
def test_load_snapshot_wire_format_is_msgpack_slots(self):
path = _temp_path()
writer = ShmLoadSnapshotWriter(path, dp_size=2, dp_rank=1)
try:
writer.write(
LoadSnapshot(
dp_rank=1,
num_running_reqs=3,
num_waiting_reqs=2,
token_usage=0.25,
)
) )
self.assertAlmostEqual(agg["avg_token_usage"], 0.7)
self.assertAlmostEqual(agg["avg_throughput"], 150.0)
self.assertAlmostEqual(agg["avg_utilization"], 0.6)
def test_total_tokens_differs_from_total_used_tokens(self): with open(path, "rb") as f:
# Regression: total_tokens sums num_total_tokens, NOT num_used_tokens. data = f.read()
# Gateway reads aggregate.total_tokens for DP load estimation, so a
# silent swap would under-report load. self.assertEqual(len(data), HEADER_STRUCT.size + 2 * SLOT_SIZE)
agg = _compute_aggregate([_load(used=10, total=30), _load(used=20, total=45)]) magic, version, dp_size, slot_size = HEADER_STRUCT.unpack_from(data, 0)
self.assertEqual(agg["total_used_tokens"], 30) self.assertEqual(magic, MAGIC)
self.assertEqual(agg["total_tokens"], 75) self.assertEqual(version, VERSION)
self.assertEqual(dp_size, 2)
self.assertEqual(slot_size, SLOT_SIZE)
offset = slot_offset(1, slot_size)
(payload_len,) = SLOT_LEN_STRUCT.unpack_from(data, offset)
payload_start = offset + SLOT_LEN_STRUCT.size
payload = data[payload_start : payload_start + payload_len]
decoded = msgspec.msgpack.decode(payload)
self.assertEqual(decoded["dp_rank"], 1)
self.assertEqual(decoded["num_running_reqs"], 3)
self.assertEqual(decoded["num_waiting_reqs"], 2)
self.assertEqual(decoded["token_usage"], 0.25)
finally:
writer.close()
if os.path.exists(path):
os.unlink(path)
def test_reads_snapshot_and_filters_sections(self):
path = _temp_path()
writer = ShmLoadSnapshotWriter(path, dp_size=1, dp_rank=0)
reader = ShmLoadSnapshotReader(path, dp_size=1)
try:
initial_load = reader.read(0)
self.assertIsNotNone(initial_load)
self.assertEqual(initial_load.num_total_tokens, 0)
writer.write(
LoadSnapshot(
dp_rank=0,
timestamp=1.25,
num_running_reqs=3,
num_waiting_reqs=2,
num_used_tokens=128,
num_total_tokens=256,
max_total_num_tokens=4096,
token_usage=0.125,
gen_throughput=99.5,
cache_hit_rate=0.75,
utilization=0.5,
max_running_requests=128,
has_disaggregation=1,
disagg_mode=2,
decode_transfer_queue_reqs=4,
has_queues=1,
queue_waiting=2,
queue_grammar=1,
queue_paused=0,
queue_retracted=3,
)
)
manager = _FakeTokenizerManager(reader, dp_size=1)
loads = asyncio.run(manager.get_loads(include=["core"], dp_rank=0))
self.assertEqual(len(loads), 1)
self.assertEqual(loads[0].num_total_tokens, 256)
d = loads[0].to_dict({"core"})
self.assertNotIn("disaggregation", d)
self.assertNotIn("queues", d)
loads_all = asyncio.run(manager.get_loads(include=["all"], dp_rank=0))
d_all = loads_all[0].to_dict()
self.assertIn("disaggregation", d_all)
self.assertIn("queues", d_all)
finally:
reader.close()
writer.close()
if os.path.exists(path):
os.unlink(path)
if __name__ == "__main__": if __name__ == "__main__":
@@ -11,11 +11,12 @@ if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
is exercised as the real method, no mock. is exercised as the real method, no mock.
""" """
import dataclasses
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock from unittest.mock import MagicMock
import msgspec.structs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
@@ -26,29 +27,20 @@ from sglang.srt.managers.data_parallel_controller import (
DPBudget, DPBudget,
LoadBalanceMethod, LoadBalanceMethod,
) )
from sglang.srt.managers.io_struct import GetLoadsReqOutput, WatchLoadUpdateReq from sglang.srt.managers.load_snapshot import LoadSnapshot
register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_cpu_ci(est_time=11, suite="base-a-test-cpu")
_BASE_LOAD = GetLoadsReqOutput( _BASE_LOAD = msgspec.structs.replace(
dp_rank=0, LoadSnapshot(dp_rank=0),
timestamp=0.0,
num_running_reqs=0,
num_waiting_reqs=0,
num_used_tokens=0,
num_total_tokens=0,
max_total_num_tokens=4096, max_total_num_tokens=4096,
token_usage=0.0,
gen_throughput=0.0,
cache_hit_rate=0.0,
utilization=0.0,
max_running_requests=128, max_running_requests=128,
) )
def _load(**overrides) -> GetLoadsReqOutput: def _load(**overrides) -> LoadSnapshot:
return dataclasses.replace(_BASE_LOAD, **overrides) return msgspec.structs.replace(_BASE_LOAD, **overrides)
def _make_controller(dp_size: int) -> DataParallelController: def _make_controller(dp_size: int) -> DataParallelController:
@@ -74,45 +66,53 @@ class TestDPBudgetUpdateBudget(CustomTestCase):
def test_maps_running_plus_waiting_to_total_requests(self): def test_maps_running_plus_waiting_to_total_requests(self):
budget = DPBudget(dp_size=2) budget = DPBudget(dp_size=2)
budget.update_budget( budget.update_budget(
WatchLoadUpdateReq( [
loads=[ _load(dp_rank=0, timestamp=1.0, num_running_reqs=3, num_waiting_reqs=2),
_load(dp_rank=0, num_running_reqs=3, num_waiting_reqs=2), _load(dp_rank=1, timestamp=1.0, num_running_reqs=5, num_waiting_reqs=1),
_load(dp_rank=1, num_running_reqs=5, num_waiting_reqs=1),
] ]
) )
)
self.assertEqual(budget.total_requests, [5, 6]) self.assertEqual(budget.total_requests, [5, 6])
def test_maps_num_total_tokens_not_num_used_tokens(self): def test_maps_num_total_tokens_not_num_used_tokens(self):
# Reads num_total_tokens (used + pending prefill), NOT num_used_tokens.
# A silent swap here would break DP balance for long-prompt workloads.
budget = DPBudget(dp_size=2) budget = DPBudget(dp_size=2)
budget.update_budget( budget.update_budget(
WatchLoadUpdateReq( [
loads=[ _load(
_load(dp_rank=0, num_used_tokens=100, num_total_tokens=150), dp_rank=0, timestamp=1.0, num_used_tokens=100, num_total_tokens=150
_load(dp_rank=1, num_used_tokens=80, num_total_tokens=80), ),
_load(
dp_rank=1, timestamp=1.0, num_used_tokens=80, num_total_tokens=80
),
] ]
) )
)
self.assertEqual(budget.total_tokens, [150, 80]) self.assertEqual(budget.total_tokens, [150, 80])
def test_partial_update_only_affects_reported_rank(self): def test_partial_update_only_affects_reported_rank(self):
budget = DPBudget(dp_size=3) budget = DPBudget(dp_size=3)
budget.total_requests = [10, 20, 30]
budget.total_tokens = [100, 200, 300]
budget.update_budget( budget.update_budget(
WatchLoadUpdateReq( [
loads=[ _load(
dp_rank=0, timestamp=1.0, num_running_reqs=10, num_total_tokens=100
),
_load(
dp_rank=1, timestamp=1.0, num_running_reqs=20, num_total_tokens=200
),
_load(
dp_rank=2, timestamp=1.0, num_running_reqs=30, num_total_tokens=300
),
]
)
budget.update_budget(
[
_load( _load(
dp_rank=1, dp_rank=1,
timestamp=2.0,
num_running_reqs=1, num_running_reqs=1,
num_waiting_reqs=1, num_waiting_reqs=1,
num_total_tokens=50, num_total_tokens=50,
) )
] ]
) )
)
self.assertEqual(budget.total_requests, [10, 2, 30]) self.assertEqual(budget.total_requests, [10, 2, 30])
self.assertEqual(budget.total_tokens, [100, 50, 300]) self.assertEqual(budget.total_tokens, [100, 50, 300])
@@ -0,0 +1,355 @@
"""Unit tests for LoadSnapshot SHM and ZMQ backends."""
import os
import tempfile
import time
import unittest
from types import SimpleNamespace
from sglang.srt.managers.load_snapshot import (
LoadSnapshot,
ShmLoadSnapshotReader,
ShmLoadSnapshotWriter,
ZmqLoadSnapshotWriter,
ZmqShmLoadSnapshotReader,
_zmq_addr_for,
create_load_snapshot_reader,
create_load_snapshot_writer,
should_use_zmq,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
def _temp_path() -> str:
fd, path = tempfile.mkstemp()
os.close(fd)
os.unlink(path)
return path
def _ipc_addr() -> str:
fd, path = tempfile.mkstemp(prefix="sglang_test_zmq_", suffix=".sock")
os.close(fd)
os.unlink(path)
return f"ipc://{path}"
def _warmup_zmq(writers, reader, attempts=20, interval=0.05):
"""Send warmup messages until the reader receives from all writers."""
expected = {w.dp_rank for w in writers}
received = set()
for _ in range(attempts):
for w in writers:
w.write(LoadSnapshot(dp_rank=w.dp_rank, timestamp=-1.0, num_running_reqs=0))
time.sleep(interval)
for rank in expected:
load = reader.read(rank)
if load is not None:
received.add(rank)
if received >= expected:
return
raise RuntimeError(f"warmup failed: expected {expected}, received {received}")
class TestShmRoundTrip(CustomTestCase):
def test_single_rank_write_read(self):
path = _temp_path()
writer = ShmLoadSnapshotWriter(path, dp_size=1, dp_rank=0)
reader = ShmLoadSnapshotReader(path, dp_size=1)
try:
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=5, timestamp=1.0))
load = reader.read(0)
self.assertIsNotNone(load)
self.assertEqual(load.num_running_reqs, 5)
self.assertEqual(load.timestamp, 1.0)
finally:
reader.close()
writer.close()
if os.path.exists(path):
os.unlink(path)
def test_multi_rank_write_read_all(self):
path = _temp_path()
writers = []
try:
for rank in range(4):
w = ShmLoadSnapshotWriter(path, dp_size=4, dp_rank=rank)
w.write(
LoadSnapshot(
dp_rank=rank,
num_running_reqs=rank * 10,
timestamp=1.0,
)
)
writers.append(w)
reader = ShmLoadSnapshotReader(path, dp_size=4)
loads = reader.read_all()
self.assertEqual(len(loads), 4)
for i, load in enumerate(loads):
self.assertEqual(load.dp_rank, i)
self.assertEqual(load.num_running_reqs, i * 10)
reader.close()
finally:
for w in writers:
w.close()
if os.path.exists(path):
os.unlink(path)
def test_reader_empty_before_writer(self):
path = _temp_path()
reader = ShmLoadSnapshotReader(path, dp_size=2)
self.assertEqual(reader.read_all(), [])
self.assertIsNone(reader.read(0))
reader.close()
class TestZmqRoundTrip(CustomTestCase):
def test_single_rank_zmq_to_shm(self):
shm_path = _temp_path()
addr = _ipc_addr()
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=2)
writer = ZmqLoadSnapshotWriter(addr, dp_size=2, dp_rank=0)
try:
_warmup_zmq([writer], reader)
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=7, timestamp=2.0))
time.sleep(0.05)
load = reader.read(0)
self.assertIsNotNone(load)
self.assertEqual(load.num_running_reqs, 7)
self.assertEqual(load.timestamp, 2.0)
finally:
writer.close()
reader.close()
if os.path.exists(shm_path):
os.unlink(shm_path)
def test_multi_rank_zmq(self):
shm_path = _temp_path()
addr = _ipc_addr()
dp_size = 4
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size)
writers = []
try:
for rank in range(dp_size):
w = ZmqLoadSnapshotWriter(addr, dp_size, dp_rank=rank)
writers.append(w)
_warmup_zmq(writers, reader)
for rank, w in enumerate(writers):
w.write(
LoadSnapshot(dp_rank=rank, num_running_reqs=rank + 1, timestamp=3.0)
)
time.sleep(0.05)
loads = reader.read_all()
self.assertEqual(len(loads), dp_size)
for load in loads:
self.assertEqual(load.num_running_reqs, load.dp_rank + 1)
finally:
for w in writers:
w.close()
reader.close()
if os.path.exists(shm_path):
os.unlink(shm_path)
def test_read_returns_latest(self):
shm_path = _temp_path()
addr = _ipc_addr()
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=1)
writer = ZmqLoadSnapshotWriter(addr, dp_size=1, dp_rank=0)
try:
_warmup_zmq([writer], reader)
for i in range(10):
writer.write(
LoadSnapshot(dp_rank=0, num_running_reqs=i, timestamp=float(i))
)
time.sleep(0.05)
load = reader.read(0)
self.assertIsNotNone(load)
self.assertEqual(load.num_running_reqs, 9)
self.assertEqual(load.timestamp, 9.0)
finally:
writer.close()
reader.close()
if os.path.exists(shm_path):
os.unlink(shm_path)
def test_zmq_writer_noblock_without_reader(self):
addr = _ipc_addr()
writer = ZmqLoadSnapshotWriter(addr, dp_size=1, dp_rank=0)
try:
writer.write(LoadSnapshot(dp_rank=0, num_running_reqs=1, timestamp=1.0))
finally:
writer.close()
ipc_path = addr[len("ipc://") :]
if os.path.exists(ipc_path):
os.unlink(ipc_path)
def test_reader_ipc_cleanup(self):
addr = _ipc_addr()
shm_path = _temp_path()
ipc_path = addr[len("ipc://") :]
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size=1)
self.assertTrue(os.path.exists(ipc_path))
reader.close()
self.assertFalse(os.path.exists(ipc_path))
if os.path.exists(shm_path):
os.unlink(shm_path)
class TestFactoryFunctions(CustomTestCase):
def test_shm_mode(self):
server_args = SimpleNamespace(
enable_dp_attention=False,
nnodes=1,
dp_size=1,
load_balance_method="round_robin",
node_rank=0,
)
port_args = SimpleNamespace(instance_id="test_shm_factory")
writer = create_load_snapshot_writer(
server_args, port_args, dp_size=1, dp_rank=0
)
self.assertIsInstance(writer, ShmLoadSnapshotWriter)
reader = create_load_snapshot_reader(server_args, port_args, caller="tokenizer")
self.assertIsInstance(reader, ShmLoadSnapshotReader)
reader.close()
writer.close()
from sglang.srt.managers.load_snapshot import shm_path_for
path = shm_path_for("test_shm_factory")
if os.path.exists(path):
os.unlink(path)
def test_zmq_mode_via_env(self):
server_args = SimpleNamespace(
enable_dp_attention=False,
nnodes=1,
dp_size=1,
load_balance_method="round_robin",
node_rank=0,
)
port_args = SimpleNamespace(instance_id="test_zmq_factory")
os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"] = "1"
try:
writer = create_load_snapshot_writer(
server_args, port_args, dp_size=1, dp_rank=0
)
self.assertIsInstance(writer, ZmqLoadSnapshotWriter)
reader = create_load_snapshot_reader(
server_args, port_args, caller="tokenizer"
)
self.assertIsInstance(reader, ZmqShmLoadSnapshotReader)
reader.close()
writer.close()
finally:
del os.environ["SGLANG_LOAD_SNAPSHOT_USE_ZMQ"]
def test_should_use_zmq_multinode_dp_attention(self):
args = SimpleNamespace(enable_dp_attention=True, nnodes=2)
self.assertTrue(should_use_zmq(args))
def test_should_use_zmq_single_node(self):
args = SimpleNamespace(enable_dp_attention=False, nnodes=1)
self.assertFalse(should_use_zmq(args))
def test_should_use_zmq_dp_attention_single_node(self):
args = SimpleNamespace(enable_dp_attention=True, nnodes=1)
self.assertFalse(should_use_zmq(args))
class TestZmqAddr(CustomTestCase):
def test_ipc_for_single_node(self):
port_args = SimpleNamespace(instance_id="myinstance")
addr = _zmq_addr_for(port_args)
self.assertTrue(addr.startswith("ipc://"))
self.assertIn("myinstance", addr)
def test_tcp_from_port_args(self):
from sglang.srt.utils.network import NetworkAddress
port_args = SimpleNamespace(
instance_id="myinstance",
load_collector_ipc_name=NetworkAddress("10.0.0.1", 29506).to_tcp(),
)
addr = _zmq_addr_for(port_args)
self.assertTrue(addr.startswith("tcp://"))
self.assertIn("10.0.0.1", addr)
class TestEndToEndZmqSimulation(CustomTestCase):
"""Simulate multi-node DP attention on single machine using IPC."""
def test_full_flow_dp_size_2(self):
shm_path = _temp_path()
addr = _ipc_addr()
dp_size = 2
reader = ZmqShmLoadSnapshotReader(addr, shm_path, dp_size)
writers = []
for rank in range(dp_size):
w = ZmqLoadSnapshotWriter(addr, dp_size, dp_rank=rank)
writers.append(w)
try:
_warmup_zmq(writers, reader)
for rank, w in enumerate(writers):
w.write(
LoadSnapshot(
dp_rank=rank,
timestamp=1.0,
num_running_reqs=10 + rank,
num_waiting_reqs=5 + rank,
num_total_tokens=100 + rank * 50,
)
)
time.sleep(0.05)
loads = reader.read_all()
self.assertEqual(len(loads), dp_size)
self.assertEqual(loads[0].num_running_reqs, 10)
self.assertEqual(loads[1].num_running_reqs, 11)
self.assertEqual(loads[0].num_total_tokens, 100)
self.assertEqual(loads[1].num_total_tokens, 150)
for rank, w in enumerate(writers):
w.write(
LoadSnapshot(
dp_rank=rank,
timestamp=2.0,
num_running_reqs=20 + rank,
num_waiting_reqs=0,
num_total_tokens=200 + rank * 50,
)
)
time.sleep(0.05)
loads = reader.read_all()
self.assertEqual(loads[0].num_running_reqs, 20)
self.assertEqual(loads[1].num_running_reqs, 21)
self.assertEqual(loads[0].num_total_tokens, 200)
self.assertEqual(loads[1].num_total_tokens, 250)
finally:
for w in writers:
w.close()
reader.close()
if os.path.exists(shm_path):
os.unlink(shm_path)
if __name__ == "__main__":
unittest.main()