diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 4f297a32d..c6424e017 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -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) diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 0e9d314e3..f0a455460 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: - 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) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index bca1c31e6..1db8a3df8 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py index e5d42bed8..b395bb184 100644 --- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py +++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py @@ -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) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index f4fc29e29..e40d701b1 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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,