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

This commit is contained in:
Yuhao Yang
2026-01-04 22:05:26 -08:00
committed by GitHub
parent e53160bb31
commit 2138ff48c6
5 changed files with 23 additions and 271 deletions
@@ -16,7 +16,6 @@
import faulthandler import faulthandler
import logging import logging
import multiprocessing as mp import multiprocessing as mp
import pickle
import signal import signal
import threading import threading
import time import time
@@ -26,7 +25,6 @@ from typing import Callable, List, Optional
import psutil import psutil
import setproctitle import setproctitle
import torch
import zmq import zmq
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -210,15 +208,13 @@ class DataParallelController:
) )
def send_to_all_workers(self, obj): def send_to_all_workers(self, obj):
msg = [b"NORM", pickle.dumps(obj)]
for worker in self.workers: for worker in self.workers:
worker.send_multipart(msg, copy=False) worker.send_pyobj(obj)
def send_control_message(self, obj): def send_control_message(self, obj):
msg = [b"NORM", pickle.dumps(obj)]
# Send control messages to first worker of tp group # Send control messages to first worker of tp group
for worker in self.workers[:: self.control_message_step]: 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): def handle_load_update_req(self, obj):
self.dp_budget.update_budget(obj) self.dp_budget.update_budget(obj)
@@ -503,9 +499,8 @@ class DataParallelController:
def maybe_external_dp_rank_routing(self, req: Req): def maybe_external_dp_rank_routing(self, req: Req):
if req.data_parallel_rank is not None: 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}") 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 True
return False return False
@@ -513,8 +508,7 @@ class DataParallelController:
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
return return
msg = [b"NORM", pickle.dumps(req)] self.workers[self.round_robin_counter].send_pyobj(req)
self.workers[self.round_robin_counter].send_multipart(msg, copy=False)
self.round_robin_counter = (self.round_robin_counter + 1) % len(self.workers) self.round_robin_counter = (self.round_robin_counter + 1) % len(self.workers)
def follow_bootstrap_room_scheduler(self, req: Req): def follow_bootstrap_room_scheduler(self, req: Req):
@@ -536,8 +530,7 @@ class DataParallelController:
"prefill or decode instances; send to the router instead." "prefill or decode instances; send to the router instead."
) )
target_rank = req.bootstrap_room % len(self.workers) target_rank = req.bootstrap_room % len(self.workers)
msg = [b"NORM", pickle.dumps(req)] self.workers[target_rank].send_pyobj(req)
self.workers[target_rank].send_multipart(msg, copy=False)
def shortest_queue_scheduler(self, req): def shortest_queue_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
@@ -549,8 +542,7 @@ class DataParallelController:
else: else:
self.follow_bootstrap_room_scheduler(req) self.follow_bootstrap_room_scheduler(req)
else: else:
msg = [b"NORM", pickle.dumps(req)] self.workers[target_worker].send_pyobj(req)
self.workers[target_worker].send_multipart(msg, copy=False)
def minimum_tokens_scheduler(self, req): def minimum_tokens_scheduler(self, req):
if self.maybe_external_dp_rank_routing(req): if self.maybe_external_dp_rank_routing(req):
@@ -565,87 +557,12 @@ class DataParallelController:
else: else:
self.follow_bootstrap_room_scheduler(req) 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): def event_loop(self):
while True: while True:
while True: while True:
self.soft_watchdog.feed() self.soft_watchdog.feed()
try: try:
parts = self.recv_from_tokenizer.recv_multipart( recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
flags=zmq.NOBLOCK, copy=False
)
if not parts:
break
recv_req = self._parse_multipart_message(parts)
except zmq.ZMQError: except zmq.ZMQError:
break break
self._request_dispatcher(recv_req) self._request_dispatcher(recv_req)
@@ -369,8 +369,8 @@ class MultiTokenizerRouter:
async def router_worker_obj(self): async def router_worker_obj(self):
while True: while True:
parts = await self.receive_from_worker.recv_multipart(copy=False) recv_obj = await self.receive_from_worker.recv_pyobj()
await self.send_to_scheduler.send_multipart(parts, copy=False) await self.send_to_scheduler.send_pyobj(recv_obj)
async def handle_loop(self): async def handle_loop(self):
# special reqs will recv from scheduler, need to route to right worker # special reqs will recv from scheduler, need to route to right worker
@@ -525,10 +525,3 @@ class SenderWrapper:
if isinstance(obj, BaseReq): if isinstance(obj, BaseReq):
obj.http_worker_ipc = self.port_args.tokenizer_ipc_name obj.http_worker_ipc = self.port_args.tokenizer_ipc_name
self.send_to_scheduler.send_pyobj(obj) 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)
+3 -44
View File
@@ -16,7 +16,6 @@
import faulthandler import faulthandler
import logging import logging
import os import os
import pickle
import signal import signal
import sys import sys
import time import time
@@ -1189,41 +1188,6 @@ class Scheduler(
return False return False
return num_recv_reqs >= self.max_recv_per_poll 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( def recv_requests(
self, self,
) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]: ) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]:
@@ -1244,15 +1208,10 @@ class Scheduler(
try: try:
if self.recv_limit_reached(len(recv_reqs)): if self.recv_limit_reached(len(recv_reqs)):
break break
parts = self.recv_from_tokenizer.recv_multipart( recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
flags=zmq.NOBLOCK, copy=False except zmq.ZMQError:
)
if not parts:
break
recv_req = self._parse_multipart_message(parts)
recv_reqs.append(recv_req)
except zmq.ZMQError as e:
break break
recv_reqs.append(recv_req)
while True: while True:
try: try:
@@ -3,7 +3,6 @@ from __future__ import annotations
import asyncio import asyncio
import copy import copy
import logging import logging
import pickle
import time import time
import uuid import uuid
from collections import deque from collections import deque
@@ -106,12 +105,7 @@ class _Communicator(Generic[T]):
assert self._result_values is None assert self._result_values is None
if obj: if obj:
self._sender.send_multipart( self._sender.send_pyobj(obj)
[
b"NORM",
pickle.dumps(obj),
]
)
self._result_event = asyncio.Event() self._result_event = asyncio.Event()
self._result_values = [] self._result_values = []
@@ -131,12 +125,7 @@ class _Communicator(Generic[T]):
self._result_event = asyncio.Event() self._result_event = asyncio.Event()
if obj: if obj:
self._sender.send_multipart( self._sender.send_pyobj(obj)
[
b"NORM",
pickle.dumps(obj),
]
)
await self._result_event.wait() await self._result_event.wait()
result_values = copy.deepcopy(self._result_values) result_values = copy.deepcopy(self._result_values)
+9 -115
View File
@@ -15,7 +15,6 @@
import asyncio import asyncio
import copy import copy
import ctypes
import dataclasses import dataclasses
import logging import logging
import os import os
@@ -33,7 +32,6 @@ from http import HTTPStatus
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union from typing import Any, Awaitable, Dict, List, Optional, Tuple, Union
import fastapi import fastapi
import torch
import uvloop import uvloop
import zmq import zmq
import zmq.asyncio import zmq.asyncio
@@ -121,29 +119,6 @@ asyncio.set_event_loop_policy(uvloop.EventLoopPolicy())
logger = logging.getLogger(__name__) 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 @dataclasses.dataclass
class ReqState: class ReqState:
"""Store the state a request.""" """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( def _send_one_request(
self, self,
obj: Union[GenerateReqInput, EmbeddingReqInput], obj: Union[GenerateReqInput, EmbeddingReqInput],
@@ -1145,7 +1038,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
): ):
trace_slice_start(RequestStage.TOKENIZER_DISPATCH, obj.rid) trace_slice_start(RequestStage.TOKENIZER_DISPATCH, obj.rid)
tokenized_obj.trace_context = trace_get_proc_propagate_context(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 = ReqState([], False, asyncio.Event(), obj, created_time=created_time)
state.request_sent_to_scheduler_ts = time.time() state.request_sent_to_scheduler_ts = time.time()
self.rid_to_state[obj.rid] = state self.rid_to_state[obj.rid] = state
@@ -1167,7 +1060,8 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs) batch_req = BatchTokenizedGenerateReqInput(batch=tokenized_objs)
else: else:
batch_req = BatchTokenizedEmbeddingReqInput(batch=tokenized_objs) 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 # Create states for each individual request in the batch
for i, tokenized_obj in enumerate(tokenized_objs): for i, tokenized_obj in enumerate(tokenized_objs):
tmp_obj = obj[i] tmp_obj = obj[i]
@@ -1393,7 +1287,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
if not abort_all and rid not in self.rid_to_state: if not abort_all and rid not in self.rid_to_state:
return return
req = AbortReq(rid=rid, abort_all=abort_all) 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: if self.enable_metrics:
# TODO: also use custom_labels from the request # TODO: also use custom_labels from the request
self.metrics_collector.observe_one_aborted_request( self.metrics_collector.observe_one_aborted_request(
@@ -1404,7 +1298,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async with self.is_pause_cond: async with self.is_pause_cond:
self.is_pause = True self.is_pause = True
if obj.mode != "abort": if obj.mode != "abort":
await self._send_multi_parts_async(self.send_to_scheduler, obj) await self.send_to_scheduler.send_pyobj(obj)
else: else:
# we are using the model_update_lock to check if there is still on-going requests. # we are using the model_update_lock to check if there is still on-going requests.
while True: while True:
@@ -1418,7 +1312,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def continue_generation(self, obj: ContinueGenerationReqInput): async def continue_generation(self, obj: ContinueGenerationReqInput):
async with self.is_pause_cond: async with self.is_pause_cond:
self.is_pause = False 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() self.is_pause_cond.notify_all()
async def update_weights_from_disk( async def update_weights_from_disk(
@@ -1463,7 +1357,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def _wait_for_model_update_from_disk( async def _wait_for_model_update_from_disk(
self, obj: UpdateWeightFromDiskReqInput self, obj: UpdateWeightFromDiskReqInput
) -> Tuple[bool, str]: ) -> 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() self.model_update_result = asyncio.Future()
if self.server_args.dp_size == 1: if self.server_args.dp_size == 1:
result = await self.model_update_result result = await self.model_update_result
@@ -1498,7 +1392,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
async def freeze_gc(self): async def freeze_gc(self):
"""Send a freeze_gc message to the scheduler first, then freeze locally.""" """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") freeze_gc("Tokenizer Manager")
return None return None
@@ -1701,7 +1595,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
and recv_obj.load is not None and recv_obj.load is not None
): ):
load_update_req = WatchLoadUpdateReq(loads=[recv_obj.load]) 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( def add_logprob_to_meta_info(
self, self,