[Cleanup] IPC struct renames, better typing, and SenderWrapper removal (#29214)

This commit is contained in:
Lianmin Zheng
2026-06-24 17:25:24 -07:00
committed by GitHub
parent 563c3418a7
commit 6c839368e0
11 changed files with 357 additions and 344 deletions
@@ -14,7 +14,7 @@ import traceback
import uuid
from collections import defaultdict
from http import HTTPStatus
from typing import Dict, List, Optional, Set, Tuple, Union
from typing import Annotated, Dict, List, Optional, Set, Tuple, Union
import aiohttp
import numpy as np
@@ -23,7 +23,7 @@ import torch
import uvicorn
import zmq
import zmq.asyncio
from fastapi import FastAPI
from fastapi import Body, FastAPI
from fastapi.responses import ORJSONResponse, Response
from transformers import AutoProcessor
@@ -3856,7 +3856,7 @@ async def health_generate():
@app.api_route("/start_profile", methods=["GET", "POST"])
async def start_profile_async(obj: Optional[ProfileReq] = None):
async def start_profile_async(obj: Annotated[Optional[ProfileReq], Body()] = None):
if dp_dispatcher is not None:
if obj is not None:
obj.req_type = ProfileReqType.START_PROFILE
+65 -33
View File
@@ -28,6 +28,7 @@ import uuid
from contextlib import asynccontextmanager
from http import HTTPStatus
from typing import (
Annotated,
Any,
AsyncGenerator,
AsyncIterator,
@@ -43,6 +44,7 @@ import requests
import uvicorn
import uvloop
from fastapi import (
Body,
Depends,
FastAPI,
File,
@@ -739,7 +741,9 @@ async def get_load():
# curl -s -X POST http://localhost:30000/set_internal_state -H "Content-Type: application/json" -d '{"server_args": {"pp_max_micro_batch_size": 8}}'
@app.api_route("/set_internal_state", methods=["POST", "PUT"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def set_internal_state(obj: SetInternalStateReq, request: Request):
async def set_internal_state(
obj: Annotated[SetInternalStateReq, Body()], request: Request
):
res = await _global_state.tokenizer_manager.set_internal_state(obj)
return res
@@ -950,7 +954,9 @@ async def clear_hicache_storage_backend():
# }'
@app.api_route("/hicache/storage-backend", methods=["PUT"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def attach_hicache_storage_backend(obj: AttachHiCacheStorageReqInput):
async def attach_hicache_storage_backend(
obj: Annotated[AttachHiCacheStorageReqInput, Body()],
):
"""Attach (enable) HiCache storage backend at runtime.
Only allowed when there are NO running / queued requests.
@@ -964,7 +970,7 @@ async def attach_hicache_storage_backend(obj: AttachHiCacheStorageReqInput):
hicache_storage_prefetch_policy=obj.hicache_storage_prefetch_policy,
hicache_write_policy=obj.hicache_write_policy,
)
msg = getattr(ret, "message", "")
msg = ret.message
return Response(
content=(
(
@@ -991,7 +997,7 @@ async def detach_hicache_storage_backend():
return _admin_api_key_missing_response()
ret = await _global_state.tokenizer_manager.detach_hicache_storage()
msg = getattr(ret, "message", "")
msg = ret.message
return Response(
content=(
(
@@ -1024,7 +1030,7 @@ async def hicache_storage_backend_status():
@app.api_route("/start_profile", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def start_profile_async(obj: Optional[ProfileReq] = None):
async def start_profile_async(obj: Annotated[Optional[ProfileReq], Body()] = None):
"""Start profiling."""
await _global_state.tokenizer_manager.start_profile(obj or ProfileReq())
return Response(
@@ -1102,7 +1108,9 @@ async def dump_expert_distribution_record_async():
@app.post("/update_weights_from_disk")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def update_weights_from_disk(obj: UpdateWeightFromDiskReqInput, request: Request):
async def update_weights_from_disk(
obj: Annotated[UpdateWeightFromDiskReqInput, Body()], request: Request
):
"""Update the weights from disk inplace without re-launching the server."""
(
success,
@@ -1130,7 +1138,8 @@ async def update_weights_from_disk(obj: UpdateWeightFromDiskReqInput, request: R
@app.post("/init_weights_send_group_for_remote_instance")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def init_weights_send_group_for_remote_instance(
obj: InitWeightsSendGroupForRemoteInstanceReqInput, request: Request
obj: Annotated[InitWeightsSendGroupForRemoteInstanceReqInput, Body()],
request: Request,
):
(
success,
@@ -1148,7 +1157,7 @@ async def init_weights_send_group_for_remote_instance(
@app.post("/send_weights_to_remote_instance")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def send_weights_to_remote_instance(
obj: SendWeightsToRemoteInstanceReqInput, request: Request
obj: Annotated[SendWeightsToRemoteInstanceReqInput, Body()], request: Request
):
(
success,
@@ -1204,7 +1213,7 @@ async def remote_instance_transfer_engine_info(rank: int = None):
@app.post("/init_weights_update_group")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def init_weights_update_group(
obj: InitWeightsUpdateGroupReqInput, request: Request
obj: Annotated[InitWeightsUpdateGroupReqInput, Body()], request: Request
):
"""Initialize the parameter update group."""
success, message = await _global_state.tokenizer_manager.init_weights_update_group(
@@ -1220,7 +1229,7 @@ async def init_weights_update_group(
@app.post("/destroy_weights_update_group")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def destroy_weights_update_group(
obj: DestroyWeightsUpdateGroupReqInput, request: Request
obj: Annotated[DestroyWeightsUpdateGroupReqInput, Body()], request: Request
):
"""Destroy the parameter update group."""
(
@@ -1236,7 +1245,7 @@ async def destroy_weights_update_group(
@app.post("/update_weights_from_tensor")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def update_weights_from_tensor(
obj: UpdateWeightsFromTensorReqInput, request: Request
obj: Annotated[UpdateWeightsFromTensorReqInput, Body()], request: Request
):
"""Update the weights from tensor inplace without re-launching the server.
Notes:
@@ -1258,7 +1267,7 @@ async def update_weights_from_tensor(
@app.post("/update_weights_from_distributed")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def update_weights_from_distributed(
obj: UpdateWeightsFromDistributedReqInput, request: Request
obj: Annotated[UpdateWeightsFromDistributedReqInput, Body()], request: Request
):
"""Update model parameter from distributed online."""
(
@@ -1277,7 +1286,9 @@ async def update_weights_from_distributed(
@app.post("/update_weights_from_ipc")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Request):
async def update_weights_from_ipc(
obj: Annotated[UpdateWeightsFromIPCReqInput, Body()], request: Request
):
"""Update the weights from IPC (Inter-Process Communication) for checkpoint-engine integration."""
success, message = await _global_state.tokenizer_manager.update_weights_from_ipc(
obj, request
@@ -1294,7 +1305,9 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re
@app.post("/update_weight_version")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request):
async def update_weight_version(
obj: Annotated[UpdateWeightVersionReqInput, Body()], request: Request
):
"""Update the weight version. This operation requires no active requests."""
if obj.abort_all_requests:
_global_state.tokenizer_manager.abort_request(abort_all=True)
@@ -1325,7 +1338,9 @@ async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Reque
@app.api_route("/get_weights_by_name", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def get_weights_by_name(obj: GetWeightsByNameReqInput, request: Request):
async def get_weights_by_name(
obj: Annotated[GetWeightsByNameReqInput, Body()], request: Request
):
"""Get model parameter by name."""
try:
ret = await _global_state.tokenizer_manager.get_weights_by_name(obj, request)
@@ -1340,7 +1355,7 @@ async def get_weights_by_name(obj: GetWeightsByNameReqInput, request: Request):
@app.api_route("/release_memory_occupation", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def release_memory_occupation(
obj: ReleaseMemoryOccupationReqInput, request: Request
obj: Annotated[ReleaseMemoryOccupationReqInput, Body()], request: Request
):
"""Release GPU memory occupation temporarily."""
try:
@@ -1352,7 +1367,7 @@ async def release_memory_occupation(
@app.api_route("/resume_memory_occupation", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def resume_memory_occupation(
obj: ResumeMemoryOccupationReqInput, request: Request
obj: Annotated[ResumeMemoryOccupationReqInput, Body()], request: Request
):
"""Resume GPU memory occupation."""
try:
@@ -1364,7 +1379,8 @@ async def resume_memory_occupation(
@app.api_route("/weights_checker", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def check_weights(
obj: Optional[CheckWeightsReqInput] = None, request: Request = None
obj: Annotated[Optional[CheckWeightsReqInput], Body()] = None,
request: Request = None,
):
if obj is None:
obj = CheckWeightsReqInput()
@@ -1381,7 +1397,7 @@ async def check_weights(
@app.api_route("/slow_down", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def slow_down(obj: SlowDownReqInput, request: Request):
async def slow_down(obj: Annotated[SlowDownReqInput, Body()], request: Request):
"""Slow down the system deliberately. Only for testing. Example scenario:
when we want to test performance of D in large-scale PD disaggregation and have no enough nodes for P,
we can use this to slow down D to let it have enough running sequences, and then disable slowdown
@@ -1395,7 +1411,9 @@ async def slow_down(obj: SlowDownReqInput, request: Request):
@app.api_route("/load_lora_adapter", methods=["POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def load_lora_adapter(obj: LoadLoRAAdapterReqInput, request: Request):
async def load_lora_adapter(
obj: Annotated[LoadLoRAAdapterReqInput, Body()], request: Request
):
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.load_lora_adapter(obj, request)
@@ -1413,7 +1431,7 @@ async def load_lora_adapter(obj: LoadLoRAAdapterReqInput, request: Request):
@app.api_route("/load_lora_adapter_from_tensors", methods=["POST"])
async def load_lora_adapter_from_tensors(
obj: LoadLoRAAdapterFromTensorsReqInput, request: Request
obj: Annotated[LoadLoRAAdapterFromTensorsReqInput, Body()], request: Request
):
"""Load a new LoRA adapter from tensors without re-launching the server."""
result = await _global_state.tokenizer_manager.load_lora_adapter_from_tensors(
@@ -1428,7 +1446,9 @@ async def load_lora_adapter_from_tensors(
@app.api_route("/unload_lora_adapter", methods=["POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def unload_lora_adapter(obj: UnloadLoRAAdapterReqInput, request: Request):
async def unload_lora_adapter(
obj: Annotated[UnloadLoRAAdapterReqInput, Body()], request: Request
):
"""Load a new LoRA adapter without re-launching the server."""
result = await _global_state.tokenizer_manager.unload_lora_adapter(obj, request)
@@ -1445,7 +1465,7 @@ async def unload_lora_adapter(obj: UnloadLoRAAdapterReqInput, request: Request):
@app.api_route("/open_session", methods=["GET", "POST"])
async def open_session(obj: OpenSessionReqInput, request: Request):
async def open_session(obj: Annotated[OpenSessionReqInput, Body()], request: Request):
"""Open a session, and return its unique session id."""
try:
session_id = await _global_state.tokenizer_manager.open_session(obj, request)
@@ -1459,7 +1479,7 @@ async def open_session(obj: OpenSessionReqInput, request: Request):
@app.api_route("/close_session", methods=["GET", "POST"])
async def close_session(obj: CloseSessionReqInput, request: Request):
async def close_session(obj: Annotated[CloseSessionReqInput, Body()], request: Request):
"""Close the session."""
try:
await _global_state.tokenizer_manager.close_session(obj, request)
@@ -1470,7 +1490,9 @@ async def close_session(obj: CloseSessionReqInput, request: Request):
@app.api_route("/configure_logging", methods=["GET", "POST"])
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def configure_logging(obj: ConfigureLoggingReq, request: Request):
async def configure_logging(
obj: Annotated[ConfigureLoggingReq, Body()], request: Request
):
"""Configure the request logging options."""
_global_state.tokenizer_manager.configure_logging(obj)
return Response(status_code=200)
@@ -1478,7 +1500,7 @@ async def configure_logging(obj: ConfigureLoggingReq, request: Request):
@app.post("/abort_request")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def abort_request(obj: AbortReq, request: Request):
async def abort_request(obj: Annotated[AbortReq, Body()], request: Request):
"""Abort a request."""
try:
_global_state.tokenizer_manager.abort_request(
@@ -1490,7 +1512,9 @@ async def abort_request(obj: AbortReq, request: Request):
@app.post("/parse_function_call")
async def parse_function_call_request(obj: ParseFunctionCallReq, request: Request):
async def parse_function_call_request(
obj: Annotated[ParseFunctionCallReq, Body()], request: Request
):
"""
A native API endpoint to parse function calls from a text.
"""
@@ -1512,7 +1536,9 @@ async def parse_function_call_request(obj: ParseFunctionCallReq, request: Reques
@app.post("/separate_reasoning")
async def separate_reasoning_request(obj: SeparateReasoningReqInput, request: Request):
async def separate_reasoning_request(
obj: Annotated[SeparateReasoningReqInput, Body()], request: Request
):
"""
A native API endpoint to separate reasoning from a text.
"""
@@ -1520,7 +1546,7 @@ async def separate_reasoning_request(obj: SeparateReasoningReqInput, request: Re
parser = ReasoningParser(model_type=obj.reasoning_parser, request=request)
# 2) Call the non-stream parsing method (non-stream)
if getattr(obj, "return_blocks", False):
if obj.return_blocks:
blocks = parser.parse_non_stream_blocks(obj.text)
reasoning_blocks = [b["text"] for b in blocks if b["type"] == "reasoning"]
text_blocks = [b["text"] for b in blocks if b["type"] == "text"]
@@ -1534,7 +1560,7 @@ async def separate_reasoning_request(obj: SeparateReasoningReqInput, request: Re
"reasoning_text": reasoning_text,
"text": normal_text,
}
if getattr(obj, "return_blocks", False):
if obj.return_blocks:
response_data["reasoning_blocks"] = reasoning_blocks
response_data["text_blocks"] = text_blocks
response_data["blocks"] = blocks
@@ -1544,7 +1570,9 @@ async def separate_reasoning_request(obj: SeparateReasoningReqInput, request: Re
@app.post("/pause_generation")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def pause_generation(obj: PauseGenerationReqInput, request: Request):
async def pause_generation(
obj: Annotated[PauseGenerationReqInput, Body()], request: Request
):
"""Pause generation."""
await _global_state.tokenizer_manager.pause_generation(obj)
return ORJSONResponse(
@@ -1555,7 +1583,9 @@ async def pause_generation(obj: PauseGenerationReqInput, request: Request):
@app.post("/continue_generation")
@auth_level(AuthLevel.ADMIN_OPTIONAL)
async def continue_generation(obj: ContinueGenerationReqInput, request: Request):
async def continue_generation(
obj: Annotated[ContinueGenerationReqInput, Body()], request: Request
):
"""Continue generation."""
await _global_state.tokenizer_manager.continue_generation(obj)
return ORJSONResponse(
@@ -1894,7 +1924,9 @@ async def sagemaker_chat_completions(
## Vertex AI API
@app.post(os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate"))
async def vertex_generate(vertex_req: VertexGenerateReqInput, raw_request: Request):
async def vertex_generate(
vertex_req: Annotated[VertexGenerateReqInput, Body()], raw_request: Request
):
if not vertex_req.instances:
return []
inputs = {}
+10 -7
View File
@@ -3,9 +3,7 @@ from __future__ import annotations
import asyncio
import copy
from collections import deque
from typing import Deque, Generic, List, Optional, TypeVar
from sglang.srt.managers.io_struct import sock_send
from typing import Callable, Deque, Generic, List, Optional, TypeVar
T = TypeVar("T")
@@ -22,8 +20,13 @@ class FanOutCommunicator(Generic[T]):
Only one request is in-flight at any time in either mode.
"""
def __init__(self, sender, fan_out: int, mode="queueing"):
self._sender = sender
def __init__(
self,
send: Callable[[T], None],
fan_out: int,
mode: str = "queueing",
):
self._send = send
self._fan_out = fan_out
self._mode = mode
self._result_event: Optional[asyncio.Event] = None
@@ -41,7 +44,7 @@ class FanOutCommunicator(Generic[T]):
assert self._result_values is None
if obj is not None:
sock_send(self._sender, obj)
self._send(obj)
self._result_event = asyncio.Event()
self._result_values = []
@@ -61,7 +64,7 @@ class FanOutCommunicator(Generic[T]):
self._result_event = asyncio.Event()
if obj is not None:
sock_send(self._sender, obj)
self._send(obj)
# Capture local refs before await -- after event fires, the first
# awakened coroutine clears shared state; later awaiters use local refs.
+2 -2
View File
@@ -19,7 +19,7 @@ io_struct.py and schedule_batch.py.
"""
from dataclasses import dataclass
from typing import List, Union
from typing import List
import torch
@@ -37,7 +37,7 @@ class PositionalEmbeds:
positions: List of positions where embeddings should be injected.
"""
embeds: Union[List[torch.Tensor], torch.Tensor]
embeds: torch.Tensor
positions: List[int]
def __post_init__(self):
+205 -218
View File
@@ -14,13 +14,16 @@
"""
The definition of objects transferred between different
processes (TokenizerManager, DetokenizerManager, Scheduler).
Keep this file focused on IPC struct definitions so it stays concise. Put
normalizers, helper utilities, and future non-struct logic in the owning module
instead, such as sglang.srt.utils.common.
"""
from __future__ import annotations
import copy
import uuid
from abc import ABC
from array import array
from collections import Counter
from dataclasses import dataclass, field
@@ -43,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 BaseFinishReason, Modality
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.mm_utils import has_valid_data
from sglang.srt.observability.req_time_stats import (
APIServerReqTimeStats,
@@ -57,39 +60,20 @@ from sglang.srt.utils.field_validators import validate_optional_list_i64_1d_2d
# Handle serialization of Image for pydantic
if TYPE_CHECKING:
from PIL.Image import Image
from sglang.srt.managers.tokenizer_manager import SenderWrapper
else:
Image = Any
@dataclass
class BaseReq(ABC):
rid: Optional[Union[str, List[str]]] = field(default=None, kw_only=True)
class BaseReq:
rid: Optional[str] = field(default=None, kw_only=True)
http_worker_ipc: Optional[str] = field(default=None, kw_only=True)
def regenerate_rid(self):
"""Generate a new request ID and return it."""
if isinstance(self.rid, list):
self.rid = [uuid.uuid4().hex for _ in range(len(self.rid))]
else:
self.rid = uuid.uuid4().hex
return self.rid
def _validate_rid_uniqueness(self):
"""Validate that request IDs within a batch are unique."""
if isinstance(self.rid, list) and len(set(self.rid)) != len(self.rid):
counts = Counter(self.rid)
duplicates = [rid for rid, count in counts.items() if count > 1]
raise ValueError(
f"Duplicate request IDs detected within the request: {duplicates}"
)
@dataclass
class BaseBatchReq(ABC):
class BaseBatchReq:
rids: Optional[List[str]] = field(default=None, kw_only=True)
http_worker_ipcs: Optional[List[str]] = field(default=None, kw_only=True)
http_worker_ipcs: Optional[List[Optional[str]]] = field(default=None, kw_only=True)
def regenerate_rids(self):
"""Generate new request IDs and return them."""
@@ -97,29 +81,6 @@ class BaseBatchReq(ABC):
return self.rids
@dataclass
class SpeculativeDecodingMetricsMixin:
"""
Mixin class containing speculative decoding metrics.
This class consolidates speculative decoding metrics that are shared across
batch output types that support speculative decoding to avoid code duplication.
"""
# Verify count: number of verification forward passes
spec_verify_ct: List[int]
# Accepted drafts: Number of accepted draft tokens during speculative decoding
# (strict drafts-only count, excludes the bonus token).
spec_num_correct_drafts: List[int]
# Acceptance histogram: List of lists, where each inner list represents histogram counts.
# List index = number of accepted tokens in a step, List value = count of steps with that many accepted tokens.
# Example: histogram[0] = 5 means 5 steps with 0 accepted tokens, histogram[3] = 10 means 10 steps with 3 accepted tokens.
# Empty list [] when speculative decoding is disabled.
spec_correct_drafts_histogram: List[List[int]]
# Parameters for a session
@dataclass
class SessionParams:
@@ -132,9 +93,9 @@ class SessionParams:
# Type definitions for multimodal input data
# Individual data item types for each modality
ImageDataInputItem = Union[Image, str, ImageData, Dict]
ImageDataInputItem = Union[str, Dict, ImageData, Image]
AudioDataInputItem = Union[str, Dict]
VideoDataInputItem = Union[str, VideoData, Dict]
VideoDataInputItem = Union[str, Dict, VideoData]
# Union type for any multimodal data item
MultimodalDataInputItem = Union[
ImageDataInputItem, VideoDataInputItem, AudioDataInputItem
@@ -146,13 +107,22 @@ 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(BaseReq):
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.
#
# Use C-loop validator to replace Pydantic per-element type check for efficiency.
input_ids: Annotated[
Optional[Union[List[List[int]], List[int]]],
@@ -235,31 +205,27 @@ class GenerateReqInput(BaseReq):
bootstrap_pair_key: Optional[Union[List[str], str]] = None
decode_tp_size: Optional[Union[List[Optional[int]], int]] = None
# Require reasoning for the request (hybrid reasoning model only)
require_reasoning: bool = False
# 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
# Deprecated: use routed_dp_rank instead
data_parallel_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
# For background responses (OpenAI responses API)
background: bool = False
# Conversation id used for tracking requests
conversation_id: Optional[str] = None
# Require reasoning for the request (hybrid reasoning model only)
require_reasoning: bool = False
# Priority for the request
priority: Optional[int] = None
# Extra 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
# Routing key for routing-key schedule policy
routing_key: Optional[str] = None
# Whether to disallow logging for this request (e.g. due to ZDR)
no_logs: bool = False
@@ -268,10 +234,8 @@ class GenerateReqInput(BaseReq):
# (Internal) Whether to return bytes for image generation
return_bytes: bool = False
# Whether to return entropy
return_entropy: bool = False
# Whether to return prompt token IDs without computing logprobs
return_prompt_token_ids: bool = False
@@ -297,6 +261,23 @@ class GenerateReqInput(BaseReq):
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def regenerate_rid(self):
"""Generate a new request ID and return it."""
if isinstance(self.rid, list):
self.rid = [uuid.uuid4().hex for _ in range(len(self.rid))]
else:
self.rid = uuid.uuid4().hex
return self.rid
def _validate_rid_uniqueness(self):
"""Validate that request IDs within a batch are unique."""
if isinstance(self.rid, list) and len(set(self.rid)) != len(self.rid):
counts = Counter(self.rid)
duplicates = [rid for rid, count in counts.items() if count > 1]
raise ValueError(
f"Duplicate request IDs detected within the request: {duplicates}"
)
def contains_mm_input(self) -> bool:
return (
has_valid_data(self.image_data)
@@ -317,18 +298,6 @@ class GenerateReqInput(BaseReq):
ValueError: If inputs are not properly specified (e.g., none or all of
text, input_ids, input_embeds are provided)
"""
if self.data_parallel_rank is not None:
import warnings
warnings.warn(
"'data_parallel_rank' is deprecated, use 'routed_dp_rank' instead.",
DeprecationWarning,
stacklevel=2,
)
if self.routed_dp_rank is None:
self.routed_dp_rank = self.data_parallel_rank
self.data_parallel_rank = None
self._validate_inputs()
self._determine_batch_size()
self._handle_parallel_sampling()
@@ -767,8 +736,11 @@ class TokenizedGenerateReqInput(BaseReq):
input_text: str
# The input token ids
input_ids: Optional[array[int]]
# The input embeds
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]]
# The multimodal inputs
mm_inputs: object
token_type_ids: Optional[List[int]]
# The sampling parameters
sampling_params: SamplingParams
# Whether to return the logprobs
@@ -778,7 +750,7 @@ class TokenizedGenerateReqInput(BaseReq):
# If return logprobs, the number of top logprobs to return at each position.
top_logprobs_num: int
# If return logprobs, the token id to return logprob for
token_ids_logprob: List[int]
token_ids_logprob: Optional[List[int]]
# Whether to stream output
stream: bool
@@ -787,17 +759,10 @@ 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
# The input embeds
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None
# Embedding overrides to place at specific token positions.
positional_embed_overrides: Optional[PositionalEmbeds] = None
# Session info for continual prompting
session_params: Optional[SessionParams] = None
@@ -808,6 +773,8 @@ class TokenizedGenerateReqInput(BaseReq):
# of `CustomLogitProcessor` in python/sglang/srt/sampling/custom_logit_processor.py
# Use the processor's `to_str()` method to generate the serialized string.
custom_logit_processor: Optional[str] = None
# Embedding overrides to place at specific token positions.
positional_embed_overrides: Optional[PositionalEmbeds] = None
# For disaggregated inference
bootstrap_host: Optional[str] = None
@@ -816,35 +783,31 @@ class TokenizedGenerateReqInput(BaseReq):
bootstrap_pair_key: Optional[str] = None
decode_tp_size: Optional[int] = None
# Require reasoning for the request (hybrid reasoning model only)
require_reasoning: bool = False
# For DP routing
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
# Require reasoning for the request (hybrid reasoning model only)
require_reasoning: bool = False
# Priority for the request
priority: Optional[int] = None
# Extra key for classifying the request (e.g. cache_salt)
# Extra cache key for classifying the request (e.g. cache_salt)
extra_key: Optional[str] = None
# Routing key for routing-key schedule policy
routing_key: Optional[str] = None
# Whether to disallow logging for this request (e.g. due to ZDR)
no_logs: bool = False
# (Internal) Whether to return bytes for image generation
return_bytes: bool = False
# Whether to return entropy
return_entropy: bool = False
token_type_ids: Optional[List[int]] = None
need_wait_for_mm_inputs: bool = False
need_wait_for_mm_inputs: Optional[bool] = None
num_items_assigned: Optional[Dict[Modality, List[int]]] = None
mm_data_mooncake: Optional[List] = None
# Encoder URL snapshot frozen at tokenizer-side dispatch time so that
@@ -875,9 +838,19 @@ class BatchTokenizedGenerateReqInput(BaseBatchReq):
@dataclass
class EmbeddingReqInput(BaseReq):
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.
input_ids: Optional[Union[List[List[int]], List[int]]] = None
# Dummy input embeds for compatibility
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None
# The image input. It can be an image instance, file name, URL, or base64 encoded string.
# Can be formatted as:
# - Single image for a single request
@@ -889,8 +862,6 @@ class EmbeddingReqInput(BaseReq):
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
# The token ids for text; one can either specify text or input_ids.
input_ids: Optional[Union[List[List[int]], List[int]]] = None
# Placeholder token ID used to locate embedding override positions in input token IDs.
embed_override_token_id: Optional[int] = None
# Unresolved embedding overrides: per-input list of tensors.
@@ -900,49 +871,61 @@ class EmbeddingReqInput(BaseReq):
# 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
# Dummy input embeds for compatibility
input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = 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
# Priority for the request
priority: Optional[int] = None
# Routing key for routing-key schedule policy
routing_key: Optional[str] = None
# For background responses (OpenAI responses API)
background: bool = False
# Priority for the request
priority: Optional[int] = None
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
dimensions: Optional[int] = None
# Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False
# Whether to return prompt token IDs without computing logprobs
return_prompt_token_ids: bool = False
# Propagates trace context via Engine.encode/async_encode
external_trace_header: Optional[Dict] = None
received_time: Optional[float] = None
# The number of dimensions the resulting output embeddings should have. It is applicable for Matryoshka Embeddings.
dimensions: Optional[int] = 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
# Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False
# Whether to return prompt token IDs without computing logprobs
return_prompt_token_ids: bool = False
# Pre-computed delimiter indices for multi-item scoring.
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
def regenerate_rid(self):
"""Generate a new request ID and return it."""
if isinstance(self.rid, list):
self.rid = [uuid.uuid4().hex for _ in range(len(self.rid))]
else:
self.rid = uuid.uuid4().hex
return self.rid
def _validate_rid_uniqueness(self):
"""Validate that request IDs within a batch are unique."""
if isinstance(self.rid, list) and len(set(self.rid)) != len(self.rid):
counts = Counter(self.rid)
duplicates = [rid for rid, count in counts.items() if count > 1]
raise ValueError(
f"Duplicate request IDs detected within the request: {duplicates}"
)
def normalize_batch_and_arguments(self):
# at least one of text, input_ids, or image should be provided
if self.text is None and self.input_ids is None and self.image_data is None:
@@ -1092,12 +1075,14 @@ class TokenizedEmbeddingReqInput(BaseReq):
input_text: str
# The input token ids
input_ids: array[int]
# The image inputs
image_inputs: dict
# The multimodal inputs
mm_inputs: object
# The token type ids
token_type_ids: List[int]
token_type_ids: Optional[List[int]]
# Dummy sampling params for compatibility
sampling_params: SamplingParams
# LoRA related
lora_id: Optional[str] = None # None means just use the base model
# Embedding overrides to place at specific token positions.
positional_embed_overrides: Optional[PositionalEmbeds] = None
# For DP routing
@@ -1106,17 +1091,14 @@ 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
# LoRA related
lora_id: Optional[str] = None # None means just use the base model
# 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
# Whether to return pooled hidden states (pre-head transformer output)
return_pooled_hidden_states: bool = False
# For observability
time_stats: Optional[Union[APIServerReqTimeStats, DPControllerReqTimeStats]] = None
@dataclass
class BatchTokenizedEmbeddingReqInput(BaseBatchReq):
@@ -1133,10 +1115,18 @@ class BatchTokenizedEmbeddingReqInput(BaseBatchReq):
return iter(self.batch)
TokenLogprobValues = Optional[List[List[Optional[float]]]]
TokenLogprobIndices = Optional[List[List[Optional[int]]]]
TopLogprobValues = Optional[List[Optional[List[Optional[List[float]]]]]]
TopLogprobIndices = Optional[List[Optional[List[Optional[List[int]]]]]]
HiddenStateChunk = List[Optional[Union[float, List[float]]]]
OutputHiddenStates = Optional[List[Optional[List[HiddenStateChunk]]]]
@dataclass
class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
class BatchTokenIDOutput(BaseBatchReq):
# The finish reason
finished_reasons: List[BaseFinishReason]
finished_reasons: List[Optional[FinishReasonDict]]
# For incremental decoding
decoded_texts: List[str]
decode_ids: List[array[int]]
@@ -1155,49 +1145,49 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
cached_tokens: List[int]
# Logprobs
input_token_logprobs_val: List[float]
input_token_logprobs_idx: List[int]
output_token_logprobs_val: List[float]
output_token_logprobs_idx: List[int]
input_top_logprobs_val: List[List]
input_top_logprobs_idx: List[List]
output_top_logprobs_val: List[List]
output_top_logprobs_idx: List[List]
input_token_ids_logprobs_val: List[List]
input_token_ids_logprobs_idx: List[List]
output_token_ids_logprobs_val: List[List]
output_token_ids_logprobs_idx: List[List]
output_token_entropy_val: List[float]
input_token_logprobs_val: TokenLogprobValues
input_token_logprobs_idx: TokenLogprobIndices
output_token_logprobs_val: TokenLogprobValues
output_token_logprobs_idx: TokenLogprobIndices
input_top_logprobs_val: TopLogprobValues
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
output_token_entropy_val: Optional[List[Optional[float]]]
# Hidden states
output_hidden_states: List[List[float]]
output_hidden_states: OutputHiddenStates
# Per-request routed experts (input + output tokens), shape
# (token, layer, top_k). DetokenizerManager encodes to base64 into
# BatchStrOutput; on the skip_tokenizer_init path the scheduler sends this
# straight to TokenizerManager, which encodes on demand.
routed_experts: List[Optional[torch.Tensor]]
routed_experts: Optional[List[Optional[torch.Tensor]]]
indexer_topk: List[Optional[torch.Tensor]]
indexer_topk: Optional[List[Optional[torch.Tensor]]]
# The information of placeholder tokens (e.g., image token)
# idx is the index of the token in the prompt after expansion.
# val is the length of padded tokens after expansion.
placeholder_tokens_idx: List[Optional[List[int]]]
placeholder_tokens_val: List[Optional[List[int]]]
placeholder_tokens_idx: Optional[List[Optional[List[int]]]]
placeholder_tokens_val: Optional[List[Optional[List[int]]]]
# Number of times each request was retracted.
retraction_counts: List[int]
retraction_counts: Optional[List[int]] = None
# The trainer step id. Used to know which step's weights are used for sampling.
token_steps: List[List[int]] = None
token_steps: Optional[List[List[int]]] = None
# Customized info
customized_info: Optional[Dict[str, List[Any]]] = None
# Detailed breakdown of cached tokens by source (device/host/storage)
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
cached_tokens_details: Optional[List[Optional[CachedTokensDetails]]] = None
# DP rank of the scheduler that processed each request
dp_ranks: Optional[List[int]] = None
dp_ranks: Optional[List[Optional[int]]] = None
# For observability
time_stats: Optional[List[SchedulerReqTimeStats]] = None
@@ -1207,15 +1197,22 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
audio_tokens: Optional[List[int]] = None
video_tokens: Optional[List[int]] = None
# Verify count: number of verification forward passes
spec_verify_ct: Optional[List[int]] = None
# Accepted drafts
spec_num_correct_drafts: Optional[List[int]] = None
# Acceptance histogram
spec_correct_drafts_histogram: Optional[List[List[int]]] = None
@dataclass
class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
class BatchStrOutput(BaseBatchReq):
# The finish reason
finished_reasons: List[dict]
finished_reasons: List[Optional[FinishReasonDict]]
# The output decoded strings
output_strs: List[str]
# The token ids
output_ids: Optional[List[int]]
output_ids: Optional[List[array]]
# Token counts
prompt_tokens: List[int]
@@ -1224,48 +1221,48 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
cached_tokens: List[int]
# Logprobs
input_token_logprobs_val: List[float]
input_token_logprobs_idx: List[int]
output_token_logprobs_val: List[float]
output_token_logprobs_idx: List[int]
input_top_logprobs_val: List[List]
input_top_logprobs_idx: List[List]
output_top_logprobs_val: List[List]
output_top_logprobs_idx: List[List]
input_token_ids_logprobs_val: List[List]
input_token_ids_logprobs_idx: List[List]
output_token_ids_logprobs_val: List[List]
output_token_ids_logprobs_idx: List[List]
output_token_entropy_val: List[float]
input_token_logprobs_val: TokenLogprobValues
input_token_logprobs_idx: TokenLogprobIndices
output_token_logprobs_val: TokenLogprobValues
output_token_logprobs_idx: TokenLogprobIndices
input_top_logprobs_val: TopLogprobValues
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
output_token_entropy_val: Optional[List[Optional[float]]]
# Hidden states
output_hidden_states: List[List[float]]
output_hidden_states: OutputHiddenStates
# Per-request routed experts, base64-encoded by DetokenizerManager off the
# tokenizer hot path. Underlying tensor shape is (token, layer, top_k);
# see BatchTokenIDOutput.routed_experts.
routed_experts: List[Optional[str]]
routed_experts: Optional[List[Optional[str]]]
indexer_topk: List[Optional[str]]
indexer_topk: Optional[List[Optional[str]]]
# The information of placeholder tokens (e.g., image token)
# idx is the index of the token in the prompt after expansion.
# val is the length of padded tokens after expansion.
placeholder_tokens_idx: List[Optional[List[int]]]
placeholder_tokens_val: List[Optional[List[int]]]
placeholder_tokens_idx: Optional[List[Optional[List[int]]]]
placeholder_tokens_val: Optional[List[Optional[List[int]]]]
# Number of times each request was retracted.
retraction_counts: List[int]
retraction_counts: Optional[List[int]] = None
# The trainer step id. Used to know which step's weights are used for sampling.
token_steps: List[List[int]] = None
token_steps: Optional[List[List[int]]] = None
# Customized info
customized_info: Optional[Dict[str, List[Any]]] = None
# Detailed breakdown of cached tokens by source (device/host/storage)
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
cached_tokens_details: Optional[List[Optional[CachedTokensDetails]]] = None
# DP rank of the scheduler that processed each request
dp_ranks: Optional[List[int]] = None
dp_ranks: Optional[List[Optional[int]]] = None
# For observability
time_stats: Optional[List[SchedulerReqTimeStats]] = None
@@ -1275,24 +1272,31 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
audio_tokens: Optional[List[int]] = None
video_tokens: Optional[List[int]] = None
# Verify count: number of verification forward passes
spec_verify_ct: Optional[List[int]] = None
# Accepted drafts
spec_num_correct_drafts: Optional[List[int]] = None
# Acceptance histogram
spec_correct_drafts_histogram: Optional[List[List[int]]] = None
@dataclass
class BatchEmbeddingOutput(BaseBatchReq):
# The finish reason
finished_reasons: List[BaseFinishReason]
finished_reasons: List[Optional[FinishReasonDict]]
# The output embedding
embeddings: Union[List[List[float]], List[Dict[int, float]]]
# Token counts
prompt_tokens: List[int]
cached_tokens: List[int]
# Placeholder token info
placeholder_tokens_idx: List[Optional[List[int]]]
placeholder_tokens_val: List[Optional[List[int]]]
placeholder_tokens_idx: Optional[List[Optional[List[int]]]]
placeholder_tokens_val: Optional[List[Optional[List[int]]]]
# Number of times each request was retracted.
retraction_counts: List[int]
retraction_counts: Optional[List[int]] = None
# Detailed breakdown of cached tokens by source (device/host/storage)
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
cached_tokens_details: Optional[List[Optional[CachedTokensDetails]]] = None
# For observability
time_stats: Optional[List[SchedulerReqTimeStats]] = None
@@ -1378,26 +1382,6 @@ class AttachHiCacheStorageReqInput(BaseReq):
hicache_storage_prefetch_policy: Optional[str] = None
hicache_write_policy: Optional[str] = None
def __post_init__(self):
if self.hicache_storage_prefetch_policy is None:
pass
else:
allowed = ["best_effort", "wait_complete", "timeout"]
if self.hicache_storage_prefetch_policy not in allowed:
raise ValueError(
f"Invalid hicache_storage_prefetch_policy: {self.hicache_storage_prefetch_policy!r}. "
f"Expected one of {allowed}."
)
if self.hicache_write_policy is None:
return
allowed = ["write_back", "write_through", "write_through_selective"]
if self.hicache_write_policy not in allowed:
raise ValueError(
f"Invalid hicache_write_policy: {self.hicache_write_policy!r}. "
f"Expected one of {allowed}."
)
@dataclass
class AttachHiCacheStorageReqOutput(BaseReq):
@@ -1440,13 +1424,6 @@ class PauseGenerationReqInput(BaseReq):
mode: Literal["abort", "retract", "in_place"] = "abort"
def __post_init__(self):
allowed = ["abort", "retract", "in_place"]
if self.mode not in allowed:
raise ValueError(
f"Invalid mode: {self.mode!r}. " f"Expected one of {allowed}."
)
@dataclass
class ContinueGenerationReqInput(BaseReq):
@@ -1459,14 +1436,14 @@ class ContinueGenerationReqInput(BaseReq):
@dataclass
class TokenizerWorkerRegistration:
class TokenizerWorkerRegistrationReq(BaseReq):
"""Sent by each TokenizerWorker on startup to register its IPC name with the router."""
worker_ipc_name: str
@dataclass
class PauseContinueBroadcast:
class PauseContinueBroadcastReq(BaseReq):
"""Broadcast from router to all workers to set is_pause state."""
is_pause: bool
@@ -1636,7 +1613,7 @@ class InitWeightsUpdateGroupReqInput(BaseReq):
# The master address
master_address: str
# The master port
master_port: int
master_port: Union[int, str]
# The rank offset
rank_offset: int
# The world size
@@ -1733,8 +1710,8 @@ class SlowDownReqOutput(BaseReq):
class AbortReq(BaseReq):
# Whether to abort all requests
abort_all: bool = False
# The finished reason data
finished_reason: Optional[Dict[str, Any]] = None
# The finished reason data (from BaseFinishReason.to_json())
finished_reason: Optional[FinishReasonDict] = None
abort_message: Optional[str] = None
def __post_init__(self):
@@ -1876,7 +1853,7 @@ class ExpertDistributionReqOutput(BaseReq):
class Function:
description: Optional[str] = None
name: Optional[str] = None
parameters: Optional[object] = None
parameters: Optional[Any] = None
@dataclass
@@ -1993,7 +1970,7 @@ class BlockReqType(Enum):
@dataclass
class BlockReqInput(BaseReq):
type: BlockReqType
req_type: BlockReqType
@dataclass
@@ -2191,7 +2168,7 @@ class DumperControlReqOutput(BaseReq):
def sock_send(
sender: Union[zmq.Socket, zmq.asyncio.Socket, SenderWrapper],
sender: Union[zmq.Socket, zmq.asyncio.Socket],
obj: Any,
flags: int = 0,
) -> None:
@@ -2203,7 +2180,7 @@ def sock_recv(socket, flags=0):
async def async_sock_send(
sender: Union[zmq.asyncio.Socket, SenderWrapper],
sender: zmq.asyncio.Socket,
obj: Any,
flags: int = 0,
) -> None:
@@ -2214,6 +2191,14 @@ async def async_sock_recv(socket, flags=0):
return await socket.recv_pyobj(flags=flags)
# The following request types are either defined in other files,
# or not subclasses of BaseReq/BaseBatchReq, so we skip the check for them.
_IGNORE_REQ_TYPES_CHECK = (
GenerateReqInput.__name__,
EmbeddingReqInput.__name__,
)
def _check_all_req_types():
"""A helper function to check all request types are defined in this file."""
import inspect
@@ -2223,6 +2208,8 @@ def _check_all_req_types():
for class_type in all_classes:
# check its name
name = class_type[0]
if name in _IGNORE_REQ_TYPES_CHECK:
continue
is_io_struct = (
name.endswith("Req") or name.endswith("Input") or name.endswith("Output")
)
@@ -29,7 +29,7 @@ import sys
import threading
import zlib
from multiprocessing import shared_memory
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, List, Optional
import psutil
import setproctitle
@@ -46,9 +46,9 @@ from sglang.srt.managers.io_struct import (
BatchTokenIDOutput,
ContinueGenerationReqInput,
FreezeGCReq,
PauseContinueBroadcast,
PauseContinueBroadcastReq,
PauseGenerationReqInput,
TokenizerWorkerRegistration,
TokenizerWorkerRegistrationReq,
async_sock_recv,
async_sock_send,
sock_recv,
@@ -437,7 +437,7 @@ class MultiTokenizerRouter:
while True:
recv_obj = await async_sock_recv(self.receive_from_worker)
if isinstance(recv_obj, TokenizerWorkerRegistration):
if isinstance(recv_obj, TokenizerWorkerRegistrationReq):
if recv_obj.worker_ipc_name not in self.all_worker_ipcs:
self.all_worker_ipcs.add(recv_obj.worker_ipc_name)
logger.info(
@@ -451,7 +451,7 @@ class MultiTokenizerRouter:
):
# Broadcast to ALL workers so every worker's is_pause is set
is_pause = isinstance(recv_obj, PauseGenerationReqInput)
broadcast = PauseContinueBroadcast(is_pause=is_pause)
broadcast = PauseContinueBroadcastReq(is_pause=is_pause)
for ipc_name in self.all_worker_ipcs:
self.socket_mapping.send_output(ipc_name, broadcast)
# Forward to scheduler rank 0 (it broadcasts to all TP/PP/DP
@@ -603,18 +603,18 @@ class TokenizerWorker(TokenizerManager):
)
# Register this worker with the router for pause/continue broadcasting
reg = TokenizerWorkerRegistration(worker_ipc_name=self.tokenizer_ipc_name)
sock_send(self.send_to_scheduler, reg)
reg = TokenizerWorkerRegistrationReq(worker_ipc_name=self.tokenizer_ipc_name)
self._dispatch_to_scheduler(reg)
# Future for awaiting pause/continue broadcast confirmation
self._pause_continue_future: Optional[asyncio.Future] = None
# Register PauseContinueBroadcast in the result dispatcher so
# Register PauseContinueBroadcastReq in the result dispatcher so
# handle_loop routes it to _handle_pause_continue_broadcast
from sglang.utils import TypeBasedDispatcher
self._result_dispatcher += TypeBasedDispatcher(
[(PauseContinueBroadcast, self._handle_pause_continue_broadcast)]
[(PauseContinueBroadcastReq, self._handle_pause_continue_broadcast)]
)
async def pause_generation(self, obj: PauseGenerationReqInput):
@@ -622,7 +622,7 @@ class TokenizerWorker(TokenizerManager):
self._pause_continue_future = loop.create_future()
# Send to router which will broadcast to all workers
# (router also handles forwarding to scheduler for non-abort modes)
sock_send(self.send_to_scheduler, obj)
self._dispatch_to_scheduler(obj)
await self._pause_continue_future
if obj.mode == "abort":
@@ -637,15 +637,15 @@ class TokenizerWorker(TokenizerManager):
async def continue_generation(self, obj: ContinueGenerationReqInput):
loop = asyncio.get_event_loop()
self._pause_continue_future = loop.create_future()
sock_send(self.send_to_scheduler, obj)
self._dispatch_to_scheduler(obj)
await self._pause_continue_future
def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcast):
def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcastReq):
"""Called from handle_loop when a broadcast arrives from the router."""
loop = asyncio.get_event_loop()
loop.create_task(self._apply_pause_continue_broadcast(obj))
async def _apply_pause_continue_broadcast(self, obj: PauseContinueBroadcast):
async def _apply_pause_continue_broadcast(self, obj: PauseContinueBroadcastReq):
"""Apply pause/continue state under the condition lock."""
async with self.is_pause_cond:
if obj.is_pause:
@@ -659,15 +659,6 @@ class TokenizerWorker(TokenizerManager):
self._pause_continue_future.set_result(True)
self._pause_continue_future = None
def _attach_multi_http_worker_info(self, req: Union[BaseReq, BaseBatchReq]):
if isinstance(req, BaseReq):
req.http_worker_ipc = self.tokenizer_ipc_name
elif isinstance(req, BaseBatchReq):
req.http_worker_ipcs = [self.tokenizer_ipc_name] * len(req.rids)
else:
raise ValueError(f"Unknown req type: {type(req)}")
async def print_exception_wrapper(func):
"""
+8 -4
View File
@@ -2448,8 +2448,8 @@ class Scheduler(
req.tokenizer = self.tokenizer
# Handle multimodal inputs
if recv_req.image_inputs is not None:
image_inputs = self._get_multimodal_inputs(recv_req.image_inputs)
if recv_req.mm_inputs is not None:
image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs)
# Expand a single image token into multiple dummy tokens for receiving image embeddings
# The `pad_input_ids_func` is model-specific and may be None for
# embedding models or models not requiring special padding.
@@ -3793,9 +3793,13 @@ class Scheduler(
for k, v in server_args_dict.items():
setattr(get_global_server_args(), k, v)
logger.info(f"Global server args updated! {get_global_server_args()=}")
server_args = dict(vars(get_global_server_args()))
# This field is not serializable.
server_args.pop("model_config", None)
return SetInternalStateReqOutput(
updated=True,
server_args=vars(get_global_server_args()),
updated=if_success,
server_args=server_args,
)
def save_remote_model(self, **kwargs):
@@ -18,6 +18,7 @@ from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput,
BatchTokenIDOutput,
CachedTokensDetails,
)
from sglang.srt.managers.schedule_batch import (
BaseFinishReason,
@@ -55,7 +56,7 @@ class SchedulerOutputStreamer:
storage_backend_type = type(storage_backend).__name__
return storage_backend_type
def get_cached_tokens_details(self, req: Req) -> Optional[dict]:
def get_cached_tokens_details(self, req: Req) -> Optional[CachedTokensDetails]:
"""Get detailed cache breakdown for a request, if available.
Returns:
@@ -246,7 +247,7 @@ class _GenerationStreamAccumulator:
disaggregation_mode: DisaggregationMode
default_stream_interval: int
default_force_stream_interval: int
get_cached_tokens_details: Callable[[Req], Optional[dict]]
get_cached_tokens_details: Callable[[Req], Optional[CachedTokensDetails]]
rids: list = field(default_factory=list)
http_worker_ipcs: list = field(default_factory=list)
@@ -14,9 +14,9 @@
import logging
from contextlib import contextmanager
from enum import Enum, auto
from typing import Any, List, Optional
from typing import Any, Callable, List, Optional
from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType, sock_send
from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType
from sglang.srt.utils.poll_based_barrier import PollBasedBarrier
logger = logging.getLogger(__name__)
@@ -51,10 +51,10 @@ class SchedulerInputBlocker:
def _handle_recv_req(self, recv_req):
if isinstance(recv_req, BlockReqInput):
if recv_req.type == BlockReqType.BLOCK:
if recv_req.req_type == BlockReqType.BLOCK:
self._execute_block_req()
return []
elif recv_req.type == BlockReqType.UNBLOCK:
elif recv_req.req_type == BlockReqType.UNBLOCK:
self._execute_unblock_req()
return []
else:
@@ -98,9 +98,9 @@ class _State(Enum):
@contextmanager
def input_blocker_guard_region(send_to_scheduler):
sock_send(send_to_scheduler, BlockReqInput(BlockReqType.BLOCK))
def input_blocker_guard_region(dispatch_to_scheduler: Callable[[BlockReqInput], None]):
dispatch_to_scheduler(BlockReqInput(req_type=BlockReqType.BLOCK))
try:
yield
finally:
sock_send(send_to_scheduler, BlockReqInput(BlockReqType.UNBLOCK))
dispatch_to_scheduler(BlockReqInput(req_type=BlockReqType.UNBLOCK))
@@ -71,8 +71,6 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromIPCReqOutput,
UpdateWeightsFromTensorReqInput,
UpdateWeightsFromTensorReqOutput,
async_sock_send,
sock_send,
)
from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.server_args import LoRARef, ServerArgs
@@ -134,7 +132,11 @@ class TokenizerControlMixin:
for spec in _COMMUNICATOR_SPECS:
name, resp_type = spec[0], spec[1]
mode = spec[2] if len(spec) > 2 else "queueing"
comm = FanOutCommunicator(self.send_to_scheduler, server_args.dp_size, mode)
comm = FanOutCommunicator(
self._dispatch_to_scheduler,
server_args.dp_size,
mode,
)
setattr(self, f"{name}_communicator", comm)
dispatch_pairs.append((resp_type, comm.handle_recv))
self._result_dispatcher += TypeBasedDispatcher(dispatch_pairs)
@@ -848,7 +850,7 @@ class TokenizerControlMixin:
future = asyncio.Future()
self.session_futures[obj.session_id] = future
sock_send(self.send_to_scheduler, obj)
self._dispatch_to_scheduler(obj)
try:
return await future
@@ -860,7 +862,7 @@ class TokenizerControlMixin:
obj: CloseSessionReqInput,
request: Optional[fastapi.Request] = None,
):
await async_sock_send(self.send_to_scheduler, obj)
await self._async_dispatch_to_scheduler(obj)
def _update_weight_version_if_provided(
self: TokenizerManager, weight_version: Optional[str]
+34 -41
View File
@@ -55,6 +55,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import (
AbortReq,
ActiveRanksOutput,
BaseBatchReq,
BaseReq,
BatchEmbeddingOutput,
BatchStrOutput,
@@ -383,19 +384,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
context, zmq.PULL, port_args.tokenizer_ipc_name, True
)
if self.server_args.tokenizer_worker_num == 1:
send_to_scheduler = get_zmq_socket(
self.send_to_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.scheduler_input_ipc_name, True
)
self.send_to_scheduler = SenderWrapper(port_args, send_to_scheduler)
self.tokenizer_ipc_name = None
else:
# Use tokenizer_worker_ipc_name in multi-tokenizer mode
send_to_scheduler = get_zmq_socket(
self.send_to_scheduler = get_zmq_socket(
context, zmq.PUSH, port_args.tokenizer_worker_ipc_name, False
)
# Make sure that each request carries the tokenizer_ipc_name for response routing
self.send_to_scheduler = SenderWrapper(
port_args, send_to_scheduler, attach_multi_http_worker_info=True
)
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
self.load_snapshot_reader = create_load_snapshot_reader(
self.server_args,
@@ -403,6 +401,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
caller="TokenizerManager",
)
def _dispatch_to_scheduler(self, obj: Any) -> None:
if self.tokenizer_ipc_name is not None:
stamp_http_worker_ipc(obj, self.tokenizer_ipc_name)
sock_send(self.send_to_scheduler, obj)
async def _async_dispatch_to_scheduler(self, obj: Any) -> None:
if self.tokenizer_ipc_name is not None:
stamp_http_worker_ipc(obj, self.tokenizer_ipc_name)
await async_sock_send(self.send_to_scheduler, obj)
def init_running_status(self):
# Request states
self.rid_to_state: Dict[str, ReqState] = {}
@@ -599,8 +607,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
f"routed_dp_rank={obj.routed_dp_rank} out of range [0, {dp_size})"
)
if self.server_args.tokenizer_worker_num > 1:
self._attach_multi_http_worker_info(obj)
self._init_req_state(obj, request)
try:
if self.server_args.language_only:
@@ -1192,7 +1198,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
tokenized_obj = TokenizedEmbeddingReqInput(
input_text=input_text,
input_ids=input_ids_arr,
image_inputs=mm_inputs,
mm_inputs=mm_inputs,
token_type_ids=token_type_ids,
sampling_params=sampling_params,
positional_embed_overrides=positional_embed_overrides,
@@ -1326,7 +1332,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
sock_send(self.send_to_scheduler, tokenized_obj)
self._dispatch_to_scheduler(tokenized_obj)
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
def _send_batch_request(
@@ -1342,7 +1348,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
sock_send(self.send_to_scheduler, batch_req)
self._dispatch_to_scheduler(batch_req)
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
def _coalesce_streaming_chunks(
@@ -1561,7 +1567,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
else:
# Sequential tokenization and processing
with (
input_blocker_guard_region(send_to_scheduler=self.send_to_scheduler)
input_blocker_guard_region(
dispatch_to_scheduler=self._dispatch_to_scheduler,
)
if get_bool_env_var("SGLANG_ENABLE_COLOCATED_BATCH_GEN")
else nullcontext()
):
@@ -1667,7 +1675,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
return
req = AbortReq(rid=rid, abort_all=abort_all)
sock_send(self.send_to_scheduler, req)
self._dispatch_to_scheduler(req)
if self.enable_metrics:
# TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request(
@@ -1678,7 +1686,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async with self.is_pause_cond:
self.is_pause = True
if obj.mode != "abort":
await async_sock_send(self.send_to_scheduler, obj)
await self._async_dispatch_to_scheduler(obj)
else:
# we are using the model_update_lock to check if there is still on-going requests.
while True:
@@ -1692,7 +1700,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def continue_generation(self, obj: ContinueGenerationReqInput):
async with self.is_pause_cond:
self.is_pause = False
await async_sock_send(self.send_to_scheduler, obj)
await self._async_dispatch_to_scheduler(obj)
self.is_pause_cond.notify_all()
async def update_weights_from_disk(
@@ -1737,7 +1745,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]:
sock_send(self.send_to_scheduler, obj)
self._dispatch_to_scheduler(obj)
self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1:
result = await self.model_update_result
@@ -1777,12 +1785,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Let the exception propagate to the caller.
# Only legal requests will be sent to scheduler.
logging.getLogger().setLevel(obj.log_level.upper())
sock_send(self.send_to_scheduler, obj)
self._dispatch_to_scheduler(obj)
logging.info(f"Config logging: {obj=}")
async def freeze_gc(self):
"""Send a freeze_gc message to the scheduler first, then freeze locally."""
sock_send(self.send_to_scheduler, FreezeGCReq())
self._dispatch_to_scheduler(FreezeGCReq())
freeze_gc("Tokenizer Manager")
return None
@@ -2652,7 +2660,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self._subprocess_watchdog.stop()
# Ask schedulers to release resources in userspace and exit (see
# ShutdownReq), then wait for them before hard-killing the rest.
sock_send(self.send_to_scheduler, ShutdownReq())
self._dispatch_to_scheduler(ShutdownReq())
deadline = time.monotonic() + 15
while time.monotonic() < deadline and collect_scheduler_processes():
time.sleep(0.1)
@@ -2721,7 +2729,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.event.set()
def update_active_ranks(self, ranks: ActiveRanksOutput):
sock_send(self.send_to_scheduler, ranks)
self._dispatch_to_scheduler(ranks)
def _handle_open_session_req_output(self, recv_obj):
future = self.session_futures.get(recv_obj.session_id)
@@ -3121,23 +3129,8 @@ class SignalHandler:
#
class SenderWrapper:
def __init__(
self,
port_args,
send_to_scheduler,
attach_multi_http_worker_info=False,
):
self.port_args = port_args
self.send_to_scheduler = send_to_scheduler
self.attach_multi_http_worker_info = attach_multi_http_worker_info
def _stamp_http_worker_ipc(self, obj):
if not self.attach_multi_http_worker_info:
return
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
def send_pyobj(self, obj, flags=0):
self._stamp_http_worker_ipc(obj)
return self.send_to_scheduler.send_pyobj(obj, flags=flags)
def stamp_http_worker_ipc(obj: Any, ipc_name: str) -> None:
if isinstance(obj, BaseReq):
obj.http_worker_ipc = ipc_name
elif isinstance(obj, BaseBatchReq):
obj.http_worker_ipcs = [ipc_name] * len(obj.rids)