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