[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)

View File

@@ -369,8 +369,8 @@ class MultiTokenizerRouter:
async def router_worker_obj(self):
while True:
recv_obj = await self.receive_from_worker.recv_pyobj()
await self.send_to_scheduler.send_pyobj(recv_obj)
parts = await self.receive_from_worker.recv_multipart(copy=False)
await self.send_to_scheduler.send_multipart(parts, copy=False)
async def handle_loop(self):
# special reqs will recv from scheduler, need to route to right worker
@@ -525,3 +525,10 @@ class SenderWrapper:
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
self.send_to_scheduler.send_pyobj(obj)
def send_multipart(self, parts, copy=False):
obj = pickle.loads(parts[1])
if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
parts = [parts[0], pickle.dumps(obj)] + list(parts[2:])
self.send_to_scheduler.send_multipart(parts, copy=copy)

View File

@@ -16,6 +16,7 @@
import faulthandler
import logging
import os
import pickle
import signal
import sys
import time
@@ -1184,6 +1185,41 @@ class Scheduler(
return False
return num_recv_reqs >= self.max_recv_per_poll
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 recv_requests(
self,
) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]:
@@ -1204,10 +1240,15 @@ class Scheduler(
try:
if self.recv_limit_reached(len(recv_reqs)):
break
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
except zmq.ZMQError:
parts = self.recv_from_tokenizer.recv_multipart(
flags=zmq.NOBLOCK, copy=False
)
if not parts:
break
recv_req = self._parse_multipart_message(parts)
recv_reqs.append(recv_req)
except zmq.ZMQError as e:
break
recv_reqs.append(recv_req)
while True:
try:

View File

@@ -3,6 +3,7 @@ from __future__ import annotations
import asyncio
import copy
import logging
import pickle
import time
import uuid
from collections import deque
@@ -105,7 +106,12 @@ class _Communicator(Generic[T]):
assert self._result_values is None
if obj:
self._sender.send_pyobj(obj)
self._sender.send_multipart(
[
b"NORM",
pickle.dumps(obj),
]
)
self._result_event = asyncio.Event()
self._result_values = []
@@ -125,7 +131,12 @@ class _Communicator(Generic[T]):
self._result_event = asyncio.Event()
if obj:
self._sender.send_pyobj(obj)
self._sender.send_multipart(
[
b"NORM",
pickle.dumps(obj),
]
)
await self._result_event.wait()
result_values = copy.deepcopy(self._result_values)

View File

@@ -15,6 +15,7 @@
import asyncio
import copy
import ctypes
import dataclasses
import logging
import os
@@ -32,6 +33,7 @@ from http import HTTPStatus
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
import fastapi
import torch
import uvloop
import zmq
import zmq.asyncio
@@ -119,6 +121,29 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
logger = logging.getLogger(__name__)
class TensorWrapper:
"""Wrapper to keep tensor alive while exposing buffer for zero-copy."""
def __init__(self, tensor):
# Ensure tensor is on CPU and contiguous
if tensor.is_cuda:
tensor = tensor.cpu()
if not tensor.is_contiguous():
tensor = tensor.contiguous()
# Keep tensor reference
self.tensor = tensor
self.shape = list(tensor.shape)
self.dtype = tensor.dtype
def __buffer__(self):
data_ptr = self.tensor.data_ptr()
total_bytes = self.tensor.numel() * self.tensor.element_size()
c_obj = (ctypes.c_char * total_bytes).from_address(data_ptr)
c_obj._keep_alive_ref = self
return memoryview(c_obj)
@dataclasses.dataclass
class ReqState:
"""Store the state a request."""
@@ -1029,6 +1054,88 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
)
)
@staticmethod
def extract_feature_tensors(tokenized_obj):
if not isinstance(tokenized_obj, TokenizedGenerateReqInput):
return False, None, None
has_feature_tensors = False
feature_wrappers = []
feature_infos = []
if hasattr(tokenized_obj, "mm_inputs") and tokenized_obj.mm_inputs:
mm_items = tokenized_obj.mm_inputs.get("mm_items", [])
for idx, item in enumerate(mm_items):
if (
hasattr(item, "feature")
and item.feature is not None
and isinstance(item.feature, torch.Tensor)
):
has_feature_tensors = True
# Create wrapper (handles CPU/contiguous conversion and keeps tensor alive)
wrapper = TensorWrapper(item.feature)
feature_wrappers.append(wrapper)
# Store metadata (from wrapper for consistency)
feature_info = {
"idx": idx,
"shape": wrapper.shape,
"dtype": wrapper.dtype,
}
feature_infos.append(feature_info)
# Clear original reference for pickling
item.feature = None
return has_feature_tensors, feature_wrappers, feature_infos
def _send_multi_parts(self, sender, obj, copy=False):
has_feature_tensors = False
feature_wrappers = None
feature_infos = None
if not self.server_args.skip_tokenizer_init:
has_feature_tensors, feature_wrappers, feature_infos = (
TokenizerManager.extract_feature_tensors(obj)
)
if has_feature_tensors:
parts = [
b"FEAT",
pickle.dumps(obj),
pickle.dumps(feature_infos),
]
# Add wrappers - they keep tensors alive and provide buffer interface
for wrapper in feature_wrappers:
parts.append(wrapper.__buffer__())
sender.send_multipart(parts, copy=copy)
else:
sender.send_multipart(
[
b"NORM",
pickle.dumps(obj),
],
copy=False,
)
async def _send_multi_parts_async(self, sender, obj, copy=False):
has_feature_tensors = False
feature_wrappers = None
feature_infos = None
if not self.server_args.skip_tokenizer_init:
has_feature_tensors, feature_wrappers, feature_infos = (
TokenizerManager.extract_feature_tensors(obj)
)
if has_feature_tensors:
parts = [b"FEAT", pickle.dumps(obj), pickle.dumps(feature_infos)]
for wrapper in feature_wrappers:
parts.append(wrapper.__buffer__())
await sender.send_multipart(parts, copy=copy)
else:
await sender.send_multipart([b"NORM", pickle.dumps(obj)], copy=copy)
def _send_one_request(
self,
obj: Union[GenerateReqInput, EmbeddingReqInput],
@@ -1037,7 +1144,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
):
trace_slice_start(RequestStage.TOKENIZER_DISPATCH, obj.rid)
tokenized_obj.trace_context = trace_get_proc_propagate_context(obj.rid)
self.send_to_scheduler.send_pyobj(tokenized_obj)
self._send_multi_parts(self.send_to_scheduler, tokenized_obj)
state = ReqState([], False, asyncio.Event(), obj, created_time=created_time)
state.request_sent_to_scheduler_ts = time.time()
self.rid_to_state[obj.rid] = state
@@ -1059,8 +1166,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
self.send_to_scheduler.send_pyobj(batch_req)
self._send_multi_parts(self.send_to_scheduler, batch_req)
# Create states for each individual request in the batch
for i, tokenized_obj in enumerate(tokenized_objs):
tmp_obj = obj[i]
@@ -1283,7 +1389,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if not abort_all and rid not in self.rid_to_state:
return
req = AbortReq(rid=rid, abort_all=abort_all)
self.send_to_scheduler.send_pyobj(req)
self._send_multi_parts(self.send_to_scheduler, req)
if self.enable_metrics:
# TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request(
@@ -1294,7 +1400,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async with self.is_pause_cond:
self.is_pause = True
if obj.mode != "abort":
await self.send_to_scheduler.send_pyobj(obj)
await self._send_multi_parts_async(self.send_to_scheduler, obj)
else:
# we are using the model_update_lock to check if there is still on-going requests.
while True:
@@ -1308,7 +1414,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
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 self._send_multi_parts_async(self.send_to_scheduler, obj)
self.is_pause_cond.notify_all()
async def update_weights_from_disk(
@@ -1353,7 +1459,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]:
self.send_to_scheduler.send_pyobj(obj)
self._send_multi_parts(self.send_to_scheduler, obj)
self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1:
result = await self.model_update_result
@@ -1388,7 +1494,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def freeze_gc(self):
"""Send a freeze_gc message to the scheduler first, then freeze locally."""
self.send_to_scheduler.send_pyobj(FreezeGCReq())
self._send_multi_parts(self.send_to_scheduler, FreezeGCReq())
freeze_gc("Tokenizer Manager")
return None
@@ -1591,7 +1697,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
and recv_obj.load is not None
):
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
self.send_to_scheduler.send_pyobj(load_update_req)
self._send_multi_parts(self.send_to_scheduler, load_update_req)
def add_logprob_to_meta_info(
self,