Revert "[FEAT] optimize tensor zmq transfer for multimodal inputs" (#16386)

This commit is contained in:
Yuhao Yang
2026-01-05 14:05:26 +08:00
committed by GitHub
parent e53160bb31
commit 2138ff48c6
5 changed files with 23 additions and 271 deletions

View File

@@ -16,7 +16,6 @@
import faulthandler
import logging
import multiprocessing as mp
import pickle
import signal
import threading
import time
@@ -26,7 +25,6 @@ from typing import Callable, List, Optional
import psutil
import setproctitle
import torch
import zmq
from sglang.srt.environ import envs
@@ -210,15 +208,13 @@ class DataParallelController:
)
def send_to_all_workers(self, obj):
msg = [b"NORM", pickle.dumps(obj)]
for worker in self.workers:
worker.send_multipart(msg, copy=False)
worker.send_pyobj(obj)
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_multipart(msg, copy=False)
worker.send_pyobj(obj)
def handle_load_update_req(self, obj):
self.dp_budget.update_budget(obj)
@@ -503,9 +499,8 @@ 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_multipart(msg, copy=False)
self.workers[req.data_parallel_rank].send_pyobj(req)
return True
return False
@@ -513,8 +508,7 @@ class DataParallelController:
if self.maybe_external_dp_rank_routing(req):
return
msg = [b"NORM", pickle.dumps(req)]
self.workers[self.round_robin_counter].send_multipart(msg, copy=False)
self.workers[self.round_robin_counter].send_pyobj(req)
self.round_robin_counter = (self.round_robin_counter + 1) % len(self.workers)
def follow_bootstrap_room_scheduler(self, req: Req):
@@ -536,8 +530,7 @@ class DataParallelController:
"prefill or decode instances; send to the router instead."
)
target_rank = req.bootstrap_room % len(self.workers)
msg = [b"NORM", pickle.dumps(req)]
self.workers[target_rank].send_multipart(msg, copy=False)
self.workers[target_rank].send_pyobj(req)
def shortest_queue_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req):
@@ -549,8 +542,7 @@ class DataParallelController:
else:
self.follow_bootstrap_room_scheduler(req)
else:
msg = [b"NORM", pickle.dumps(req)]
self.workers[target_worker].send_multipart(msg, copy=False)
self.workers[target_worker].send_pyobj(req)
def minimum_tokens_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req):
@@ -565,87 +557,12 @@ 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:
parts = self.recv_from_tokenizer.recv_multipart(
flags=zmq.NOBLOCK, copy=False
)
if not parts:
break
recv_req = self._parse_multipart_message(parts)
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
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:
parts = await self.receive_from_worker.recv_multipart(copy=False)
await self.send_to_scheduler.send_multipart(parts, copy=False)
recv_obj = await self.receive_from_worker.recv_pyobj()
await self.send_to_scheduler.send_pyobj(recv_obj)
async def handle_loop(self):
# special reqs will recv from scheduler, need to route to right worker
@@ -525,10 +525,3 @@ 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,7 +16,6 @@
import faulthandler
import logging
import os
import pickle
import signal
import sys
import time
@@ -1189,41 +1188,6 @@ 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]]:
@@ -1244,15 +1208,10 @@ class Scheduler(
try:
if self.recv_limit_reached(len(recv_reqs)):
break
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:
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
except zmq.ZMQError:
break
recv_reqs.append(recv_req)
while True:
try:

View File

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

View File

@@ -15,7 +15,6 @@
import asyncio
import copy
import ctypes
import dataclasses
import logging
import os
@@ -33,7 +32,6 @@ 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
@@ -121,29 +119,6 @@ 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."""
@@ -1055,88 +1030,6 @@ 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],
@@ -1145,7 +1038,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_multi_parts(self.send_to_scheduler, tokenized_obj)
self.send_to_scheduler.send_pyobj(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
@@ -1167,7 +1060,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs)
self._send_multi_parts(self.send_to_scheduler, batch_req)
self.send_to_scheduler.send_pyobj(batch_req)
# Create states for each individual request in the batch
for i, tokenized_obj in enumerate(tokenized_objs):
tmp_obj = obj[i]
@@ -1393,7 +1287,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_multi_parts(self.send_to_scheduler, req)
self.send_to_scheduler.send_pyobj(req)
if self.enable_metrics:
# TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request(
@@ -1404,7 +1298,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async with self.is_pause_cond:
self.is_pause = True
if obj.mode != "abort":
await self._send_multi_parts_async(self.send_to_scheduler, obj)
await self.send_to_scheduler.send_pyobj(obj)
else:
# we are using the model_update_lock to check if there is still on-going requests.
while True:
@@ -1418,7 +1312,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_multi_parts_async(self.send_to_scheduler, obj)
await self.send_to_scheduler.send_pyobj(obj)
self.is_pause_cond.notify_all()
async def update_weights_from_disk(
@@ -1463,7 +1357,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]:
self._send_multi_parts(self.send_to_scheduler, obj)
self.send_to_scheduler.send_pyobj(obj)
self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1:
result = await self.model_update_result
@@ -1498,7 +1392,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def freeze_gc(self):
"""Send a freeze_gc message to the scheduler first, then freeze locally."""
self._send_multi_parts(self.send_to_scheduler, FreezeGCReq())
self.send_to_scheduler.send_pyobj(FreezeGCReq())
freeze_gc("Tokenizer Manager")
return None
@@ -1701,7 +1595,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
and recv_obj.load is not None
):
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load])
self._send_multi_parts(self.send_to_scheduler, load_update_req)
self.send_to_scheduler.send_pyobj(load_update_req)
def add_logprob_to_meta_info(
self,