[Refactor] Introduce sock_send/sock_recv wrappers for zmq IPC (#29012)

This commit is contained in:
Lianmin Zheng
2026-06-23 15:54:36 -07:00
committed by GitHub
parent ecab3f322e
commit 34dd9c28ca
22 changed files with 258 additions and 154 deletions
+9 -7
View File
@@ -20,6 +20,9 @@ from typing import Any, List, Literal, Optional, Union, get_args, get_type_hints
import torch
import torch.distributed as dist
import zmq
from sglang.srt.managers.io_struct import sock_recv, sock_send
# -------------------------------------- config base ------------------------------------------
@@ -1419,8 +1422,6 @@ def _create_zmq_rpc_broadcast(
handler, timeout_seconds: int = 60
) -> Optional["_ZmqRpcBroadcast"]:
"""A general-purpose minimal RPC to support broadcasting executions to multi processes"""
import zmq
rank = _get_rank()
world_size = dist.get_world_size() if dist.is_initialized() else 1
@@ -1433,13 +1434,13 @@ def _create_zmq_rpc_broadcast(
def serve_loop():
while True:
try:
req = sock.recv_pyobj()
req = sock_recv(sock)
result = getattr(handler, req["method"])(*req["args"], **req["kwargs"])
resp = {"result": result, "error": None}
except Exception as e:
_log(f"[ZmqRpc] error inside handler: {e}")
resp = {"result": None, "error": str(e)}
sock.send_pyobj(resp)
sock_send(sock, resp)
thread = threading.Thread(target=serve_loop, daemon=True)
thread.start()
@@ -1476,14 +1477,15 @@ class _ZmqRpcHandle:
def __getattr__(self, method_name: str):
def call(*args, **kwargs):
self._socket.send_pyobj(
sock_send(
self._socket,
{
"method": method_name,
"args": args,
"kwargs": kwargs,
}
},
)
response = self._socket.recv_pyobj()
response = sock_recv(self._socket)
if response["error"]:
raise RuntimeError(
f"RPC error on {self._debug_name}: {response['error']}"
@@ -26,6 +26,7 @@ from sglang.srt.disaggregation.encode_server import (
handle_scheduler_receive_url_request,
launch_encoder,
)
from sglang.srt.managers.io_struct import async_sock_send
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import random_uuid
@@ -95,7 +96,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
"part_idx": request.part_idx,
}
for socket in self.send_sockets:
await socket.send_pyobj(request_dict)
await async_sock_send(socket, request_dict)
# gRPC encode is image-only; encoder.encode() requires modality
(
@@ -44,7 +44,14 @@ from sglang.srt.distributed.parallel_state import (
)
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileReqType
from sglang.srt.managers.io_struct import (
ProfileReq,
ProfileReqInput,
ProfileReqType,
async_sock_recv,
async_sock_send,
sock_send,
)
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
from sglang.srt.model_loader import get_model
@@ -2390,13 +2397,14 @@ class EncoderScheduler:
requests = [p.request for p in group]
start = time.time()
for sock in self.send_sockets:
sock.send_pyobj(
sock_send(
sock,
{
"type": "batch_encode",
"modality": modality.name,
"requests": requests,
"enter_time": start,
}
},
)
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
@@ -2442,7 +2450,7 @@ class EncoderScheduler:
req = p.request
try:
for sock in self.send_sockets:
sock.send_pyobj(req)
sock_send(sock, req)
result = await self.encoder.encode_request(req, modality)
if not p.future.done():
p.future.set_result(result)
@@ -2731,7 +2739,7 @@ class DPDispatcher:
)
try:
await self.dispatch_sockets[rank].send_pyobj(request)
await async_sock_send(self.dispatch_sockets[rank], request)
# An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip
# the watchdog, so bound the wait explicitly.
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
@@ -2780,7 +2788,7 @@ class DPDispatcher:
f"dp_rank={rank}, pending={self.pending_counts}"
)
try:
await self.dispatch_sockets[rank].send_pyobj(request)
await async_sock_send(self.dispatch_sockets[rank], request)
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
except asyncio.TimeoutError:
self.pending_futures[rank].pop(key, None)
@@ -2821,7 +2829,7 @@ class DPDispatcher:
self.req_id_to_rank[req_id] = rank
rank_keys.append((rank, req_id))
request_copy = {**request, "req_id": req_id}
await self.dispatch_sockets[rank].send_pyobj(request_copy)
await async_sock_send(self.dispatch_sockets[rank], request_copy)
futures.append(future)
# Concurrent wait → total bounded by eff_timeout, not
# dp_size × eff_timeout.
@@ -2901,7 +2909,7 @@ class DPDispatcher:
consecutive_errors = 0
while True:
try:
msg = await self.result_socket.recv_pyobj()
msg = await async_sock_recv(self.result_socket)
consecutive_errors = 0
except asyncio.CancelledError:
raise
@@ -3051,10 +3059,10 @@ async def _dp_worker_handle_request(
"_error_code": err_code,
}
# pyzmq async send_pyobj isn't safe for concurrent senders.
# pyzmq async send isn't safe for concurrent senders.
try:
async with send_lock:
await send_sock.send_pyobj(envelope)
await async_sock_send(send_sock, envelope)
except Exception:
logger.error(
f"DP worker {dp_rank} failed to send envelope for "
@@ -3113,7 +3121,7 @@ async def run_dp_worker(
spawned = False
try:
try:
request = await recv_sock.recv_pyobj()
request = await async_sock_recv(recv_sock)
except asyncio.CancelledError:
raise
except Exception:
@@ -3197,7 +3205,7 @@ async def run_encoder(
):
encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank)
while True:
request = await encoder.schedule_socket.recv_pyobj()
request = await async_sock_recv(encoder.schedule_socket)
if isinstance(request, ProfileReq):
if request.type == ProfileReqType.START_PROFILE:
if encoder.profiler is None:
@@ -3601,7 +3609,7 @@ async def handle_encode_request(request: dict):
)
else:
for socket in send_sockets:
socket.send_pyobj(request)
sock_send(socket, request)
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
await encoder.encode_request(request, modality)
)
@@ -3822,7 +3830,7 @@ async def health_generate():
# Broadcast to other TP ranks so distributed ops stay in sync
for socket in send_sockets:
socket.send_pyobj(dummy_request)
sock_send(socket, dummy_request)
# Run encode on rank 0 with timeout
_, _, _, error_msg, _ = await asyncio.wait_for(
@@ -3900,7 +3908,7 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
profile_stages=obj.profile_stages,
)
for socket in send_sockets:
socket.send_pyobj(req)
sock_send(socket, req)
if encoder.profiler is None:
encoder.profiler = EncoderProfiler(encoder.rank)
ok, msg = encoder.profiler.start(req)
@@ -3931,7 +3939,7 @@ async def stop_profile_async():
)
req = ProfileReq(ProfileReqType.STOP_PROFILE)
for socket in send_sockets:
socket.send_pyobj(req)
sock_send(socket, req)
ok, msg = encoder.profiler.stop()
if ok:
return Response(content="Stop profiling.\n", status_code=200)
@@ -12,7 +12,7 @@ from sglang.srt.distributed.parallel_state import (
)
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.managers.io_struct import UpdateExpertBackupReq
from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_local_ip_auto
@@ -65,7 +65,7 @@ class ExpertBackupClient:
self.ready_sockets[i].connect(
f"tcp://{all_ips[i * get_world_size() // server_args.nnodes]}:{PORT_BASE + i * 2}"
)
self.ready_sockets[i].send_pyobj(UpdateExpertBackupReq())
sock_send(self.ready_sockets[i], UpdateExpertBackupReq())
self._receive_thread = threading.Thread(target=self._receive_loop, daemon=True)
self._receive_thread.start()
@@ -73,7 +73,7 @@ class ExpertBackupClient:
def _receive_loop(self):
cnt = 0
while cnt < self.engine_num:
response = self.recv_list[cnt].recv_pyobj()
response = sock_recv(self.recv_list[cnt])
self.dram_map_list[response.rank] = response.weight_pointer_map
self.session_id_list[response.rank] = response.session_id
self.buffer_size = max(self.buffer_size, response.buffer_size)
@@ -9,7 +9,7 @@ import zmq
from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import BackupDramReq
from sglang.srt.managers.io_struct import BackupDramReq, sock_recv, sock_send
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.server_args import (
@@ -62,7 +62,7 @@ class ExpertBackupManager:
num_ready_clients = 0
while num_ready_clients < server_args.tp_size:
self.recv_from_expert_backup_client.recv_pyobj()
sock_recv(self.recv_from_expert_backup_client)
num_ready_clients += 1
back_req = BackupDramReq(
@@ -72,7 +72,7 @@ class ExpertBackupManager:
buffer_size=self.continuous_buffer.numel()
* self.continuous_buffer.element_size(),
)
self.send_to_expert_backup_client.send_pyobj(back_req)
sock_send(self.send_to_expert_backup_client, back_req)
# Keep the manager subprocess alive until signals
signal.pause()
+4 -2
View File
@@ -77,6 +77,8 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromDistributedReqInput,
UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput,
sock_recv,
sock_send,
)
from sglang.srt.managers.multi_tokenizer_mixin import (
MultiTokenizerRouter,
@@ -1220,8 +1222,8 @@ class Engine(EngineScoreMixin, EngineBase):
def collective_rpc(self, method: str, **kwargs):
obj = RpcReqInput(method=method, parameters=kwargs)
self.send_to_rpc.send_pyobj(obj)
recv_req = self.send_to_rpc.recv_pyobj(zmq.BLOCKY)
sock_send(self.send_to_rpc, obj)
recv_req = sock_recv(self.send_to_rpc, flags=zmq.BLOCKY)
assert isinstance(recv_req, RpcReqOutput)
assert recv_req.success, recv_req.message
+4 -4
View File
@@ -5,7 +5,7 @@ import copy
from collections import deque
from typing import Deque, Generic, List, Optional, TypeVar
import zmq
from sglang.srt.managers.io_struct import sock_send
T = TypeVar("T")
@@ -22,7 +22,7 @@ class FanOutCommunicator(Generic[T]):
Only one request is in-flight at any time in either mode.
"""
def __init__(self, sender: zmq.Socket, fan_out: int, mode="queueing"):
def __init__(self, sender, fan_out: int, mode="queueing"):
self._sender = sender
self._fan_out = fan_out
self._mode = mode
@@ -41,7 +41,7 @@ class FanOutCommunicator(Generic[T]):
assert self._result_values is None
if obj is not None:
self._sender.send_pyobj(obj)
sock_send(self._sender, obj)
self._result_event = asyncio.Event()
self._result_values = []
@@ -61,7 +61,7 @@ class FanOutCommunicator(Generic[T]):
self._result_event = asyncio.Event()
if obj is not None:
self._sender.send_pyobj(obj)
sock_send(self._sender, obj)
# Capture local refs before await -- after event fires, the first
# awakened coroutine clears shared state; later awaiters use local refs.
@@ -36,6 +36,8 @@ from sglang.srt.managers.io_struct import (
ProfileReq,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
sock_recv,
sock_send,
)
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.schedule_batch import Req
@@ -205,12 +207,12 @@ class DataParallelController:
def send_to_all_workers(self, obj):
for i, worker in enumerate(self.workers):
if self.status[i]:
worker.send_pyobj(obj)
sock_send(worker, obj)
def send_control_message(self, obj):
# Send control messages to first worker of tp group
for worker in self.workers[:: self.control_message_step]:
worker.send_pyobj(obj)
sock_send(worker, obj)
def update_active_ranks(self, ranks: ActiveRanksOutput):
self.status = ranks.status
@@ -384,7 +386,7 @@ class DataParallelController:
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
rep_socket.send_pyobj(worker_ports)
sock_send(rep_socket, worker_ports)
connected_clients += 1
logger.debug(
f"Sent worker ports to {connected_clients}/{expected_clients} nodes"
@@ -418,7 +420,7 @@ class DataParallelController:
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
rep_socket.send_pyobj(worker_ports)
sock_send(rep_socket, worker_ports)
logger.debug(f"Sent worker ports to node {client_rank}")
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]:
@@ -434,7 +436,7 @@ class DataParallelController:
req_socket.send(str(node_rank).encode())
# Receive worker ports
worker_ports = req_socket.recv_pyobj()
worker_ports = sock_recv(req_socket)
logger.debug(f"Received {len(worker_ports)} worker ports from node 0")
return worker_ports
except zmq.Again:
@@ -597,7 +599,7 @@ class DataParallelController:
def maybe_external_dp_rank_routing(self, req: Req):
if req.routed_dp_rank is not None:
logger.debug(f"Direct routing to DP rank {req.routed_dp_rank}")
self.workers[req.routed_dp_rank].send_pyobj(req)
sock_send(self.workers[req.routed_dp_rank], req)
return True
return False
@@ -608,7 +610,7 @@ class DataParallelController:
while True:
if self.status[self.round_robin_counter]:
logger.debug(f"Choose worker {self.round_robin_counter}")
self.workers[self.round_robin_counter].send_pyobj(req)
sock_send(self.workers[self.round_robin_counter], req)
self.round_robin_counter = (self.round_robin_counter + 1) % len(
self.workers
)
@@ -626,13 +628,13 @@ class DataParallelController:
"prefill or decode instances; send to the router instead."
)
target_rank = req.bootstrap_room % len(self.workers)
self.workers[target_rank].send_pyobj(req)
sock_send(self.workers[target_rank], req)
def total_requests_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req):
return
target_worker = self.dp_budget.dispatch(LoadBalanceMethod.TOTAL_REQUESTS)
self.workers[target_worker].send_pyobj(req)
sock_send(self.workers[target_worker], req)
def total_tokens_scheduler(self, req: Req):
if self.maybe_external_dp_rank_routing(req):
@@ -641,14 +643,14 @@ class DataParallelController:
target_worker = self.dp_budget.dispatch(
LoadBalanceMethod.TOTAL_TOKENS, estimated_tokens=estimated_tokens
)
self.workers[target_worker].send_pyobj(req)
sock_send(self.workers[target_worker], req)
def event_loop(self):
while True:
while True:
self.soft_watchdog.feed()
try:
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
recv_req = sock_recv(self.recv_from_tokenizer, flags=zmq.NOBLOCK)
except zmq.ZMQError:
break
self._request_dispatcher(recv_req)
@@ -34,6 +34,8 @@ from sglang.srt.managers.io_struct import (
BatchTokenIDOutput,
ConfigureLoggingReq,
FreezeGCReq,
sock_recv,
sock_send,
)
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
@@ -160,10 +162,10 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
"""The event loop that handles requests"""
while True:
with self.soft_watchdog.disable():
recv_obj = self.recv_from_scheduler.recv_pyobj()
recv_obj = sock_recv(self.recv_from_scheduler)
output = self._request_dispatcher(recv_obj)
if output is not None:
self.send_to_tokenizer.send_pyobj(output)
sock_send(self.send_to_tokenizer, output)
self.soft_watchdog.feed()
def trim_matched_stop(
+38 -1
View File
@@ -25,9 +25,20 @@ from array import array
from collections import Counter
from dataclasses import dataclass, field
from enum import Enum
from typing import TYPE_CHECKING, Annotated, Any, Dict, List, Literal, Optional, Union
from typing import (
TYPE_CHECKING,
Annotated,
Any,
Dict,
List,
Literal,
Optional,
Union,
)
import torch
import zmq
import zmq.asyncio
from pydantic import PlainValidator
from sglang.srt.lora.lora_registry import LoRARef
@@ -46,6 +57,8 @@ 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
@@ -2191,6 +2204,30 @@ class DumperControlReqOutput(BaseReq):
error: str = ""
def sock_send(
sender: Union[zmq.Socket, zmq.asyncio.Socket, SenderWrapper],
obj: Any,
flags: int = 0,
) -> None:
sender.send_pyobj(obj, flags=flags)
def sock_recv(socket, flags=0):
return socket.recv_pyobj(flags=flags)
async def async_sock_send(
sender: Union[zmq.asyncio.Socket, SenderWrapper],
obj: Any,
flags: int = 0,
) -> None:
await sender.send_pyobj(obj, flags=flags)
async def async_sock_recv(socket, flags=0):
return await socket.recv_pyobj(flags=flags)
def _check_all_req_types():
"""A helper function to check all request types are defined in this file."""
import inspect
@@ -49,6 +49,10 @@ from sglang.srt.managers.io_struct import (
PauseContinueBroadcast,
PauseGenerationReqInput,
TokenizerWorkerRegistration,
async_sock_recv,
async_sock_send,
sock_recv,
sock_send,
)
from sglang.srt.managers.load_snapshot import (
create_load_snapshot_reader,
@@ -97,7 +101,7 @@ class SocketMapping:
if ipc_name not in self._mapping:
self._register_ipc_mapping(ipc_name, is_tokenizer=is_tokenizer)
self._mapping[ipc_name].send_pyobj(output)
sock_send(self._mapping[ipc_name], output)
def _extract_field_by_index(
@@ -334,7 +338,7 @@ class MultiHttpWorkerDetokenizerMixin:
"""The event loop that handles requests, for multi multi-http-worker mode"""
self.socket_mapping = SocketMapping()
while True:
recv_obj = self.recv_from_scheduler.recv_pyobj()
recv_obj = sock_recv(self.recv_from_scheduler)
output = self._request_dispatcher(recv_obj)
if output is None:
continue
@@ -431,7 +435,7 @@ class MultiTokenizerRouter:
async def router_worker_obj(self):
"""Forward path: workers → scheduler, with pause/continue broadcast."""
while True:
recv_obj = await self.receive_from_worker.recv_pyobj()
recv_obj = await async_sock_recv(self.receive_from_worker)
if isinstance(recv_obj, TokenizerWorkerRegistration):
if recv_obj.worker_ipc_name not in self.all_worker_ipcs:
@@ -456,15 +460,15 @@ class MultiTokenizerRouter:
isinstance(recv_obj, PauseGenerationReqInput)
and recv_obj.mode == "abort"
):
await self.send_to_scheduler.send_pyobj(recv_obj)
await async_sock_send(self.send_to_scheduler, recv_obj)
continue
await self.send_to_scheduler.send_pyobj(recv_obj)
await async_sock_send(self.send_to_scheduler, recv_obj)
async def handle_loop(self):
"""Backward path: detokenizer → route results to correct worker."""
while True:
recv_obj = await self.recv_from_detokenizer.recv_pyobj()
recv_obj = await async_sock_recv(self.recv_from_detokenizer)
await self._distribute_result_to_workers(recv_obj)
async def _distribute_result_to_workers(self, recv_obj):
@@ -505,7 +509,7 @@ class MultiDetokenizerRouter:
def event_loop(self):
while True:
recv_obj = self.recv_from_scheduler.recv_pyobj()
recv_obj = sock_recv(self.recv_from_scheduler)
# FreezeGCReq must freeze every detokenizer process.
if isinstance(recv_obj, FreezeGCReq):
@@ -600,7 +604,7 @@ class TokenizerWorker(TokenizerManager):
# Register this worker with the router for pause/continue broadcasting
reg = TokenizerWorkerRegistration(worker_ipc_name=self.tokenizer_ipc_name)
self.send_to_scheduler.send_pyobj(reg)
sock_send(self.send_to_scheduler, reg)
# Future for awaiting pause/continue broadcast confirmation
self._pause_continue_future: Optional[asyncio.Future] = None
@@ -618,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)
self.send_to_scheduler.send_pyobj(obj)
sock_send(self.send_to_scheduler, obj)
await self._pause_continue_future
if obj.mode == "abort":
@@ -633,7 +637,7 @@ class TokenizerWorker(TokenizerManager):
async def continue_generation(self, obj: ContinueGenerationReqInput):
loop = asyncio.get_event_loop()
self._pause_continue_future = loop.create_future()
self.send_to_scheduler.send_pyobj(obj)
sock_send(self.send_to_scheduler, obj)
await self._pause_continue_future
def _handle_pause_continue_broadcast(self, obj: PauseContinueBroadcast):
@@ -734,14 +738,3 @@ def write_data_for_multi_tokenizer(
args_shm.close()
return args_shm
class SenderWrapper:
def __init__(self, port_args: PortArgs, send_to_scheduler: zmq.Socket):
self.port_args = port_args
self.send_to_scheduler = send_to_scheduler
def send_pyobj(self, obj):
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
self.send_to_scheduler.send_pyobj(obj)
+2 -1
View File
@@ -147,6 +147,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromDistributedReqInput,
UpdateWeightsFromIPCReqInput,
UpdateWeightsFromTensorReqInput,
sock_send,
)
from sglang.srt.managers.load_snapshot import LoadSnapshot, create_load_snapshot_writer
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
@@ -1645,7 +1646,7 @@ class Scheduler(
self.ipc_channels.send_to_tokenizer.send_output(output, recv_req)
else:
if self.ipc_channels.recv_from_rpc is not None:
self.ipc_channels.recv_from_rpc.send_pyobj(output)
sock_send(self.ipc_channels.recv_from_rpc, output)
self.flush_wrapper.check_pending()
if self.external_corpus_manager is not None:
@@ -16,6 +16,7 @@ from sglang.srt.disaggregation.kv_events import (
EventPublisherFactory,
KVEventBatch,
)
from sglang.srt.managers.io_struct import sock_send
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
@@ -85,7 +86,7 @@ class SchedulerKvEventsPublisher:
)
if not self.send_metrics_from_scheduler.closed:
self.send_metrics_from_scheduler.send_pyobj(kv_metrics)
sock_send(self.send_metrics_from_scheduler, kv_metrics)
def publish_kv_events(self):
if not self.enable_kv_cache_events:
@@ -2,7 +2,7 @@ from typing import Optional, Union
import zmq
from sglang.srt.managers.io_struct import BaseBatchReq, BaseReq
from sglang.srt.managers.io_struct import BaseBatchReq, BaseReq, sock_send
class SenderWrapper:
@@ -25,4 +25,4 @@ class SenderWrapper:
# handle communicator reqs for multi-http worker case
output.http_worker_ipc = recv_obj.http_worker_ipc
self.socket.send_pyobj(output)
sock_send(self.socket, output)
@@ -20,6 +20,7 @@ from sglang.srt.managers.io_struct import (
BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
sock_recv,
)
from sglang.srt.managers.mm_utils import (
has_shm_features,
@@ -103,7 +104,7 @@ class SchedulerRequestReceiver:
try:
if self.recv_limit_reached(len(recv_reqs)):
break
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
recv_req = sock_recv(self.recv_from_tokenizer, zmq.NOBLOCK)
except zmq.ZMQError:
break
recv_reqs.append(recv_req)
@@ -112,7 +113,7 @@ class SchedulerRequestReceiver:
try:
if self.recv_limit_reached(len(recv_reqs)):
break
recv_rpc = self.recv_from_rpc.recv_pyobj(zmq.NOBLOCK)
recv_rpc = sock_recv(self.recv_from_rpc, zmq.NOBLOCK)
except zmq.ZMQError:
break
recv_reqs.append(recv_rpc)
@@ -16,7 +16,7 @@ from contextlib import contextmanager
from enum import Enum, auto
from typing import Any, List, Optional
from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType
from sglang.srt.managers.io_struct import BlockReqInput, BlockReqType, sock_send
from sglang.srt.utils.poll_based_barrier import PollBasedBarrier
logger = logging.getLogger(__name__)
@@ -99,8 +99,8 @@ class _State(Enum):
@contextmanager
def input_blocker_guard_region(send_to_scheduler):
send_to_scheduler.send_pyobj(BlockReqInput(BlockReqType.BLOCK))
sock_send(send_to_scheduler, BlockReqInput(BlockReqType.BLOCK))
try:
yield
finally:
send_to_scheduler.send_pyobj(BlockReqInput(BlockReqType.UNBLOCK))
sock_send(send_to_scheduler, BlockReqInput(BlockReqType.UNBLOCK))
@@ -71,6 +71,8 @@ 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
@@ -857,7 +859,7 @@ class TokenizerControlMixin:
future = asyncio.Future()
self.session_futures[obj.session_id] = future
self.send_to_scheduler.send_pyobj(obj)
sock_send(self.send_to_scheduler, obj)
try:
return await future
@@ -869,7 +871,7 @@ class TokenizerControlMixin:
obj: CloseSessionReqInput,
request: Optional[fastapi.Request] = None,
):
await self.send_to_scheduler.send_pyobj(obj)
await async_sock_send(self.send_to_scheduler, obj)
def _update_weight_version_if_provided(
self: TokenizerManager, weight_version: Optional[str]
+42 -16
View File
@@ -55,6 +55,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
from sglang.srt.managers.io_struct import (
AbortReq,
ActiveRanksOutput,
BaseReq,
BatchEmbeddingOutput,
BatchStrOutput,
BatchTokenIDOutput,
@@ -75,6 +76,9 @@ from sglang.srt.managers.io_struct import (
TokenizedGenerateReqInput,
UpdateWeightFromDiskReqInput,
UpdateWeightFromDiskReqOutput,
async_sock_recv,
async_sock_send,
sock_send,
)
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
@@ -379,19 +383,19 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
context, zmq.PULL, port_args.tokenizer_ipc_name, True
)
if self.server_args.tokenizer_worker_num == 1:
self.send_to_scheduler = get_zmq_socket(
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)
else:
from sglang.srt.managers.multi_tokenizer_mixin import SenderWrapper
# Use tokenizer_worker_ipc_name in multi-tokenizer mode
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)
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.server_args,
@@ -1322,7 +1326,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
tokenized_obj.time_stats.set_api_server_dispatch_time()
tokenized_obj = wrap_shm_features(tokenized_obj)
self.send_to_scheduler.send_pyobj(tokenized_obj)
sock_send(self.send_to_scheduler, tokenized_obj)
tokenized_obj.time_stats.set_api_server_dispatch_finish_time()
def _send_batch_request(
@@ -1338,7 +1342,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
set_time_batch(tokenized_objs, "set_api_server_dispatch_time")
self.send_to_scheduler.send_pyobj(batch_req)
sock_send(self.send_to_scheduler, batch_req)
set_time_batch(tokenized_objs, "set_api_server_dispatch_finish_time")
def _coalesce_streaming_chunks(
@@ -1663,7 +1667,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
return
req = AbortReq(rid=rid, abort_all=abort_all)
self.send_to_scheduler.send_pyobj(req)
sock_send(self.send_to_scheduler, req)
if self.enable_metrics:
# TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request(
@@ -1674,7 +1678,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async with self.is_pause_cond:
self.is_pause = True
if obj.mode != "abort":
await self.send_to_scheduler.send_pyobj(obj)
await async_sock_send(self.send_to_scheduler, obj)
else:
# we are using the model_update_lock to check if there is still on-going requests.
while True:
@@ -1688,7 +1692,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def continue_generation(self, obj: ContinueGenerationReqInput):
async with self.is_pause_cond:
self.is_pause = False
await self.send_to_scheduler.send_pyobj(obj)
await async_sock_send(self.send_to_scheduler, obj)
self.is_pause_cond.notify_all()
async def update_weights_from_disk(
@@ -1733,7 +1737,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]:
self.send_to_scheduler.send_pyobj(obj)
sock_send(self.send_to_scheduler, obj)
self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1:
result = await self.model_update_result
@@ -1773,12 +1777,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())
self.send_to_scheduler.send_pyobj(obj)
sock_send(self.send_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."""
self.send_to_scheduler.send_pyobj(FreezeGCReq())
sock_send(self.send_to_scheduler, FreezeGCReq())
freeze_gc("Tokenizer Manager")
return None
@@ -1825,7 +1829,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"""The event loop that handles requests"""
while True:
with self.soft_watchdog.disable():
recv_obj = await self.recv_from_detokenizer.recv_pyobj()
recv_obj = await async_sock_recv(self.recv_from_detokenizer)
if isinstance(
recv_obj,
(BatchStrOutput, BatchEmbeddingOutput, BatchTokenIDOutput),
@@ -2648,7 +2652,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.
self.send_to_scheduler.send_pyobj(ShutdownReq())
sock_send(self.send_to_scheduler, ShutdownReq())
deadline = time.monotonic() + 15
while time.monotonic() < deadline and collect_scheduler_processes():
time.sleep(0.1)
@@ -2717,7 +2721,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.event.set()
def update_active_ranks(self, ranks: ActiveRanksOutput):
self.send_to_scheduler.send_pyobj(ranks)
sock_send(self.send_to_scheduler, ranks)
def _handle_open_session_req_output(self, recv_obj):
future = self.session_futures.get(recv_obj.session_id)
@@ -3115,3 +3119,25 @@ class SignalHandler:
# | http | no | waiting queue | type 1 | type 1 exception | del in _handle_abort_req |
# | http | no | running | type 3 | type 3 exception | del in _handle_batch_output |
#
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)
@@ -12,6 +12,7 @@ import zmq
from sglang.srt.entrypoints.http_server import launch_server
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.network import get_free_port, get_zmq_socket_on_host
from sglang.test.scripted_runtime.io_struct import (
@@ -89,7 +90,7 @@ class ScriptedHttpServer:
raise RuntimeError(f"ScriptedHttpServer is dirty: {self._dirty}")
fn_path = f"{script_fn.__module__}:{script_fn.__qualname__}"
self._socket.send_pyobj(RunScript(fn_path=fn_path, args=args))
sock_send(self._socket, RunScript(fn_path=fn_path, args=args))
if not self._socket.poll(int(timeout_s * 1000)):
if not self._server_process.is_alive():
@@ -98,7 +99,7 @@ class ScriptedHttpServer:
self._dirty = f"script {fn_path!r} timed out after {timeout_s}s"
raise TimeoutError(self._dirty)
reply = self._socket.recv_pyobj()
reply = sock_recv(self._socket)
match reply:
case ScriptFailed(traceback=tb):
raise AssertionError(f"scripted-runtime script failed:\n{tb}")
@@ -116,7 +117,7 @@ class ScriptedHttpServer:
fatal_error: Optional[OutOfBandError] = None
try:
try:
self._socket.send_pyobj(Shutdown())
sock_send(self._socket, Shutdown())
except zmq.ZMQError:
pass
@@ -139,7 +140,7 @@ class ScriptedHttpServer:
f"{LISTENER_ACCEPT_TIMEOUT_S}s"
)
ready = self._socket.recv_pyobj()
ready = sock_recv(self._socket)
if not isinstance(ready, HookReady):
raise RuntimeError(
f"ScriptedHttpServer: expected HookReady handshake, got {ready!r}"
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Generator, List, Optional, Tuple
import zmq
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send
from sglang.srt.utils.network import get_zmq_socket
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.scripted_runtime.context import ScriptedContext
@@ -152,9 +153,9 @@ class ScriptedSchedulerHook:
socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
try:
yield from _drive_engine_through_warmup(self._context)
socket.send_pyobj(HookReady())
sock_send(socket, HookReady())
while True:
msg = socket.recv_pyobj()
msg = sock_recv(socket)
match msg:
case Shutdown():
return
@@ -167,11 +168,12 @@ class ScriptedSchedulerHook:
try:
yield from sub_gen
except Exception:
socket.send_pyobj(
ScriptFailed(traceback=traceback.format_exc())
sock_send(
socket,
ScriptFailed(traceback=traceback.format_exc()),
)
else:
socket.send_pyobj(ScriptSucceeded())
sock_send(socket, ScriptSucceeded())
case _:
raise ValueError(f"dispatch loop: unknown command {msg!r}")
finally:
@@ -11,6 +11,7 @@ from sglang.srt.managers.io_struct import (
BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
sock_recv,
)
_WORK_REQ_TYPES = (
@@ -64,7 +65,7 @@ class ScriptedTokenizerRecvProxy:
def _drain_underlying(self) -> None:
while True:
try:
req = self._underlying.recv_pyobj(zmq.NOBLOCK)
req = sock_recv(self._underlying, zmq.NOBLOCK)
except zmq.ZMQError:
break
if isinstance(req, _WORK_REQ_TYPES):