[FEAT] optimize tensor zmq transfer for multimodal inputs (#13592)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com>
This commit is contained in:
siyu
2026-01-03 12:23:23 +08:00
committed by GitHub
parent 8b111b20c3
commit 078d96213a
5 changed files with 271 additions and 23 deletions

View File

@@ -16,6 +16,7 @@
import faulthandler
import logging
import multiprocessing as mp
import pickle
import signal
import threading
import time
@@ -25,6 +26,7 @@ from typing import Callable, List, Optional
import psutil
import setproctitle
import torch
import zmq
from sglang.srt.environ import envs
@@ -208,13 +210,15 @@ class DataParallelController:
)
def send_to_all_workers(self, obj):
msg = [b"NORM", pickle.dumps(obj)]
for worker in self.workers:
worker.send_pyobj(obj)
worker.send_multipart(msg, copy=False)
def send_control_message(self, obj):
msg = [b"NORM", pickle.dumps(obj)]
# Send control messages to first worker of tp group
for worker in self.workers[:: self.control_message_step]:
worker.send_pyobj(obj)
worker.send_multipart(msg, copy=False)
def handle_load_update_req(self, obj):
self.dp_budget.update_budget(obj)
@@ -499,8 +503,9 @@ class DataParallelController:
def maybe_external_dp_rank_routing(self, req: Req):
if req.data_parallel_rank is not None:
msg = [b"NORM", pickle.dumps(req)]
logger.debug(f"Direct routing to DP rank {req.data_parallel_rank}")
self.workers[req.data_parallel_rank].send_pyobj(req)
self.workers[req.data_parallel_rank].send_multipart(msg, copy=False)
return True
return False
@@ -508,7 +513,8 @@ class DataParallelController:
if self.maybe_external_dp_rank_routing(req):
return
self.workers[self.round_robin_counter].send_pyobj(req)
msg = [b"NORM", pickle.dumps(req)]
self.workers[self.round_robin_counter].send_multipart(msg, copy=False)
self.round_robin_counter = (self.round_robin_counter + 1) % len(self.workers)
def follow_bootstrap_room_scheduler(self, req: Req):
@@ -530,7 +536,8 @@ 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)
msg = [b"NORM", pickle.dumps(req)]
self.workers[target_rank].send_multipart(msg, copy=False)
def shortest_queue_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req):
@@ -542,7 +549,8 @@ class DataParallelController:
else:
self.follow_bootstrap_room_scheduler(req)
else:
self.workers[target_worker].send_pyobj(req)
msg = [b"NORM", pickle.dumps(req)]
self.workers[target_worker].send_multipart(msg, copy=False)
def minimum_tokens_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req):
@@ -557,12 +565,87 @@ class DataParallelController:
else:
self.follow_bootstrap_room_scheduler(req)
def _parse_multipart_message(self, parts):
# Check message type
msg_type = bytes(parts[0])
if msg_type == b"NORM":
# Normal message
recv_req = pickle.loads(parts[1])
elif msg_type == b"FEAT":
# Message with optimized feature tensors
recv_req = pickle.loads(parts[1])
feature_infos = pickle.loads(parts[2])
# Reconstruct tensors
for i, feature_info in enumerate(feature_infos):
buffer_idx = 3 + i
buffer = (
parts[buffer_idx].buffer
if hasattr(parts[buffer_idx], "buffer")
else parts[buffer_idx]
)
dtype = feature_info["dtype"]
shape = feature_info["shape"]
tensor = torch.frombuffer(buffer, dtype=dtype).reshape(shape)
idx = feature_info["idx"]
if hasattr(recv_req, "mm_inputs") and recv_req.mm_inputs:
mm_items = recv_req.mm_inputs.get("mm_items", [])
if idx < len(mm_items):
mm_items[idx].feature = tensor
else:
logger.warning(f"Unknown message type: {msg_type}")
return recv_req
def _parse_multipart_message(self, parts):
# Check message type
msg_type = bytes(parts[0])
if msg_type == b"NORM":
# Normal message
recv_req = pickle.loads(parts[1])
elif msg_type == b"FEAT":
# Message with optimized feature tensors
recv_req = pickle.loads(parts[1])
feature_infos = pickle.loads(parts[2])
# Reconstruct tensors
for i, feature_info in enumerate(feature_infos):
buffer_idx = 3 + i
buffer = (
parts[buffer_idx].buffer
if hasattr(parts[buffer_idx], "buffer")
else parts[buffer_idx]
)
dtype = feature_info["dtype"]
shape = feature_info["shape"]
tensor = torch.frombuffer(buffer, dtype=dtype).reshape(shape)
idx = feature_info["idx"]
if hasattr(recv_req, "mm_inputs") and recv_req.mm_inputs:
mm_items = recv_req.mm_inputs.get("mm_items", [])
if idx < len(mm_items):
mm_items[idx].feature = tensor
else:
logger.warning(f"Unknown message type: {msg_type}")
return recv_req
def event_loop(self):
while True:
while True:
self.soft_watchdog.feed()
try:
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
parts = self.recv_from_tokenizer.recv_multipart(
flags=zmq.NOBLOCK, copy=False
)
if not parts:
break
recv_req = self._parse_multipart_message(parts)
except zmq.ZMQError:
break
self._request_dispatcher(recv_req)