fix: detect cross-node multimodal transport by nnodes (#35646)

This commit is contained in:
Yuanle Liu
2026-08-26 10:31:48 +08:00
committed by GitHub
parent 2d8484740d
commit 4ae30dc736
6 changed files with 69 additions and 38 deletions
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReq
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
from sglang.srt.runtime_context import get_disagg, get_exec, get_serving
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
@@ -1696,16 +1697,6 @@ def _view_pool_buffer_by_modality(raw_buffer, embedding_data, dtype):
}
def _determine_tensor_transport_mode(server_args):
is_cross_node = server_args.dist_init_addr
if is_cross_node:
# Fallback to default CPU transport for multi-node
return "default"
else:
return "cuda_ipc"
class MMReceiverBase(ABC):
def __init__(
self,
@@ -1835,7 +1826,7 @@ class MMReceiverBase(ABC):
model_config=None,
):
"""Load processor and initialize mm_processor, shared by all backends."""
transport_mode = _determine_tensor_transport_mode(server_args)
transport_mode = determine_tensor_transport_mode()
import_processors("sglang.srt.multimodal.processors")
extra_kwargs = {}
+6 -13
View File
@@ -10,7 +10,7 @@ import sys
from abc import abstractmethod
from collections import defaultdict
from multiprocessing import shared_memory
from typing import Any, Dict, List, Literal, Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import torch
@@ -38,6 +38,10 @@ from sglang.srt.managers.schedule_batch import (
MultimodalInputs,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.multimodal.transport import (
TensorTransportMode,
determine_tensor_transport_mode,
)
from sglang.srt.runtime_context import (
get_disagg,
get_server_args,
@@ -52,11 +56,6 @@ from sglang.utils import logger
# propagation that can cause some log messages (like 'server is fired up') to not appear
# in the console when multimodal support is enabled.
# TODO(mick): nccl
# cuda_ipc: for intranode tensor sharing
TensorTransportMode = Literal["cuda_ipc", "auto", "default"]
_GPU_FEATURE_BUFFER: Optional[torch.Tensor] = None
_BUFFER_OFFSET = 0
@@ -1336,13 +1335,7 @@ class ShmPointerMMData:
def _get_is_default_transport():
global _is_default_tensor_transport
if _is_default_tensor_transport is None:
from sglang.srt.managers.tokenizer_manager import (
determine_tensor_transport_mode,
)
_is_default_tensor_transport = (
determine_tensor_transport_mode(get_server_args()) == "default"
)
_is_default_tensor_transport = determine_tensor_transport_mode() == "default"
return _is_default_tensor_transport
+2 -2
View File
@@ -267,13 +267,13 @@ class NativeMmHost:
map in parallel — the transport the Python TokenizerManager already uses.
Single-rank serving stays inline, where shm would only add a copy.
"""
from sglang.srt.managers.tokenizer_manager import (
from sglang.srt.multimodal.transport import (
determine_tensor_transport_mode,
)
return (
self.server_args.tp_size > 1
and determine_tensor_transport_mode(self.server_args) != "default"
and determine_tensor_transport_mode() != "default"
and not self.server_args.skip_tokenizer_init
)
@@ -89,7 +89,7 @@ from sglang.srt.managers.io_struct import (
unwrap_from_pickle,
)
from sglang.srt.managers.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features
from sglang.srt.managers.mm_utils import wrap_shm_features
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import (
MultimodalDataItem,
@@ -105,6 +105,7 @@ from sglang.srt.managers.utils import (
from sglang.srt.model_executor.forward_batch_info import (
get_server_return_hidden_states_mode,
)
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.metrics_collector import (
STAT_LOGGER_ROLE_TOKENIZER,
@@ -486,7 +487,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
import_processors(mm_process_pkg, overwrite=True)
_processor = get_processor_wrapper(server_args)
transport_mode = determine_tensor_transport_mode(self.server_args)
transport_mode = determine_tensor_transport_mode()
# We want to parallelize the image pre-processing so we create an executor for it
# We create mm_processor for any skip_tokenizer_init to make sure we still encode
@@ -3598,16 +3599,6 @@ def get_processor_wrapper(server_args):
)
def determine_tensor_transport_mode(server_args: ServerArgs) -> TensorTransportMode:
is_cross_node = get_parallel().dist_init_addr
if is_cross_node:
# Fallback to default CPU transport for multi-node
return "default"
else:
return "cuda_ipc"
class SignalHandler:
def __init__(self, tokenizer_manager: TokenizerManager):
self.tokenizer_manager = tokenizer_manager
@@ -1 +1,20 @@
"""GPU transports for multimodal feature tensors."""
from typing import Literal
from sglang.srt.runtime_context import get_parallel
TensorTransportMode = Literal["cuda_ipc", "auto", "default"]
def determine_tensor_transport_mode() -> TensorTransportMode:
"""Select tensor transport from node topology, not rendezvous configuration.
A rendezvous address may be present for single-node TP. Conversely, Ray can
inject the address only into scheduler actors for a multi-node deployment,
and external launchers may use an environment-based rendezvous instead.
"""
if get_parallel().nnodes > 1:
# CUDA IPC and POSIX shared memory are local to one node.
return "default"
return "cuda_ipc"