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