diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index c6424e017..4f297a32d 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -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) diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index f0a455460..0e9d314e3 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -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) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index dedff1f69..b7ff7134f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index b395bb184..e5d42bed8 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -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) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 5cfefda66..fc65ce962 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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,