Skip to content

Commit 34dd9c2

Browse files
authored
[Refactor] Introduce sock_send/sock_recv wrappers for zmq IPC (#29012)
1 parent ecab3f3 commit 34dd9c2

22 files changed

Lines changed: 258 additions & 154 deletions

python/sglang/srt/debug_utils/dumper.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,9 @@
2020

2121
import torch
2222
import torch.distributed as dist
23+
import zmq
24+
25+
from sglang.srt.managers.io_struct import sock_recv, sock_send
2326

2427
# -------------------------------------- config base ------------------------------------------
2528

@@ -1419,8 +1422,6 @@ def _create_zmq_rpc_broadcast(
14191422
handler, timeout_seconds: int = 60
14201423
) -> Optional["_ZmqRpcBroadcast"]:
14211424
"""A general-purpose minimal RPC to support broadcasting executions to multi processes"""
1422-
import zmq
1423-
14241425
rank = _get_rank()
14251426
world_size = dist.get_world_size() if dist.is_initialized() else 1
14261427

@@ -1433,13 +1434,13 @@ def _create_zmq_rpc_broadcast(
14331434
def serve_loop():
14341435
while True:
14351436
try:
1436-
req = sock.recv_pyobj()
1437+
req = sock_recv(sock)
14371438
result = getattr(handler, req["method"])(*req["args"], **req["kwargs"])
14381439
resp = {"result": result, "error": None}
14391440
except Exception as e:
14401441
_log(f"[ZmqRpc] error inside handler: {e}")
14411442
resp = {"result": None, "error": str(e)}
1442-
sock.send_pyobj(resp)
1443+
sock_send(sock, resp)
14431444

14441445
thread = threading.Thread(target=serve_loop, daemon=True)
14451446
thread.start()
@@ -1476,14 +1477,15 @@ def __init__(self, socket, debug_name: str):
14761477

14771478
def __getattr__(self, method_name: str):
14781479
def call(*args, **kwargs):
1479-
self._socket.send_pyobj(
1480+
sock_send(
1481+
self._socket,
14801482
{
14811483
"method": method_name,
14821484
"args": args,
14831485
"kwargs": kwargs,
1484-
}
1486+
},
14851487
)
1486-
response = self._socket.recv_pyobj()
1488+
response = sock_recv(self._socket)
14871489
if response["error"]:
14881490
raise RuntimeError(
14891491
f"RPC error on {self._debug_name}: {response['error']}"

python/sglang/srt/disaggregation/encode_grpc_server.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
handle_scheduler_receive_url_request,
2727
launch_encoder,
2828
)
29+
from sglang.srt.managers.io_struct import async_sock_send
2930
from sglang.srt.managers.schedule_batch import Modality
3031
from sglang.srt.server_args import PortArgs, ServerArgs
3132
from sglang.srt.utils import random_uuid
@@ -95,7 +96,7 @@ async def Encode(
9596
"part_idx": request.part_idx,
9697
}
9798
for socket in self.send_sockets:
98-
await socket.send_pyobj(request_dict)
99+
await async_sock_send(socket, request_dict)
99100

100101
# gRPC encode is image-only; encoder.encode() requires modality
101102
(

python/sglang/srt/disaggregation/encode_server.py

Lines changed: 24 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,14 @@
4444
)
4545
from sglang.srt.environ import envs
4646
from sglang.srt.layers.dp_attention import initialize_dp_attention
47-
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileReqType
47+
from sglang.srt.managers.io_struct import (
48+
ProfileReq,
49+
ProfileReqInput,
50+
ProfileReqType,
51+
async_sock_recv,
52+
async_sock_send,
53+
sock_send,
54+
)
4855
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
4956
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
5057
from sglang.srt.model_loader import get_model
@@ -2390,13 +2397,14 @@ async def _dispatch_group(
23902397
requests = [p.request for p in group]
23912398
start = time.time()
23922399
for sock in self.send_sockets:
2393-
sock.send_pyobj(
2400+
sock_send(
2401+
sock,
23942402
{
23952403
"type": "batch_encode",
23962404
"modality": modality.name,
23972405
"requests": requests,
23982406
"enter_time": start,
2399-
}
2407+
},
24002408
)
24012409

24022410
logger.info(f"Dispatching batch of {len(group)} {modality.name} requests")
@@ -2442,7 +2450,7 @@ async def _dispatch_per_request(
24422450
req = p.request
24432451
try:
24442452
for sock in self.send_sockets:
2445-
sock.send_pyobj(req)
2453+
sock_send(sock, req)
24462454
result = await self.encoder.encode_request(req, modality)
24472455
if not p.future.done():
24482456
p.future.set_result(result)
@@ -2731,7 +2739,7 @@ async def dispatch(self, request: dict) -> dict:
27312739
)
27322740

27332741
try:
2734-
await self.dispatch_sockets[rank].send_pyobj(request)
2742+
await async_sock_send(self.dispatch_sockets[rank], request)
27352743
# An alive-but-stuck worker (NCCL deadlock etc.) wouldn't trip
27362744
# the watchdog, so bound the wait explicitly.
27372745
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
@@ -2780,7 +2788,7 @@ async def dispatch_send(self, request: dict) -> dict:
27802788
f"dp_rank={rank}, pending={self.pending_counts}"
27812789
)
27822790
try:
2783-
await self.dispatch_sockets[rank].send_pyobj(request)
2791+
await async_sock_send(self.dispatch_sockets[rank], request)
27842792
return await asyncio.wait_for(future, timeout=ENCODER_REQ_TIMEOUT)
27852793
except asyncio.TimeoutError:
27862794
self.pending_futures[rank].pop(key, None)
@@ -2821,7 +2829,7 @@ async def broadcast(
28212829
self.req_id_to_rank[req_id] = rank
28222830
rank_keys.append((rank, req_id))
28232831
request_copy = {**request, "req_id": req_id}
2824-
await self.dispatch_sockets[rank].send_pyobj(request_copy)
2832+
await async_sock_send(self.dispatch_sockets[rank], request_copy)
28252833
futures.append(future)
28262834
# Concurrent wait → total bounded by eff_timeout, not
28272835
# dp_size × eff_timeout.
@@ -2901,7 +2909,7 @@ async def _result_listener(self) -> None:
29012909
consecutive_errors = 0
29022910
while True:
29032911
try:
2904-
msg = await self.result_socket.recv_pyobj()
2912+
msg = await async_sock_recv(self.result_socket)
29052913
consecutive_errors = 0
29062914
except asyncio.CancelledError:
29072915
raise
@@ -3051,10 +3059,10 @@ async def _dp_worker_handle_request(
30513059
"_error_code": err_code,
30523060
}
30533061

3054-
# pyzmq async send_pyobj isn't safe for concurrent senders.
3062+
# pyzmq async send isn't safe for concurrent senders.
30553063
try:
30563064
async with send_lock:
3057-
await send_sock.send_pyobj(envelope)
3065+
await async_sock_send(send_sock, envelope)
30583066
except Exception:
30593067
logger.error(
30603068
f"DP worker {dp_rank} failed to send envelope for "
@@ -3113,7 +3121,7 @@ async def run_dp_worker(
31133121
spawned = False
31143122
try:
31153123
try:
3116-
request = await recv_sock.recv_pyobj()
3124+
request = await async_sock_recv(recv_sock)
31173125
except asyncio.CancelledError:
31183126
raise
31193127
except Exception:
@@ -3197,7 +3205,7 @@ async def run_encoder(
31973205
):
31983206
encoder = MMEncoder(server_args, schedule_path, dist_init_method, rank)
31993207
while True:
3200-
request = await encoder.schedule_socket.recv_pyobj()
3208+
request = await async_sock_recv(encoder.schedule_socket)
32013209
if isinstance(request, ProfileReq):
32023210
if request.type == ProfileReqType.START_PROFILE:
32033211
if encoder.profiler is None:
@@ -3601,7 +3609,7 @@ def start_background_send(req_id):
36013609
)
36023610
else:
36033611
for socket in send_sockets:
3604-
socket.send_pyobj(request)
3612+
sock_send(socket, request)
36053613
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
36063614
await encoder.encode_request(request, modality)
36073615
)
@@ -3822,7 +3830,7 @@ async def health_generate():
38223830

38233831
# Broadcast to other TP ranks so distributed ops stay in sync
38243832
for socket in send_sockets:
3825-
socket.send_pyobj(dummy_request)
3833+
sock_send(socket, dummy_request)
38263834

38273835
# Run encode on rank 0 with timeout
38283836
_, _, _, error_msg, _ = await asyncio.wait_for(
@@ -3900,7 +3908,7 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
39003908
profile_stages=obj.profile_stages,
39013909
)
39023910
for socket in send_sockets:
3903-
socket.send_pyobj(req)
3911+
sock_send(socket, req)
39043912
if encoder.profiler is None:
39053913
encoder.profiler = EncoderProfiler(encoder.rank)
39063914
ok, msg = encoder.profiler.start(req)
@@ -3931,7 +3939,7 @@ async def stop_profile_async():
39313939
)
39323940
req = ProfileReq(ProfileReqType.STOP_PROFILE)
39333941
for socket in send_sockets:
3934-
socket.send_pyobj(req)
3942+
sock_send(socket, req)
39353943
ok, msg = encoder.profiler.stop()
39363944
if ok:
39373945
return Response(content="Stop profiling.\n", status_code=200)

python/sglang/srt/elastic_ep/expert_backup_client.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
)
1313
from sglang.srt.environ import envs
1414
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
15-
from sglang.srt.managers.io_struct import UpdateExpertBackupReq
15+
from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
1616
from sglang.srt.server_args import ServerArgs
1717
from sglang.srt.utils.network import get_local_ip_auto
1818

@@ -65,15 +65,15 @@ def __init__(self, server_args: ServerArgs, model_runner):
6565
self.ready_sockets[i].connect(
6666
f"tcp://{all_ips[i * get_world_size() // server_args.nnodes]}:{PORT_BASE + i * 2}"
6767
)
68-
self.ready_sockets[i].send_pyobj(UpdateExpertBackupReq())
68+
sock_send(self.ready_sockets[i], UpdateExpertBackupReq())
6969

7070
self._receive_thread = threading.Thread(target=self._receive_loop, daemon=True)
7171
self._receive_thread.start()
7272

7373
def _receive_loop(self):
7474
cnt = 0
7575
while cnt < self.engine_num:
76-
response = self.recv_list[cnt].recv_pyobj()
76+
response = sock_recv(self.recv_list[cnt])
7777
self.dram_map_list[response.rank] = response.weight_pointer_map
7878
self.session_id_list[response.rank] = response.session_id
7979
self.buffer_size = max(self.buffer_size, response.buffer_size)

python/sglang/srt/elastic_ep/expert_backup_manager.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
from sglang.srt.configs.load_config import LoadConfig
1010
from sglang.srt.configs.model_config import ModelConfig
1111
from sglang.srt.environ import envs
12-
from sglang.srt.managers.io_struct import BackupDramReq
12+
from sglang.srt.managers.io_struct import BackupDramReq, sock_recv, sock_send
1313
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
1414
from sglang.srt.model_loader.utils import set_default_torch_dtype
1515
from sglang.srt.server_args import (
@@ -62,7 +62,7 @@ def __init__(self, server_args: ServerArgs, port_args: PortArgs):
6262
num_ready_clients = 0
6363

6464
while num_ready_clients < server_args.tp_size:
65-
self.recv_from_expert_backup_client.recv_pyobj()
65+
sock_recv(self.recv_from_expert_backup_client)
6666
num_ready_clients += 1
6767

6868
back_req = BackupDramReq(
@@ -72,7 +72,7 @@ def __init__(self, server_args: ServerArgs, port_args: PortArgs):
7272
buffer_size=self.continuous_buffer.numel()
7373
* self.continuous_buffer.element_size(),
7474
)
75-
self.send_to_expert_backup_client.send_pyobj(back_req)
75+
sock_send(self.send_to_expert_backup_client, back_req)
7676

7777
# Keep the manager subprocess alive until signals
7878
signal.pause()

python/sglang/srt/entrypoints/engine.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,8 @@
7777
UpdateWeightsFromDistributedReqInput,
7878
UpdateWeightsFromIPCReqInput,
7979
UpdateWeightsFromTensorReqInput,
80+
sock_recv,
81+
sock_send,
8082
)
8183
from sglang.srt.managers.multi_tokenizer_mixin import (
8284
MultiTokenizerRouter,
@@ -1220,8 +1222,8 @@ def freeze_gc(self):
12201222

12211223
def collective_rpc(self, method: str, **kwargs):
12221224
obj = RpcReqInput(method=method, parameters=kwargs)
1223-
self.send_to_rpc.send_pyobj(obj)
1224-
recv_req = self.send_to_rpc.recv_pyobj(zmq.BLOCKY)
1225+
sock_send(self.send_to_rpc, obj)
1226+
recv_req = sock_recv(self.send_to_rpc, flags=zmq.BLOCKY)
12251227
assert isinstance(recv_req, RpcReqOutput)
12261228
assert recv_req.success, recv_req.message
12271229

python/sglang/srt/managers/communicator.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
from collections import deque
66
from typing import Deque, Generic, List, Optional, TypeVar
77

8-
import zmq
8+
from sglang.srt.managers.io_struct import sock_send
99

1010
T = TypeVar("T")
1111

@@ -22,7 +22,7 @@ class FanOutCommunicator(Generic[T]):
2222
Only one request is in-flight at any time in either mode.
2323
"""
2424

25-
def __init__(self, sender: zmq.Socket, fan_out: int, mode="queueing"):
25+
def __init__(self, sender, fan_out: int, mode="queueing"):
2626
self._sender = sender
2727
self._fan_out = fan_out
2828
self._mode = mode
@@ -41,7 +41,7 @@ async def queueing_call(self, obj: T):
4141
assert self._result_values is None
4242

4343
if obj is not None:
44-
self._sender.send_pyobj(obj)
44+
sock_send(self._sender, obj)
4545

4646
self._result_event = asyncio.Event()
4747
self._result_values = []
@@ -61,7 +61,7 @@ async def watching_call(self, obj):
6161
self._result_event = asyncio.Event()
6262

6363
if obj is not None:
64-
self._sender.send_pyobj(obj)
64+
sock_send(self._sender, obj)
6565

6666
# Capture local refs before await -- after event fires, the first
6767
# awakened coroutine clears shared state; later awaiters use local refs.

0 commit comments

Comments
 (0)