diff --git a/python/sglang/srt/disaggregation/encoder/receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py index 0ddb26ed4..50b75ccae 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -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 = {} diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 2927824fa..7e2611587 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -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 diff --git a/python/sglang/srt/managers/rust_server.py b/python/sglang/srt/managers/rust_server.py index 3e699d579..79c1a48aa 100644 --- a/python/sglang/srt/managers/rust_server.py +++ b/python/sglang/srt/managers/rust_server.py @@ -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 ) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 51b214b54..b9a102087 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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 diff --git a/python/sglang/srt/multimodal/transport/__init__.py b/python/sglang/srt/multimodal/transport/__init__.py index 659092775..3509b7400 100644 --- a/python/sglang/srt/multimodal/transport/__init__.py +++ b/python/sglang/srt/multimodal/transport/__init__.py @@ -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" diff --git a/test/registered/unit/multimodal/test_tensor_transport_mode.py b/test/registered/unit/multimodal/test_tensor_transport_mode.py new file mode 100644 index 000000000..efa2a32bf --- /dev/null +++ b/test/registered/unit/multimodal/test_tensor_transport_mode.py @@ -0,0 +1,37 @@ +"""Tests for multimodal tensor transport topology detection.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from sglang.srt.multimodal.transport import determine_tensor_transport_mode +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestTensorTransportMode(CustomTestCase): + def test_transport_mode_uses_published_node_topology(self): + cases = ( + (1, None, "cuda_ipc"), + (1, "127.0.0.1:20000", "cuda_ipc"), + (2, None, "default"), + (2, "10.0.0.1:20000", "default"), + ) + + for nnodes, dist_init_addr, expected in cases: + with self.subTest(nnodes=nnodes, dist_init_addr=dist_init_addr): + parallel = SimpleNamespace( + nnodes=nnodes, + dist_init_addr=dist_init_addr, + ) + with patch( + "sglang.srt.multimodal.transport.get_parallel", + return_value=parallel, + ): + self.assertEqual(determine_tensor_transport_mode(), expected) + + +if __name__ == "__main__": + unittest.main()