[gRPC] Add GetLoads RPC for comprehensive load metrics (#17087)

This commit is contained in:
Simo Lin
2026-01-14 12:05:10 -08:00
committed by GitHub
parent 5b1215d9da
commit f091858304
4 changed files with 439 additions and 0 deletions
@@ -23,6 +23,8 @@ from sglang.srt.managers.io_struct import (
AbortReq,
BatchEmbeddingOutput,
BatchTokenIDOutput,
GetLoadsReqInput,
GetLoadsReqOutput,
HealthCheckOutput,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
@@ -34,6 +36,68 @@ from sglang.utils import get_exception_traceback
logger = logging.getLogger(__name__)
class _GrpcCommunicator:
"""
Communicator for request/response patterns with scheduler.
Thread-safe and handles the async request/response cycle with proper
timeout handling to prevent hangs if the scheduler becomes unresponsive.
"""
DEFAULT_TIMEOUT = 30.0 # seconds
def __init__(self, sender: zmq.Socket, fan_out: int = 1):
self._sender = sender
self._fan_out = fan_out
self._result_event: Optional[asyncio.Event] = None
self._result_values: Optional[List[Any]] = None
self._lock = asyncio.Lock()
async def __call__(self, obj, timeout: float = DEFAULT_TIMEOUT) -> List[Any]:
"""
Send request and wait for response(s).
Args:
obj: Request object to send to scheduler
timeout: Maximum time to wait for response (seconds)
Returns:
List of response objects from scheduler(s)
Raises:
asyncio.TimeoutError: If no response within timeout
"""
async with self._lock:
# Initialize state BEFORE sending to avoid race condition
self._result_event = asyncio.Event()
self._result_values = []
# Send request to scheduler
if obj:
self._sender.send_pyobj(obj)
try:
# Wait for response(s) with timeout
await asyncio.wait_for(self._result_event.wait(), timeout=timeout)
return self._result_values
finally:
# Always clean up state
self._result_event = None
self._result_values = None
def handle_recv(self, recv_obj: Any):
"""
Handle received response from scheduler.
Called by handle_loop when a matching response type is received.
Safe to call even if no request is pending (will be ignored).
"""
if self._result_values is not None and self._result_event is not None:
self._result_values.append(recv_obj)
if len(self._result_values) >= self._fan_out:
self._result_event.set()
class GrpcSignalHandler:
"""Minimal signal handler for gRPC server - delegates real crash handling to scheduler."""
@@ -154,6 +218,12 @@ class GrpcRequestManager:
# Bootstrap server (passed from serve_grpc, not started here)
self.bootstrap_server = bootstrap_server
# Communicators for request/response patterns with scheduler
# Note: These must be initialized after send_to_scheduler socket is created
self.get_loads_communicator = _GrpcCommunicator(
self.send_to_scheduler, fan_out=server_args.dp_size
)
logger.info(
f"GrpcRequestManager initialized with ZMQ IPC: "
f"recv={port_args.detokenizer_ipc_name}, "
@@ -462,6 +532,9 @@ class GrpcRequestManager:
await self._handle_health_check_output(recv_obj)
elif isinstance(recv_obj, AbortReq):
await self._handle_abort_req(recv_obj)
elif isinstance(recv_obj, GetLoadsReqOutput):
# Route to communicator for request/response pattern
self.get_loads_communicator.handle_recv(recv_obj)
else:
logger.warning(f"Unknown output type: {type(recv_obj)}")
@@ -872,6 +945,31 @@ class GrpcRequestManager:
"last_receive_time": self.last_receive_tstamp,
}
async def get_loads(
self, include: List[str], dp_rank: Optional[int] = None
) -> List[GetLoadsReqOutput]:
"""
Get comprehensive load metrics from the scheduler.
This method uses the communicator pattern to send GetLoadsReqInput to the
scheduler and wait for GetLoadsReqOutput responses.
Args:
include: List of metric sections to include (core, memory, spec, lora, disagg, queues, all)
dp_rank: Optional DP rank filter (None for all ranks)
Returns:
List of GetLoadsReqOutput objects, one per scheduler/DP rank
"""
req = GetLoadsReqInput(include=include, dp_rank=dp_rank)
results = await self.get_loads_communicator(req)
# Filter by dp_rank if specified
if dp_rank is not None:
results = [r for r in results if r.dp_rank == dp_rank]
return results
def auto_create_handle_loop(self):
"""Automatically create and start the handle_loop task, matching TokenizerManager pattern."""
if self.no_create_loop:
@@ -26,6 +26,9 @@ service SglangScheduler {
// Get server information
rpc GetServerInfo(GetServerInfoRequest) returns (GetServerInfoResponse);
// Get comprehensive load metrics
rpc GetLoads(GetLoadsRequest) returns (GetLoadsResponse);
}
// =====================
@@ -464,3 +467,100 @@ message GetServerInfoResponse {
// bidirectional communicator infrastructure not available in gRPC.
// Use HTTP /get_server_info if scheduler internal state is needed.
}
// =====================
// Load Metrics (v1/loads)
// =====================
message GetLoadsRequest {
// Optional: filter to specific DP rank
optional int32 dp_rank = 1;
// Sections to include: core, memory, spec, lora, disagg, queues, all
repeated string include = 2;
}
message GetLoadsResponse {
// ISO 8601 timestamp
string timestamp = 1;
// SGLang version
string version = 2;
// Number of DP ranks
int32 dp_rank_count = 3;
// Per-DP-rank load metrics
repeated SchedulerLoad loads = 4;
// Aggregate metrics across all DP ranks
AggregateMetrics aggregate = 5;
}
message SchedulerLoad {
int32 dp_rank = 1;
// Core metrics (always included)
int32 num_running_reqs = 2;
int32 num_waiting_reqs = 3;
int32 num_total_reqs = 4;
int32 num_used_tokens = 5;
int32 max_total_num_tokens = 6;
double token_usage = 7;
double gen_throughput = 8;
double cache_hit_rate = 9;
double utilization = 10;
int32 max_running_requests = 11;
// Optional sections
optional MemoryMetrics memory = 12;
optional SpeculativeMetrics speculative = 13;
optional LoRAMetrics lora = 14;
optional DisaggregationMetrics disaggregation = 15;
optional QueueMetrics queues = 16;
}
message MemoryMetrics {
double weight_gb = 1;
double kv_cache_gb = 2;
double graph_gb = 3;
int32 token_capacity = 4;
}
message SpeculativeMetrics {
double accept_length = 1;
double accept_rate = 2;
}
message LoRAMetrics {
int32 slots_used = 1;
int32 slots_total = 2;
double utilization = 3;
}
message DisaggregationMetrics {
string mode = 1; // "prefill", "decode", or "null"
int32 prefill_prealloc_queue_reqs = 2;
int32 prefill_inflight_queue_reqs = 3;
int32 decode_prealloc_queue_reqs = 4;
int32 decode_transfer_queue_reqs = 5;
int32 decode_retracted_queue_reqs = 6;
double kv_transfer_speed_gb_s = 7;
double kv_transfer_latency_ms = 8;
}
message QueueMetrics {
int32 waiting = 1;
int32 grammar = 2;
int32 paused = 3;
int32 retracted = 4;
}
message AggregateMetrics {
int32 total_running_reqs = 1;
int32 total_waiting_reqs = 2;
int32 total_reqs = 3;
double avg_token_usage = 4;
double avg_throughput = 5;
double avg_utilization = 6;
}