[Cleanup] Style and type annotation improvements extracted from #28688 (#29224)

This commit is contained in:
Lianmin Zheng
2026-06-25 12:51:15 -07:00
committed by GitHub
parent 3344b73c80
commit 118d6b2e5e
5 changed files with 217 additions and 247 deletions
+6 -27
View File
@@ -1416,17 +1416,8 @@ async def load_lora_adapter(
):
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.load_lora_adapter(obj, request)
if result.success:
return ORJSONResponse(
result,
status_code=HTTPStatus.OK,
)
else:
return ORJSONResponse(
result,
status_code=HTTPStatus.BAD_REQUEST,
)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
@app.api_route("/load_lora_adapter_from_tensors", methods=["POST"])
@@ -1437,11 +1428,8 @@ async def load_lora_adapter_from_tensors(
result = await _global_state.tokenizer_manager.load_lora_adapter_from_tensors(
obj, request
)
if result.success:
return ORJSONResponse(result, status_code=HTTPStatus.OK)
else:
return ORJSONResponse(result, status_code=HTTPStatus.BAD_REQUEST)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
@app.api_route("/unload_lora_adapter", methods=["POST"])
@@ -1451,17 +1439,8 @@ async def unload_lora_adapter(
):
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.unload_lora_adapter(obj, request)
if result.success:
return ORJSONResponse(
result,
status_code=HTTPStatus.OK,
)
else:
return ORJSONResponse(
result,
status_code=HTTPStatus.BAD_REQUEST,
)
status_code = HTTPStatus.OK if result.success else HTTPStatus.BAD_REQUEST
return ORJSONResponse(result, status_code=status_code)
@app.api_route("/open_session", methods=["GET", "POST"])
@@ -382,7 +382,7 @@ class DataParallelController:
connected_clients = 0
while connected_clients < expected_clients:
# Wait for client handshake
client_rank = rep_socket.recv().decode()
client_rank = sock_recv(rep_socket).decode()
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
@@ -411,7 +411,7 @@ class DataParallelController:
while True:
# Wait for client handshake
try:
client_rank = rep_socket.recv().decode()
client_rank = sock_recv(rep_socket).decode()
except Exception:
logger.exception(
"Failed to recv/decode handshake in reply thread; continue"
@@ -433,7 +433,7 @@ class DataParallelController:
try:
# Send handshake with our node rank
req_socket.send(str(node_rank).encode())
sock_send(req_socket, str(node_rank).encode())
# Receive worker ports
worker_ports = sock_recv(req_socket)
+183 -204
View File
@@ -46,7 +46,7 @@ from pydantic import PlainValidator
from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.managers.schedule_batch import Modality, MultimodalInputs
from sglang.srt.multimodal.mm_utils import has_valid_data
from sglang.srt.observability.req_time_stats import (
APIServerReqTimeStats,
@@ -84,18 +84,34 @@ class BaseBatchReq:
# Parameters for a session
@dataclass
class SessionParams:
# The session identifier. Used by the scheduler to look up or create the
# Session object that groups all requests in a multi-turn conversation.
id: Optional[str] = None
# A request identifier *within* the session. In non-streaming sessions the
# session maintains a tree of request nodes keyed by rid; this field selects
# which node to continue from (append) or replace. When None the default
# branch point is used (latest node for streaming, all nodes cleared on
# replace).
rid: Optional[str] = None
# Token-level insertion point. When set, the new request's tokens are
# spliced into the accumulated context at this position instead of being
# appended at the end (i.e. ``context[:offset] + new_tokens``).
offset: Optional[int] = None
# When True, the request node identified by ``rid`` (or all nodes if
# ``rid`` is None) is aborted and its children are cleared before the new
# request is inserted. Not supported in streaming sessions.
replace: Optional[bool] = None
# When True, the previous request's generated output tokens are excluded
# from the accumulated context so the new turn sees only the original input.
# Not supported in streaming sessions.
drop_previous_output: Optional[bool] = None
# Type definitions for multimodal input data
# Individual data item types for each modality
ImageDataInputItem = Union[str, Dict, ImageData, Image]
AudioDataInputItem = Union[str, Dict]
VideoDataInputItem = Union[str, Dict, VideoData]
ImageDataInputItem = Union[str, bytes, Dict[str, Any], ImageData, Image]
AudioDataInputItem = Union[str, bytes, Dict[str, Any]]
VideoDataInputItem = Union[str, bytes, Dict[str, Any], VideoData]
# Union type for any multimodal data item
MultimodalDataInputItem = Union[
ImageDataInputItem, VideoDataInputItem, AudioDataInputItem
@@ -107,19 +123,12 @@ MultimodalDataInputFormat = Union[
MultimodalDataInputItem,
]
# Serialized form of BaseFinishReason.to_json() — all values are primitives.
FinishReasonDict = Dict[str, Optional[Union[str, int, List[int]]]]
CachedTokensDetails = Dict[str, Union[int, str]]
@dataclass
class GenerateReqInput:
# Request ID(s). If omitted, generated during normalization. For batch
# requests, a string is expanded to per-item IDs using it as a prefix.
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
# Internal IPC endpoint of the HTTP/tokenizer worker that owns this request.
# Used to route outputs back in multi-tokenizer mode.
http_worker_ipc: Optional[str] = field(default=None, kw_only=True)
# The input prompt. It can be a single prompt or a batch of prompts.
text: Optional[Union[List[str], str]] = None
# The token ids for text.
@@ -141,8 +150,10 @@ class GenerateReqInput:
video_data: Optional[MultimodalDataInputFormat] = None
# The audio input. Like image data, it can be a file name, a url, or base64 encoded string.
audio_data: Optional[MultimodalDataInputFormat] = None
# Optional per-image hashes the caller has already computed (hex strings,
# one per image in `image_data`). When supplied, each MultimodalDataItem's
# Optional per-image hashes the caller has already computed (hex strings).
# Single request: one hash per image. Batch request: either one hash per
# request when each request has one image, or one list of hashes per request.
# When supplied, each MultimodalDataItem's
# `hash` is initialised from this list and `set_pad_value` skips the
# internal `hash_feature()` recompute, so the resulting `pad_value` is
# deterministic from the caller's hash. Intended for external KV routers
@@ -153,7 +164,7 @@ class GenerateReqInput:
# Whether to extract and process audio from video inputs.
use_audio_in_video: bool = False
# The sampling_params. See descriptions below.
sampling_params: Optional[Union[List[Dict], Dict]] = None
sampling_params: Optional[Union[List[Dict[str, Any]], Dict[str, Any]]] = None
# Whether to return logprobs.
return_logprob: Optional[Union[List[bool], bool]] = None
# If return logprobs, the start location in the prompt for returning logprobs.
@@ -173,21 +184,21 @@ class GenerateReqInput:
return_hidden_states: Union[List[bool], bool] = False
# Whether to return captured routed experts
return_routed_experts: bool = False
return_indexer_topk: bool = False
# Absolute start position for returned routings; response covers
# `[routed_experts_start_len, seqlen - 1)`. Must be in [0, prompt_tokens].
# 0 = full sequence.
routed_experts_start_len: int = 0
return_indexer_topk: bool = False
# The modalities of the image data [image, multi-images, video]
modalities: Optional[List[str]] = None
# Session info for continual prompting
session_params: Optional[Union[List[Dict], Dict]] = None
session_params: Optional[Dict[str, Any]] = None
# The path to the LoRA adaptors
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
lora_path: Optional[Union[List[Optional[str]], str]] = None
# The uid of LoRA adaptors, should be initialized by tokenizer manager
lora_id: Optional[Union[List[Optional[str]], Optional[str]]] = None
lora_id: Optional[Union[List[Optional[str]], str]] = None
# Custom logit processor for advanced sampling control. Must be a serialized instance
# of `CustomLogitProcessor` in python/sglang/srt/sampling/custom_logit_processor.py
@@ -199,21 +210,23 @@ class GenerateReqInput:
positional_embed_overrides: Any = None
# For disaggregated inference
bootstrap_host: Optional[Union[List[str], str]] = None
bootstrap_host: Optional[Union[List[Optional[str]], str]] = None
bootstrap_port: Optional[Union[List[Optional[int]], int]] = None
bootstrap_room: Optional[Union[List[int], int]] = None
bootstrap_pair_key: Optional[Union[List[str], str]] = None
bootstrap_room: Optional[Union[List[Optional[int]], int]] = None
bootstrap_pair_key: Optional[Union[List[Optional[str]], str]] = None
decode_tp_size: Optional[Union[List[Optional[int]], int]] = None
# For DP routing — external router assigns a specific DP worker
routed_dp_rank: Optional[int] = None
# For PD disagg — hint telling decode which prefill DP worker has the KV cache
disagg_prefill_dp_rank: Optional[int] = None
# Routing key for routing-key schedule policy
routing_key: Optional[str] = None
# Conversation id used for tracking requests
conversation_id: Optional[str] = None
# Internal IPC endpoint of the HTTP/tokenizer worker that owns this request.
# Used to route outputs back in multi-tokenizer mode.
http_worker_ipc: Optional[str] = field(default=None, kw_only=True)
# For background responses (OpenAI responses API)
background: bool = False
@@ -222,13 +235,11 @@ class GenerateReqInput:
# Priority for the request
priority: Optional[int] = None
# Extra cache key for classifying the request (e.g. cache_salt)
extra_key: Optional[Union[List[str], str]] = None
# Whether to disallow logging for this request (e.g. due to ZDR)
no_logs: bool = False
# For custom metric labels
custom_labels: Optional[Dict[str, str]] = None
@@ -240,13 +251,13 @@ class GenerateReqInput:
return_prompt_token_ids: bool = False
# Propagates trace context via Engine.generate/async_generate
external_trace_header: Optional[Dict] = None
external_trace_header: Optional[Dict[str, Any]] = None
received_time: Optional[float] = None
# For EPD-disaggregated inference
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
mm_data_mooncake: Optional[List[Any]] = None
# Snapshot of encoder URLs at the time tokenizer-side computed
# ``num_items_assigned``.
encoder_urls: Optional[List[str]] = None
@@ -629,13 +640,13 @@ class GenerateReqInput:
elif isinstance(self.bootstrap_pair_key, list):
self.bootstrap_pair_key = self.bootstrap_pair_key * self.parallel_sample_num
def _validate_session_params(self):
"""Validate that session parameters are properly formatted."""
if self.session_params is not None:
if not isinstance(self.session_params, dict) and not isinstance(
self.session_params[0], dict
):
raise ValueError("Session params must be a dict or a list of dicts.")
# Normalize decode_tp_size
if self.decode_tp_size is None:
self.decode_tp_size = [None] * num
elif not isinstance(self.decode_tp_size, list):
self.decode_tp_size = [self.decode_tp_size] * num
elif isinstance(self.decode_tp_size, list):
self.decode_tp_size = self.decode_tp_size * self.parallel_sample_num
def _get_positional_embed_overrides_item(
self, i: int
@@ -654,17 +665,16 @@ class GenerateReqInput:
if i in cache:
return cache[i]
sub = GenerateReqInput(
rid=self.rid[i],
text=self.text[i] if self.text is not None else None,
input_ids=self.input_ids[i] if self.input_ids is not None else None,
input_embeds=(
self.input_embeds[i] if self.input_embeds is not None else None
),
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
image_data=self.image_data[i],
video_data=self.video_data[i],
audio_data=self.audio_data[i],
sampling_params=self.sampling_params[i],
rid=self.rid[i],
return_logprob=self.return_logprob[i],
logprob_start_len=self.logprob_start_len[i],
top_logprobs_num=self.top_logprobs_num[i],
@@ -689,7 +699,8 @@ class GenerateReqInput:
if self.custom_logit_processor is not None
else None
),
# if `__getitem__` is called, the bootstrap_host, bootstrap_port, bootstrap_room must be a list
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
# If `__getitem__` is called, these bootstrap fields must be lists.
bootstrap_host=(
self.bootstrap_host[i] if self.bootstrap_host is not None else None
),
@@ -710,6 +721,7 @@ class GenerateReqInput:
routed_dp_rank=self.routed_dp_rank,
disagg_prefill_dp_rank=self.disagg_prefill_dp_rank,
conversation_id=self.conversation_id,
http_worker_ipc=self.http_worker_ipc,
priority=self.priority,
extra_key=self.extra_key[i] if self.extra_key is not None else None,
no_logs=self.no_logs,
@@ -718,7 +730,6 @@ class GenerateReqInput:
return_entropy=self.return_entropy,
return_prompt_token_ids=self.return_prompt_token_ids,
external_trace_header=self.external_trace_header,
http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
@@ -733,13 +744,13 @@ class GenerateReqInput:
@dataclass
class TokenizedGenerateReqInput(BaseReq):
# The input text
input_text: str
input_text: Optional[Union[str, List[Union[str, List[str]]]]]
# The input token ids
input_ids: Optional[array[int]]
input_ids: Optional[array] # Optional[array[int]]
# The input embeds
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]]
input_embeds: Optional[List[List[float]]]
# The multimodal inputs
mm_inputs: object
mm_inputs: Optional[MultimodalInputs]
token_type_ids: Optional[List[int]]
# The sampling parameters
sampling_params: SamplingParams
@@ -759,9 +770,9 @@ class TokenizedGenerateReqInput(BaseReq):
# Whether to return captured routed experts
return_routed_experts: bool = False
return_indexer_topk: bool = False
# See GenerateReqInput.routed_experts_start_len.
routed_experts_start_len: int = 0
return_indexer_topk: bool = False
# Session info for continual prompting
session_params: Optional[SessionParams] = None
@@ -809,7 +820,7 @@ class TokenizedGenerateReqInput(BaseReq):
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
mm_data_mooncake: Optional[List[Any]] = None
# Encoder URL snapshot frozen at tokenizer-side dispatch time so that
# encoder_idx assignments stay consistent in the scheduler subprocess.
# Internal IPC only.
@@ -842,9 +853,6 @@ class EmbeddingReqInput:
# Request ID(s). If omitted, generated during normalization. For batch
# requests, a string is expanded to per-item IDs using it as a prefix.
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
# Internal IPC endpoint of the HTTP/tokenizer worker that owns this request.
# Used to route outputs back in multi-tokenizer mode.
http_worker_ipc: Optional[str] = field(default=None, kw_only=True)
# The input prompt. It can be a single prompt or a batch of prompts.
text: Optional[Union[List[List[str]], List[str], str]] = None
# The token ids for text; one can either specify text or input_ids.
@@ -871,26 +879,30 @@ class EmbeddingReqInput:
# Runtime type: Optional[List[Optional[List[torch.Tensor]]]]
# Typed as Any to avoid Pydantic/FastAPI schema errors (contains torch.Tensor).
embed_overrides: Any = None
# The path to the LoRA adaptors
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None
# The uid of LoRA adaptors, should be initialized by tokenizer manager
lora_id: Optional[Union[List[Optional[str]], Optional[str]]] = None
# Resolved embedding overrides with positions (set by tokenizer manager or score mixin).
# Runtime type: Optional[Union[PositionalEmbeds, List[Optional[PositionalEmbeds]]]]
positional_embed_overrides: Any = None
# Dummy sampling params for compatibility
sampling_params: Optional[Union[List[Dict], Dict]] = None
sampling_params: Optional[Union[List[Dict[str, Any]], Dict[str, Any]]] = None
# Whether to log metrics for this request (e.g. health_generate calls do not log metrics)
log_metrics: bool = True
# The modalities of the image data [image, multi-images, video]
modalities: Optional[List[str]] = None
# For cross-encoder requests
is_cross_encoder_request: bool = False
# The path to the LoRA adaptors
lora_path: Optional[Union[List[Optional[str]], str]] = None
# The uid of LoRA adaptors, should be initialized by tokenizer manager
lora_id: Optional[Union[List[Optional[str]], str]] = None
# Resolved embedding overrides with positions (set by tokenizer manager or score mixin).
# Runtime type: Optional[Union[PositionalEmbeds, List[Optional[PositionalEmbeds]]]]
positional_embed_overrides: Any = None
# Routing key for routing-key schedule policy
routing_key: Optional[str] = None
# Internal IPC endpoint of the HTTP/tokenizer worker that owns this request.
# Used to route outputs back in multi-tokenizer mode.
http_worker_ipc: Optional[str] = field(default=None, kw_only=True)
# For background responses (OpenAI responses API)
background: bool = False
# Priority for the request
priority: Optional[int] = None
@@ -902,7 +914,7 @@ class EmbeddingReqInput:
return_prompt_token_ids: bool = False
# Propagates trace context via Engine.encode/async_encode
external_trace_header: Optional[Dict] = None
external_trace_header: Optional[Dict[str, Any]] = None
received_time: Optional[float] = None
# Pre-computed delimiter indices for multi-item scoring.
@@ -1019,13 +1031,13 @@ class EmbeddingReqInput:
if self.is_cross_encoder_request:
sub = EmbeddingReqInput(
text=[self.text[i]] if self.text is not None else None,
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
sampling_params=self.sampling_params[i],
rid=self.rid[i],
text=[self.text[i]] if self.text is not None else None,
sampling_params=self.sampling_params[i],
is_cross_encoder_request=True,
lora_path=self.lora_path[i] if self.lora_path is not None else None,
lora_id=self.lora_id[i] if self.lora_id is not None else None,
is_cross_encoder_request=True,
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
http_worker_ipc=self.http_worker_ipc,
return_pooled_hidden_states=self.return_pooled_hidden_states,
return_prompt_token_ids=self.return_prompt_token_ids,
@@ -1037,28 +1049,28 @@ class EmbeddingReqInput:
)
else:
sub = EmbeddingReqInput(
rid=self.rid[i],
text=self.text[i] if self.text is not None else None,
input_ids=self.input_ids[i] if self.input_ids is not None else None,
image_data=self.image_data[i] if self.image_data is not None else None,
video_data=self.video_data[i] if self.video_data is not None else None,
audio_data=self.audio_data[i] if self.audio_data is not None else None,
embed_override_token_id=self.embed_override_token_id,
embed_overrides=(
self.embed_overrides[i]
if self.embed_overrides is not None
else None
),
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
image_data=self.image_data[i] if self.image_data is not None else None,
audio_data=self.audio_data[i] if self.audio_data is not None else None,
video_data=self.video_data[i] if self.video_data is not None else None,
sampling_params=self.sampling_params[i],
rid=self.rid[i],
lora_path=self.lora_path[i] if self.lora_path is not None else None,
lora_id=self.lora_id[i] if self.lora_id is not None else None,
external_trace_header=self.external_trace_header,
dimensions=self.dimensions,
positional_embed_overrides=self._get_positional_embed_overrides_item(i),
http_worker_ipc=self.http_worker_ipc,
received_time=self.received_time,
dimensions=self.dimensions,
return_pooled_hidden_states=self.return_pooled_hidden_states,
return_prompt_token_ids=self.return_prompt_token_ids,
external_trace_header=self.external_trace_header,
received_time=self.received_time,
multi_item_delimiter_indices=(
self.multi_item_delimiter_indices[i]
if self.multi_item_delimiter_indices is not None
@@ -1072,11 +1084,11 @@ class EmbeddingReqInput:
@dataclass
class TokenizedEmbeddingReqInput(BaseReq):
# The input text
input_text: str
input_text: Optional[Union[str, List[Union[str, List[str]]]]]
# The input token ids
input_ids: array[int]
input_ids: Optional[array] # array[int]
# The multimodal inputs
mm_inputs: object
mm_inputs: Optional[MultimodalInputs]
# The token type ids
token_type_ids: Optional[List[int]]
# Dummy sampling params for compatibility
@@ -1091,10 +1103,10 @@ class TokenizedEmbeddingReqInput(BaseReq):
priority: Optional[int] = None
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
dimensions: Optional[int] = None
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False
# Pre-computed delimiter indices for multi-item scoring
multi_item_delimiter_indices: Optional[List[int]] = None
# For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@@ -1115,12 +1127,17 @@ class BatchTokenizedEmbeddingReqInput(BaseBatchReq):
return iter(self.batch)
TokenLogprobValues = Optional[List[List[Optional[float]]]]
TokenLogprobIndices = Optional[List[List[Optional[int]]]]
TokenLogprobValues = Optional[List[Optional[List[Optional[float]]]]]
TokenLogprobIndices = Optional[List[Optional[List[Optional[int]]]]]
TopLogprobValues = Optional[List[Optional[List[Optional[List[float]]]]]]
TopLogprobIndices = Optional[List[Optional[List[Optional[List[int]]]]]]
TokenIdsLogprobValues = Optional[List[Optional[List[Optional[List[float]]]]]]
TokenIdsLogprobIndices = Optional[List[Optional[List[Optional[List[int]]]]]]
HiddenStateChunk = List[Optional[Union[float, List[float]]]]
OutputHiddenStates = Optional[List[Optional[List[HiddenStateChunk]]]]
CachedTokensDetails = Dict[str, Union[int, str]]
# Serialized form of BaseFinishReason.to_json() — all values are primitives.
FinishReasonDict = Dict[str, Optional[Union[str, int, List[int]]]]
@dataclass
@@ -1129,10 +1146,10 @@ class BatchTokenIDOutput(BaseBatchReq):
finished_reasons: List[Optional[FinishReasonDict]]
# For incremental decoding
decoded_texts: List[str]
decode_ids: List[array[int]]
decode_ids: List[array] # List[array[int]]
read_offsets: List[int]
# Only used when `--skip-tokenizer-init` is on
output_ids: Optional[List[array[int]]]
output_ids: Optional[List[array]] # Optional[List[array[int]]]
# Detokenization configs
skip_special_tokens: List[bool]
spaces_between_special_tokens: List[bool]
@@ -1153,10 +1170,10 @@ class BatchTokenIDOutput(BaseBatchReq):
input_top_logprobs_idx: TopLogprobIndices
output_top_logprobs_val: TopLogprobValues
output_top_logprobs_idx: TopLogprobIndices
input_token_ids_logprobs_val: TokenLogprobValues
input_token_ids_logprobs_idx: TokenLogprobIndices
output_token_ids_logprobs_val: TokenLogprobValues
output_token_ids_logprobs_idx: TokenLogprobIndices
input_token_ids_logprobs_val: TokenIdsLogprobValues
input_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_ids_logprobs_val: TokenIdsLogprobValues
output_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_entropy_val: Optional[List[Optional[float]]]
# Hidden states
@@ -1229,10 +1246,10 @@ class BatchStrOutput(BaseBatchReq):
input_top_logprobs_idx: TopLogprobIndices
output_top_logprobs_val: TopLogprobValues
output_top_logprobs_idx: TopLogprobIndices
input_token_ids_logprobs_val: TokenLogprobValues
input_token_ids_logprobs_idx: TokenLogprobIndices
output_token_ids_logprobs_val: TokenLogprobValues
output_token_ids_logprobs_idx: TokenLogprobIndices
input_token_ids_logprobs_val: TokenIdsLogprobValues
input_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_ids_logprobs_val: TokenIdsLogprobValues
output_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_entropy_val: Optional[List[Optional[float]]]
# Hidden states
@@ -1285,7 +1302,7 @@ class BatchEmbeddingOutput(BaseBatchReq):
# The finish reason
finished_reasons: List[Optional[FinishReasonDict]]
# The output embedding
embeddings: Union[List[List[float]], List[Dict[int, float]]]
embeddings: List[Union[List[Union[float, List[float]]], Dict[int, float], float]]
# Token counts
prompt_tokens: List[int]
cached_tokens: List[int]
@@ -1302,10 +1319,10 @@ class BatchEmbeddingOutput(BaseBatchReq):
time_stats: Optional[List[SchedulerReqTimeStats]] = None
# Optional pooled hidden states (pre-head transformer output).
# Sent as a single stacked tensor to minimize pickle overhead.
pooled_hidden_states: Optional[
Union[List[Optional[torch.Tensor]], torch.Tensor]
] = None
# Two IPC formats, disambiguated by len vs len(rids):
# Stacked: [stacked_tensor(N, ...)] — len 1, reduces pickle overhead
# Non-stacked: [t0, t1, ..., tN] — len N, when shapes differ or None entries exist
pooled_hidden_states: Optional[List[Optional[torch.Tensor]]] = None
@dataclass
@@ -1480,7 +1497,7 @@ class UpdateWeightFromDiskReqOutput(BaseReq):
success: bool
message: str
# Number of paused requests during weight sync.
num_paused_requests: Optional[int] = 0
num_paused_requests: int = 0
@dataclass
@@ -1516,6 +1533,10 @@ class UpdateWeightsFromTensorReqInput(BaseReq):
- Data is structured in JSON for easy transmission over HTTP
"""
# Accepts both base64 str (from HTTP/JSON, which has no bytes type) and
# raw bytes (from the Python Engine API / MultiprocessingSerializer).
# Normalized to List[bytes] by normalize_serialized_named_tensor_payloads
# in tokenizer_control_mixin before forwarding over scheduler IPC.
serialized_named_tensors: List[Union[str, bytes]]
# Optional format specification for loading
load_format: Optional[str] = None
@@ -1613,7 +1634,7 @@ class InitWeightsUpdateGroupReqInput(BaseReq):
# The master address
master_address: str
# The master port
master_port: Union[int, str]
master_port: int
# The rank offset
rank_offset: int
# The world size
@@ -1657,7 +1678,7 @@ class GetWeightsByNameReqInput(BaseReq):
@dataclass
class GetWeightsByNameReqOutput(BaseReq):
parameter: list
parameter: Optional[List[Any]]
@dataclass
@@ -1693,7 +1714,7 @@ class CheckWeightsReqInput(BaseReq):
class CheckWeightsReqOutput(BaseReq):
success: bool
message: str
payload: Optional[Dict] = None
payload: Optional[Dict[str, Any]] = None
@dataclass
@@ -1732,7 +1753,7 @@ class GetInternalStateReq(BaseReq):
@dataclass
class GetInternalStateReqOutput(BaseReq):
internal_state: Dict[Any, Any]
internal_state: Dict[str, Any]
@dataclass
@@ -1853,13 +1874,13 @@ class ExpertDistributionReqOutput(BaseReq):
class Function:
description: Optional[str] = None
name: Optional[str] = None
parameters: Optional[Any] = None
parameters: Optional[Dict[str, Any]] = None
@dataclass
class Tool:
function: Function
type: Optional[str] = "function"
type: str = "function"
@dataclass
@@ -1882,14 +1903,14 @@ class SeparateReasoningReqInput(BaseReq):
@dataclass
class VertexGenerateReqInput(BaseReq):
instances: List[dict]
parameters: Optional[dict] = None
instances: List[Dict[str, Any]]
parameters: Optional[Dict[str, Any]] = None
@dataclass
class RpcReqInput(BaseReq):
method: str
parameters: Optional[Dict] = None
parameters: Optional[Dict[str, Any]] = None
@dataclass
@@ -1955,7 +1976,7 @@ class LoadLoRAAdapterFromTensorsReqInput(BaseReq):
class LoRAUpdateOutput(BaseReq):
success: bool
error_message: Optional[str] = None
loaded_adapters: Optional[Dict[str, LoRARef]] = None
loaded_adapters: Optional[Dict[str, Union[str, LoRARef]]] = None
LoadLoRAAdapterReqOutput = UnloadLoRAAdapterReqOutput = (
@@ -1977,84 +1998,51 @@ class BlockReqInput(BaseReq):
class MemoryMetrics:
"""Memory breakdown metrics."""
weight_gb: float = field(
metadata={"metric": ("gauge", "Model weight memory in GB")}
)
kv_cache_gb: float = field(metadata={"metric": ("gauge", "KV cache memory in GB")})
graph_gb: float = field(metadata={"metric": ("gauge", "CUDA graph memory in GB")})
token_capacity: int = field(
metadata={"metric": ("gauge", "Max tokens in KV cache")}
)
weight_gb: float
kv_cache_gb: float
graph_gb: float
token_capacity: int
@dataclass
class SpeculativeMetrics:
"""Speculative decoding metrics."""
accept_length: float = field(
metadata={
"metric": (
"gauge",
"Mean acceptance length (accepted drafts + bonus token per forward)",
)
}
)
accept_rate: float = field(
metadata={"metric": ("gauge", "Speculative acceptance rate")}
)
accept_length: float
accept_rate: float
@dataclass
class LoRAMetrics:
"""LoRA adapter pool metrics."""
slots_used: int = field(metadata={"metric": ("gauge", "LoRA adapter slots in use")})
slots_total: int = field(metadata={"metric": ("gauge", "Total LoRA adapter slots")})
utilization: float = field(
metadata={"metric": ("gauge", "LoRA pool utilization ratio")}
)
slots_used: int
slots_total: int
utilization: float
@dataclass
class DisaggregationMetrics:
"""PD disaggregation metrics."""
mode: str # "prefill", "decode", or "null" - not a metric
prefill_bootstrap_queue_reqs: int = field(
default=0, metadata={"metric": ("gauge", "Prefill bootstrap queue requests")}
)
prefill_inflight_queue_reqs: int = field(
default=0, metadata={"metric": ("gauge", "Prefill inflight queue requests")}
)
decode_prealloc_queue_reqs: int = field(
default=0, metadata={"metric": ("gauge", "Decode prealloc queue requests")}
)
decode_transfer_queue_reqs: int = field(
default=0, metadata={"metric": ("gauge", "Decode transfer queue requests")}
)
decode_retracted_queue_reqs: int = field(
default=0, metadata={"metric": ("gauge", "Decode retracted queue requests")}
)
kv_transfer_speed_gb_s: float = field(
default=0.0, metadata={"metric": ("gauge", "KV transfer speed in GB/s")}
)
kv_transfer_latency_ms: float = field(
default=0.0, metadata={"metric": ("gauge", "KV transfer latency in ms")}
)
mode: str # "prefill", "decode", or "null"
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
@dataclass
class QueueMetrics:
"""Detailed queue breakdown."""
waiting: int = field(metadata={"metric": ("gauge", "Main waiting queue size")})
grammar: int = field(
metadata={"metric": ("gauge", "Grammar compilation queue size")}
)
paused: int = field(
metadata={"metric": ("gauge", "Requests paused by weight sync")}
)
retracted: int = field(metadata={"metric": ("gauge", "Retracted requests count")})
waiting: int
grammar: int
paused: int
retracted: int
@dataclass
@@ -2086,46 +2074,21 @@ class GetLoadsReqOutput(BaseReq):
dp_rank: int
timestamp: float
num_running_reqs: int = field(
metadata={"metric": ("gauge", "Number of running requests")}
)
num_waiting_reqs: int = field(
metadata={"metric": ("gauge", "Number of waiting requests")}
)
num_waiting_uncached_tokens: int = field(
metadata={
"metric": (
"gauge",
"Number of uncached input tokens waiting for prefill compute",
)
}
)
num_used_tokens: int = field(
metadata={"metric": ("gauge", "Number of tokens in use")}
)
num_running_reqs: int
num_waiting_reqs: int
num_waiting_uncached_tokens: int
num_used_tokens: int
# num_used_tokens plus pending tokens not already allocated in the KV pool.
# Used for DP balance.
num_total_tokens: int = field(
metadata={"metric": ("gauge", "Used tokens plus pending unallocated tokens")}
)
max_total_num_tokens: int = field(
metadata={"metric": ("gauge", "Maximum token capacity")}
)
num_total_tokens: int
max_total_num_tokens: int
# FIXME: token_usage is actually max usage across all pools (KV, SWA, mamba),
# not just KV token usage. Rename requires API deprecation.
token_usage: float = field(metadata={"metric": ("gauge", "Token pool usage ratio")})
gen_throughput: float = field(
metadata={"metric": ("gauge", "Generation throughput tokens/sec")}
)
cache_hit_rate: float = field(
metadata={"metric": ("gauge", "Prefix cache hit rate")}
)
utilization: float = field(
metadata={"metric": ("gauge", "Overall utilization ratio")}
)
max_running_requests: int = field(
metadata={"metric": ("gauge", "Maximum running requests capacity")}
)
token_usage: float
gen_throughput: float
cache_hit_rate: float
utilization: float
max_running_requests: int
memory: Optional[MemoryMetrics] = None
speculative: Optional[SpeculativeMetrics] = None
@@ -2167,27 +2130,19 @@ class DumperControlReqOutput(BaseReq):
error: str = ""
def sock_send(
sender: Union[zmq.Socket, zmq.asyncio.Socket],
obj: Any,
flags: int = 0,
) -> None:
sender.send_pyobj(obj, flags=flags)
def sock_send(socket: zmq.Socket, obj: Any, flags: int = 0) -> None:
socket.send_pyobj(obj, flags=flags)
def sock_recv(socket, flags=0):
def sock_recv(socket: zmq.Socket, flags: int = 0) -> Any:
return socket.recv_pyobj(flags=flags)
async def async_sock_send(
sender: zmq.asyncio.Socket,
obj: Any,
flags: int = 0,
) -> None:
await sender.send_pyobj(obj, flags=flags)
async def async_sock_send(socket: zmq.asyncio.Socket, obj: Any, flags: int = 0) -> None:
await socket.send_pyobj(obj, flags=flags)
async def async_sock_recv(socket, flags=0):
async def async_sock_recv(socket: zmq.asyncio.Socket, flags: int = 0) -> Any:
return await socket.recv_pyobj(flags=flags)
@@ -2225,3 +2180,27 @@ def _check_all_req_types():
_check_all_req_types()
# IPC struct types whose fields still use opaque annotations (Any, Dict[str, Any],
# List[Any], etc.) instead of precise types. Kept as an explicit registry so
# opaque usage can be audited and gradually narrowed.
# NOTE: GenerateReqInput and EmbeddingReqInput are standalone (not BaseReq/
# BaseBatchReq subclasses) and are tracked separately.
_REQ_TYPES_WITH_OPAQUE_FIELDS = (
TokenizedGenerateReqInput, # mm_data_mooncake: Optional[List[Any]]
UpdateWeightFromDiskReqInput, # manifest: Optional[Dict[str, Any]]
BackupDramReq, # weight_pointer_map: Dict[str, Any]
GetWeightsByNameReqOutput, # parameter: Optional[List[Any]]
CheckWeightsReqOutput, # payload: Optional[Dict[str, Any]]
GetInternalStateReqOutput, # internal_state: Dict[str, Any]
SetInternalStateReq, # server_args: Dict[str, Any]
SetInternalStateReqOutput, # server_args: Dict[str, Any]
VertexGenerateReqInput, # instances, parameters: Dict[str, Any]
RpcReqInput, # parameters: Optional[Dict[str, Any]]
LoadLoRAAdapterFromTensorsReqInput, # config_dict, added_tokens_config: Dict[str, Any]
SetInjectDumpMetadataReqInput, # dump_metadata: Dict[str, Any]
DumperControlReqInput, # body: Dict[str, Any]
DumperControlReqOutput, # response: List[Dict[str, Any]]
BatchTokenIDOutput, # customized_info: Optional[Dict[str, List[Any]]]
BatchStrOutput, # customized_info: Optional[Dict[str, List[Any]]]
)
@@ -202,20 +202,27 @@ class SchedulerOutputStreamer:
if phs is not None:
has_phs = True
# Optimize PHS for pickle: torch.stack reduces N __reduce_ex__
# calls to 1 across the ZMQ IPC boundary. We can only stack when
# *every* entry is non-None (homogeneous batch); mixed batches
# (some requests want PHS, others don't) keep the raw list so
# positional indexing on the receiver side stays correct.
# Optimize pooled hidden states (PHS) for IPC serialization.
# Two formats, disambiguated on the receiver side by length:
# Stacked: [stacked_tensor(N, ...)] — len 1, N > 1 requests
# Non-stacked: [tensor_0, tensor_1, ...] — len == N
# Stacking reduces N pickle/__reduce_ex__ calls to 1.
# Only possible when all entries are non-None and same shape.
# See paired receiver logic in tokenizer_manager.py.
stacked_phs = None
if has_phs:
all_have_phs = all(t is not None for t in phs_list)
if all_have_phs:
if all(t.shape == phs_list[0].shape for t in phs_list):
stacked_phs = torch.stack(phs_list)
if len(phs_list) > 1 and all(
t.shape == phs_list[0].shape for t in phs_list
):
# Stacked: single tensor, wrapped in a list.
stacked_phs = [torch.stack(phs_list)]
else:
# Non-stacked: 1 request, mixed shapes, or mixed None.
stacked_phs = phs_list
else:
# Non-stacked: some requests don't have PHS (None entries).
stacked_phs = phs_list
self.send_to_detokenizer.send_output(
@@ -1113,7 +1113,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
obj: Union[GenerateReqInput, EmbeddingReqInput],
input_text: str,
input_ids: Optional[List[int]],
input_embeds: Optional[Union[List[float], None]] = None,
input_embeds: Optional[List[List[float]]] = None,
mm_inputs=None,
token_type_ids: Optional[List[int]] = None,
) -> Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput]:
@@ -2050,11 +2050,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"embedding": recv_obj.embeddings[i],
"meta_info": meta_info,
}
if (
recv_obj.pooled_hidden_states is not None
and recv_obj.pooled_hidden_states[i] is not None
):
out_dict["pooled_hidden_state"] = recv_obj.pooled_hidden_states[i]
# Unpack pooled hidden states (PHS).
# See paired sender logic in output_streamer.py.
# Stacked: len == 1 and N > 1 → unwrap the tensor
# Non-stacked: len == N → index directly
pooled_hidden_states = recv_obj.pooled_hidden_states
if pooled_hidden_states is not None:
if len(pooled_hidden_states) == 1 and len(recv_obj.rids) > 1:
pooled_hidden_states = pooled_hidden_states[0]
if pooled_hidden_states[i] is not None:
out_dict["pooled_hidden_state"] = pooled_hidden_states[i]
# Set first_token_time on the first output batch.
# This is the single write point for first_token_time.