From 118d6b2e5ed30d9803728b845fb44e30b5289615 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Thu, 25 Jun 2026 12:51:15 -0700 Subject: [PATCH] [Cleanup] Style and type annotation improvements extracted from #28688 (#29224) --- python/sglang/srt/entrypoints/http_server.py | 33 +- .../srt/managers/data_parallel_controller.py | 6 +- python/sglang/srt/managers/io_struct.py | 387 +++++++++--------- .../scheduler_components/output_streamer.py | 21 +- .../sglang/srt/managers/tokenizer_manager.py | 17 +- 5 files changed, 217 insertions(+), 247 deletions(-) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 4d8581afc..78a091f7a 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -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"]) diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 2073343eb..50c49ff35 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -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) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 48e6671e2..2ea058648 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -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]]] +) diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 0d1c4fb2a..1e3d33bce 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -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( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 2984b0e41..9dd16b2a3 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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.