Revert "[FEAT] optimize tensor zmq transfer for multimodal inputs" (#16386)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user