Add CUDA VMM multimodal feature transport (#33899)

Co-authored-by: Yinghai Lu <yinghai@meta.com>
This commit is contained in:
Oguz Ulgen
2026-08-07 13:39:54 -07:00
committed by GitHub
co-authored by Yinghai Lu
parent 3c51e29deb
commit 7f6b4cb94b
10 changed files with 2490 additions and 78 deletions
@@ -237,6 +237,31 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
self.assertFalse(processor.use_ipc_pool_handle_cache)
memory_pool.assert_not_called()
def test_cuda_vmm_keeps_features_on_device_without_ipc_pool(self):
from sglang.srt.multimodal.processors import base_processor
hf_processor = self._processor()
feature = torch.empty(1, device="meta")
hf_processor.return_value = {"pixel_values": feature}
with patch.object(
base_processor.BaseMultimodalProcessor, "__abstractmethods__", set()
), patch.object(base_processor, "MmItemMemoryPool") as memory_pool:
processor = base_processor.BaseMultimodalProcessor(
hf_config=MagicMock(),
server_args=self._server_args("cuda_vmm"),
_processor=hf_processor,
transport_mode=None,
)
result = processor.process_mm_data("test")
self.assertEqual(processor.mm_feature_transport, "cuda_vmm")
self.assertFalse(processor.use_cuda_ipc)
self.assertTrue(processor.keep_mm_features_on_device)
self.assertEqual(processor.cpu_executor._mp_context.get_start_method(), "spawn")
self.assertIs(result["pixel_values"], feature)
memory_pool.assert_not_called()
class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
@staticmethod
@@ -251,6 +276,7 @@ class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
processor = BaseMultimodalProcessor()
processor.precompute_hash_before_cpu_transfer = enabled
processor.use_cuda_ipc = False
processor.mm_feature_transport = "cpu"
return processor
def test_enabled_path_sets_hash_and_pad_value(self):
@@ -0,0 +1,490 @@
"""CUDA VMM multimodal feature transport regression tests."""
from __future__ import annotations
import gc
import multiprocessing as mp
import os
import pickle
import queue
import threading
import unittest
from unittest.mock import MagicMock, patch
import torch
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmMemoryPool,
CudaVmmPackedTensorTransportProxy,
_imported_pool_cache_clear,
_PosixFdBroker,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=60, stage="base-c", runner_config="4-gpu-gb300")
class _FabricUnavailableCudaVmmMemoryPool(CudaVmmMemoryPool):
def _allocate(self, memory_size: int) -> None:
if self.use_fabric:
raise RuntimeError("forced FABRIC allocation failure")
super()._allocate(memory_size)
def _produce_vmm_tensor(proxy_queue, consumer_done, result_queue, mode):
pool = source = proxy = None
try:
torch.cuda.set_device(0)
pool_cls = (
_FabricUnavailableCudaVmmMemoryPool
if mode == "posix_fallback"
else CudaVmmMemoryPool
)
pool = pool_cls(
memory_size=4 << 20,
recycle_interval=60,
base_gpu_id=0,
consumer_count=2,
allow_posix_fallback=True,
)
source = torch.arange(35, dtype=torch.float32, device="cuda").reshape(5, 7)
expected = source.cpu().tolist()
proxy = pool.wrap_tensor(source)
proxy_queue.put((proxy, expected))
if not consumer_done.wait(timeout=60):
raise TimeoutError("consumers did not release the CUDA VMM tensor")
with pool._lock:
pool._recycle_chunks()
pool._merge_chunks()
if pool.occupied_chunks:
raise RuntimeError(
"consumer acknowledgements did not recycle the slice"
)
except Exception as exc: # noqa: BLE001 # pragma: no cover
result_queue.put(("error", repr(exc)))
return
finally:
del proxy, source
if pool is not None:
pool.shutdown()
del pool
gc.collect()
result_queue.put(("ok", None))
class TestCudaVmmTransport(CustomTestCase):
@classmethod
def setUpClass(cls):
if (
not torch.cuda.is_available()
or torch.version.cuda is None
or torch.cuda.device_count() < 3
):
raise unittest.SkipTest("At least three NVIDIA CUDA GPUs are required")
def _run_round_trip(self, mode: str):
consumer_devices = (1, 2)
torch.cuda.set_device(consumer_devices[0])
ctx = mp.get_context("spawn")
proxy_queue = ctx.Queue()
producer_results = ctx.Queue()
consumer_done = ctx.Event()
producer = ctx.Process(
target=_produce_vmm_tensor,
args=(proxy_queue, consumer_done, producer_results, mode),
)
producer.start()
proxy = second_proxy = None
reconstructed = []
producer_result = None
try:
try:
proxy, expected = proxy_queue.get(timeout=60)
except queue.Empty:
producer_result = producer_results.get(timeout=5)
_status, payload = producer_result
self.fail(
f"CUDA VMM producer failed before sending its proxy: {payload}"
)
if mode == "posix_fallback":
self.assertIsNone(proxy.fabric_handle)
self.assertIsNotNone(proxy.posix_socket_path)
else:
self.assertIsNotNone(proxy.fabric_handle)
self.assertIsNone(proxy.posix_socket_path)
second_proxy = pickle.loads(pickle.dumps(proxy))
for tp_rank, (consumer_proxy, device) in enumerate(
zip((proxy, second_proxy), consumer_devices)
):
torch.cuda.set_device(device)
with get_parallel().override(
attn_tp_size=2,
attn_tp_rank=tp_rank,
attn_cp_size=1,
attn_cp_rank=0,
):
tensor = consumer_proxy.reconstruct_on_target_device(
device, consumer_count=1
)
torch.cuda.synchronize(device)
self.assertEqual(tensor.cpu().tolist(), expected)
reconstructed.append(tensor)
finally:
del reconstructed, second_proxy, proxy
_imported_pool_cache_clear()
gc.collect()
consumer_done.set()
producer.join(timeout=60)
try:
if producer_result is None:
producer_result = producer_results.get(timeout=5)
status, payload = producer_result
self.assertEqual(status, "ok", payload)
finally:
if producer.is_alive():
producer.terminate()
producer.join(timeout=10)
torch.cuda.set_device(0)
self.assertEqual(producer.exitcode, 0)
def test_posix_fd_fallback_tensor_round_trip_and_recycling(self):
self._run_round_trip(mode="posix_fallback")
def test_auto_prefers_fabric_tensor_round_trip_and_recycling(self):
self._run_round_trip(mode="auto")
def test_reused_chunk_clears_acknowledgements(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 2, allow_posix_fallback=True)
try:
old = pool.wrap_tensor(torch.ones(1024, dtype=torch.uint8, device="cuda:0"))
pool.memory_pool[old.control_offset : old.control_offset + 8].view(
torch.int32
).fill_(1)
torch.cuda.synchronize(0)
with pool._lock:
pool._recycle_chunks()
pool._merge_chunks()
pool.wrap_tensor(torch.ones(100, dtype=torch.uint8, device="cuda:0"))
live = pool.wrap_tensor(torch.ones(256, dtype=torch.uint8, device="cuda:0"))
control = pool.memory_pool[
live.control_offset : live.control_offset + 8
].view(torch.int32)
self.assertTrue(torch.equal(control, torch.zeros_like(control)))
with pool._lock:
pool._recycle_chunks()
self.assertIn(
live.control_offset,
[chunk.start for chunk in pool.occupied_chunks],
)
finally:
pool.shutdown()
def test_packed_tensors_round_trip_through_one_shared_buffer(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
sources = [
torch.arange(24, dtype=torch.float32, device="cuda:0")
.reshape(4, 6)
.transpose(0, 1),
torch.arange(7, dtype=torch.bfloat16),
torch.arange(5, dtype=torch.int64, device="cuda:0"),
]
expected = [source.contiguous().cpu() for source in sources]
proxies = reconstructed = None
try:
stream = MagicMock(wraps=torch.cuda.current_stream(0))
with patch("torch.cuda.current_stream", return_value=stream):
proxies = pool.wrap_tensors(sources)
self.assertIsNotNone(proxies)
self.assertEqual(stream.synchronize.call_count, 1)
self.assertEqual(len(pool.occupied_chunks), 1)
self.assertTrue(
all(
isinstance(proxy, CudaVmmPackedTensorTransportProxy)
for proxy in proxies
)
)
self.assertEqual(len({proxy.control_offset for proxy in proxies}), 1)
proxies = pickle.loads(pickle.dumps(proxies))
self.assertIs(proxies[0]._packed_owner, proxies[-1]._packed_owner)
with get_parallel().override(
attn_tp_size=1,
attn_tp_rank=0,
attn_cp_size=1,
attn_cp_rank=0,
):
reconstructed = [
proxy.reconstruct_on_target_device(0, consumer_count=1)
for proxy in proxies
]
torch.cuda.synchronize(0)
for actual, wanted in zip(reconstructed, expected):
self.assertTrue(torch.equal(actual.cpu(), wanted))
packed_storage = proxies[
0
]._packed_owner.reconstruct_tensor.untyped_storage()
self.assertTrue(
all(
tensor.untyped_storage().data_ptr() == packed_storage.data_ptr()
for tensor in reconstructed
)
)
with pool._lock:
pool._recycle_chunks()
self.assertFalse(pool.occupied_chunks)
finally:
del reconstructed, proxies, expected, sources
_imported_pool_cache_clear()
pool.shutdown()
def test_packed_cancel_is_shared_and_idempotent(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
proxies = pool.wrap_tensors(
[
torch.ones(8, dtype=torch.float32, device="cuda:0"),
torch.ones(8, dtype=torch.float32, device="cuda:0"),
]
)
self.assertIsNotNone(proxies)
pool.cancel_proxy(proxies[0])
pool.cancel_proxy(proxies[1])
self.assertFalse(pool.occupied_chunks)
self.assertEqual(
sum(chunk.size for chunk in pool.available_chunks),
pool.allocation_size,
)
finally:
pool.shutdown()
def test_packed_reservation_failure_returns_fallback_signal(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
source = torch.empty(
pool.allocation_size, dtype=torch.uint8, device="cuda:0"
)
self.assertIsNone(pool.wrap_tensors([source]))
self.assertFalse(pool.occupied_chunks)
self.assertEqual(
sum(chunk.size for chunk in pool.available_chunks),
pool.allocation_size,
)
finally:
pool.shutdown()
def test_oversized_tensor_falls_back_to_cpu(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
source = torch.empty(
pool.allocation_size, dtype=torch.uint8, device="cuda:0"
)
fallback = pool.wrap_tensor(source)
self.assertTrue(fallback.is_cpu)
self.assertFalse(pool.occupied_chunks)
self.assertEqual(
sum(chunk.size for chunk in pool.available_chunks),
pool.allocation_size,
)
finally:
pool.shutdown()
def test_failed_copy_rolls_back_reservation(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
with self.assertRaisesRegex(NotImplementedError, "meta tensor"):
pool.wrap_tensor(torch.ones(16, device="meta"))
with self.assertRaisesRegex(NotImplementedError, "meta tensor"):
pool.wrap_tensors(
[
torch.ones(16, device="cuda:0"),
torch.ones(16, device="meta"),
]
)
self.assertFalse(pool.occupied_chunks)
self.assertEqual(
sum(chunk.size for chunk in pool.available_chunks),
pool.allocation_size,
)
finally:
pool.shutdown()
def test_undispatched_proxy_can_be_cancelled_immediately(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
proxy = pool.wrap_tensor(torch.ones(16, device="cuda:0"))
pool.cancel_proxy(proxy)
self.assertFalse(pool.occupied_chunks)
self.assertEqual(
sum(chunk.size for chunk in pool.available_chunks),
pool.allocation_size,
)
finally:
pool.shutdown()
def test_failed_cleanup_sync_quarantines_pool(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
stream = MagicMock()
stream.synchronize.side_effect = RuntimeError("forced sync failure")
try:
with (
patch("torch.cuda.current_stream", return_value=stream),
self.assertRaisesRegex(RuntimeError, "forced sync failure"),
):
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
torch.cuda.synchronize(0)
self.assertIsNotNone(pool._pool_error)
with self.assertRaisesRegex(RuntimeError, "pool failed"):
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
finally:
pool.shutdown()
def test_shutdown_waits_for_active_publisher(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
publisher_entered = threading.Event()
allow_publisher_to_finish = threading.Event()
shutdown_entered = threading.Event()
shutdown_finished = threading.Event()
errors = []
real_stream = torch.cuda.current_stream(0)
stream = MagicMock(wraps=real_stream)
def synchronize():
publisher_entered.set()
if not allow_publisher_to_finish.wait(timeout=10):
raise TimeoutError("publisher was not released")
real_stream.synchronize()
stream.synchronize.side_effect = synchronize
def publish():
try:
with patch("torch.cuda.current_stream", return_value=stream):
pool.wrap_tensor(torch.ones(16, device="cuda:0"))
except Exception as error: # pragma: no cover
errors.append(error)
def shutdown():
shutdown_entered.set()
try:
pool.shutdown()
except Exception as error: # pragma: no cover
errors.append(error)
finally:
shutdown_finished.set()
publisher = threading.Thread(target=publish)
shutdown_thread = threading.Thread(target=shutdown)
try:
publisher.start()
self.assertTrue(publisher_entered.wait(timeout=10))
shutdown_thread.start()
self.assertTrue(shutdown_entered.wait(timeout=10))
self.assertFalse(shutdown_finished.wait(timeout=0.1))
finally:
allow_publisher_to_finish.set()
publisher.join(timeout=10)
shutdown_thread.join(timeout=10)
if not pool._closed:
pool.shutdown()
self.assertFalse(publisher.is_alive())
self.assertFalse(shutdown_thread.is_alive())
self.assertFalse(errors)
def test_consumer_copy_failure_releases_slice_without_allowing_retry(self):
pool = CudaVmmMemoryPool(4 << 20, 60, 0, 1, allow_posix_fallback=True)
try:
proxy = pool.wrap_tensor(torch.ones(1, device="cuda:0"))
proxy.shape = (2,)
with (
get_parallel().override(
attn_tp_size=1,
attn_tp_rank=0,
attn_cp_size=1,
attn_cp_rank=0,
),
self.assertRaises(RuntimeError),
):
proxy.reconstruct_on_target_device(0, consumer_count=1)
torch.cuda.synchronize(0)
with pool._lock:
pool._recycle_chunks()
self.assertFalse(pool.occupied_chunks)
with (
get_parallel().override(
attn_tp_size=1,
attn_tp_rank=0,
attn_cp_size=1,
attn_cp_rank=0,
),
self.assertRaisesRegex(RuntimeError, "already released"),
):
proxy.reconstruct_on_target_device(0, consumer_count=1)
finally:
_imported_pool_cache_clear()
pool.shutdown()
def test_posix_export_fd_closes_when_allocation_setup_fails(self):
with (
patch(
"sglang.srt.utils.cuda_vmm_transport_utils._tensor_from_pointer",
side_effect=RuntimeError("forced storage failure"),
),
patch(
"sglang.srt.utils.cuda_vmm_transport_utils.os.close",
wraps=os.close,
) as close_fd,
self.assertRaisesRegex(RuntimeError, "forced storage failure"),
):
_FabricUnavailableCudaVmmMemoryPool(
4 << 20, 60, 0, 1, allow_posix_fallback=True
)
close_fd.assert_called_once()
def test_stream_setup_failure_releases_pool_and_posix_broker(self):
release_allocation = CudaVmmMemoryPool._release_allocation
close_broker = _PosixFdBroker.close
with (
patch(
"sglang.srt.utils.cuda_vmm_transport_utils.torch.cuda.Stream",
side_effect=RuntimeError("forced stream failure"),
),
patch.object(
CudaVmmMemoryPool,
"_release_allocation",
autospec=True,
side_effect=release_allocation,
) as release_pool,
patch.object(
_PosixFdBroker,
"close",
autospec=True,
side_effect=close_broker,
) as close_fd_broker,
self.assertRaisesRegex(RuntimeError, "forced stream failure"),
):
_FabricUnavailableCudaVmmMemoryPool(
4 << 20, 60, 0, 1, allow_posix_fallback=True
)
close_fd_broker.assert_called_once()
release_pool.assert_called_once()
if __name__ == "__main__":
unittest.main(verbosity=2)
@@ -0,0 +1,666 @@
import unittest
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import MagicMock, call, patch
import torch
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_partial_pool_release_can_be_retried(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
pool = object.__new__(vmm.CudaVmmMemoryPool)
pool.memory_pool = object()
pool.use_fabric = True
pool.shareable_handle = b"handle"
pool._pool_pointer = 123
pool._allocation_handle = 456
pool._allocation_mapped = True
pool.allocation_size = 4096
pool.device_index = 0
driver = MagicMock()
driver.cuMemUnmap.return_value = "unmap"
driver.cuMemAddressFree.return_value = "address_free"
driver.cuMemRelease.return_value = "release"
failed_once = False
def check_driver(result, _operation):
nonlocal failed_once
if result == "address_free" and not failed_once:
failed_once = True
raise RuntimeError("forced address-free failure")
return result
with (
patch.object(vmm, "_get_cuda_driver", return_value=driver),
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
patch.object(vmm, "check_drv", side_effect=check_driver),
self.assertRaisesRegex(RuntimeError, "forced address-free failure"),
):
pool._release_allocation()
self.assertFalse(pool._allocation_mapped)
self.assertEqual(pool._pool_pointer, 123)
self.assertEqual(pool._allocation_handle, 456)
with (
patch.object(vmm, "_get_cuda_driver", return_value=driver),
patch.object(vmm.torch.cuda, "device", return_value=nullcontext()),
patch.object(vmm, "check_drv", side_effect=lambda result, _: result),
):
pool._release_allocation()
self.assertIsNone(pool._pool_pointer)
self.assertIsNone(pool._allocation_handle)
self.assertEqual(driver.cuMemUnmap.call_count, 1)
self.assertEqual(driver.cuMemAddressFree.call_count, 2)
self.assertEqual(driver.cuMemRelease.call_count, 1)
def test_model_class_controls_cuda_vmm_opt_in(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
class SupportedModel:
supports_cuda_vmm_feature_transport = True
class UnsupportedModel:
pass
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
manager.model_config = object()
with patch(
"sglang.srt.model_loader.utils.get_model_architecture",
return_value=(SupportedModel, "supported"),
):
manager._validate_cuda_vmm_feature_transport_support()
with (
patch(
"sglang.srt.model_loader.utils.get_model_architecture",
return_value=(UnsupportedModel, "unsupported"),
),
self.assertRaisesRegex(ValueError, "UnsupportedModel"),
):
manager._validate_cuda_vmm_feature_transport_support()
def test_cpu_transport_skips_model_opt_in_lookup(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cpu")
manager.model_config = object()
with patch(
"sglang.srt.model_loader.utils.get_model_architecture"
) as get_model_architecture:
manager._validate_cuda_vmm_feature_transport_support()
get_model_architecture.assert_not_called()
def test_vmm_transport_initializes_pool(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
server_args = SimpleNamespace(
mm_feature_transport="cuda_vmm",
tokenizer_worker_num=2,
base_gpu_id=3,
enable_dp_attention=False,
tp_size=4,
nnodes=1,
)
pool = object()
with (
patch.object(vmm, "get_mm_feature_pool_size_per_worker", return_value=123),
patch.object(vmm, "CudaVmmMemoryPool", return_value=pool) as pool_class,
):
transport = vmm.CudaVmmFeatureTransport(server_args, SimpleNamespace())
self.assertIs(transport.pool, pool)
pool_class.assert_called_once_with(
memory_size=123,
recycle_interval=vmm.MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
base_gpu_id=3,
consumer_count=4,
allow_posix_fallback=True,
)
def test_disabled_transport_is_a_noop(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = CudaVmmFeatureTransport(
SimpleNamespace(mm_feature_transport="cpu"), None
)
self.assertEqual(transport.prepare_for_dispatch([None]), [])
transport.cancel_for_dispatch([])
transport.shutdown()
self.assertIsNone(transport.pool)
def test_vmm_transport_requires_processor(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
with self.assertRaisesRegex(RuntimeError, "multimodal processor"):
CudaVmmFeatureTransport(
SimpleNamespace(mm_feature_transport="cuda_vmm"), None
)
def test_image_features_are_packed_per_request(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
features = [torch.arange(4), torch.arange(4, 8)]
proxies = [object(), object()]
transport.pool.wrap_tensors.return_value = proxies
items = [
MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
for feature in features
]
transport.wrap_items(items)
transport.pool.wrap_tensors.assert_called_once_with(features)
transport.pool.wrap_tensor.assert_not_called()
self.assertEqual([item.feature for item in items], proxies)
def test_deferred_features_are_not_packed(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_ipc_transport_utils import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
)
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
features = [torch.arange(4), torch.arange(4, 8)]
proxies = [object(), object()]
transport.pool.wrap_tensor.side_effect = proxies
items = [
MultimodalDataItem(
modality=Modality.IMAGE,
feature=feature,
model_specific_data={DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY: True},
)
for feature in features
]
transport.wrap_items(items)
transport.pool.wrap_tensors.assert_not_called()
self.assertEqual(
transport.pool.wrap_tensor.call_args_list,
[call(feature) for feature in features],
)
self.assertEqual([item.feature for item in items], proxies)
def test_tensor_containers_fail_closed(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
item = MultimodalDataItem(
modality=Modality.IMAGE,
feature=[torch.arange(4), torch.arange(4, 8)],
)
with self.assertRaisesRegex(TypeError, "single tensor"):
transport.wrap_items([item])
transport.pool.wrap_tensor.assert_not_called()
transport.pool.wrap_tensors.assert_not_called()
def test_partial_failure_restores_tensors_and_cancels_packed_chunk_once(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
CudaVmmMemoryPool,
CudaVmmPackedTensorTransportProxy,
_CudaVmmPackedTransportOwner,
)
owner = object.__new__(_CudaVmmPackedTransportOwner)
owner.control_offset = 64
owner._producer_cancelled = False
proxies = [object.__new__(CudaVmmPackedTensorTransportProxy) for _ in range(2)]
for proxy in proxies:
proxy._packed_owner = owner
pool = object.__new__(CudaVmmMemoryPool)
pool.wrap_tensors = MagicMock(return_value=proxies)
pool.wrap_tensor = MagicMock(side_effect=RuntimeError("copy failed"))
pool._cancel_control_offset = MagicMock()
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = pool
features = [torch.arange(4), torch.arange(4, 8)]
embedding = torch.arange(2)
items = [
MultimodalDataItem(
modality=Modality.IMAGE,
feature=features[0],
precomputed_embeddings=embedding,
),
MultimodalDataItem(modality=Modality.IMAGE, feature=features[1]),
]
with self.assertRaisesRegex(RuntimeError, "copy failed"):
transport.wrap_items(items)
for item, feature in zip(items, features, strict=True):
self.assertIs(item.feature, feature)
self.assertIs(items[0].precomputed_embeddings, embedding)
pool._cancel_control_offset.assert_called_once_with(owner.control_offset)
def test_text_request_uses_base_send_path(self):
from sglang.srt.managers import tokenizer_manager
from sglang.srt.managers.tokenizer_manager import TokenizerManager
manager = object.__new__(TokenizerManager)
transport = MagicMock()
transport.prepare_for_dispatch.return_value = []
manager.cuda_vmm_feature_transport = transport
manager._dispatch_to_scheduler = MagicMock()
state = SimpleNamespace(dispatched=False)
manager.rid_to_state = {"test-request": state}
tokenized_obj = SimpleNamespace(
rid="test-request",
mm_inputs=None,
time_stats=MagicMock(),
wrap_pickle_fields=MagicMock(),
)
with patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj):
manager._send_one_request(tokenized_obj)
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
transport.prepare_for_dispatch.assert_called_once_with((None,))
transport.cancel_for_dispatch.assert_not_called()
self.assertTrue(state.dispatched)
def test_failed_dispatch_cancels_published_items(self):
from sglang.srt.managers import tokenizer_manager
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
manager = object.__new__(tokenizer_manager.TokenizerManager)
transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock(
side_effect=RuntimeError("send failed")
)
state = SimpleNamespace(dispatched=False)
manager.rid_to_state = {"test-request": state}
items = [MultimodalDataItem(modality=Modality.IMAGE, feature=torch.arange(2))]
tokenized_obj = SimpleNamespace(
rid="test-request",
mm_inputs=MultimodalProcessorOutput(input_ids=[1], mm_items=items),
time_stats=MagicMock(),
wrap_pickle_fields=MagicMock(),
)
transport.prepare_for_dispatch.return_value = items
manager.cuda_vmm_feature_transport = transport
with (
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
self.assertRaisesRegex(RuntimeError, "send failed"),
):
manager._send_one_request(tokenized_obj)
transport.prepare_for_dispatch.assert_called_once_with(
(tokenized_obj.mm_inputs,)
)
transport.cancel_for_dispatch.assert_called_once_with(items)
self.assertFalse(state.dispatched)
def test_post_dispatch_failure_does_not_cancel_published_items(self):
from sglang.srt.managers import tokenizer_manager
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
manager = object.__new__(tokenizer_manager.TokenizerManager)
transport = MagicMock()
manager._dispatch_to_scheduler = MagicMock()
state = SimpleNamespace(dispatched=False)
manager.rid_to_state = {"test-request": state}
time_stats = MagicMock()
time_stats.set_api_server_dispatch_finish_time.side_effect = RuntimeError(
"bookkeeping failed"
)
items = [MultimodalDataItem(modality=Modality.IMAGE, feature=torch.arange(2))]
tokenized_obj = SimpleNamespace(
rid="test-request",
mm_inputs=MultimodalProcessorOutput(input_ids=[1], mm_items=items),
time_stats=time_stats,
wrap_pickle_fields=MagicMock(),
)
transport.prepare_for_dispatch.return_value = items
manager.cuda_vmm_feature_transport = transport
with (
patch.object(tokenizer_manager, "wrap_shm_features", lambda obj: obj),
self.assertRaisesRegex(RuntimeError, "bookkeeping failed"),
):
manager._send_one_request(tokenized_obj)
manager._dispatch_to_scheduler.assert_called_once_with(tokenized_obj)
transport.cancel_for_dispatch.assert_not_called()
self.assertTrue(state.dispatched)
def test_prepare_batch_cancels_prior_groups_on_failure(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
transport.wrap_items = MagicMock(
side_effect=[None, RuntimeError("wrap failed")]
)
transport.cancel_for_dispatch = MagicMock()
item_groups = [[object()], [object()]]
mm_inputs_batch = [SimpleNamespace(mm_items=items) for items in item_groups]
with self.assertRaisesRegex(RuntimeError, "wrap failed"):
transport.prepare_for_dispatch(mm_inputs_batch)
self.assertEqual(
transport.wrap_items.call_args_list,
[call(item_groups[0]), call(item_groups[1])],
)
transport.cancel_for_dispatch.assert_called_once_with(item_groups[0])
def test_prepare_batch_returns_flattened_items(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
transport.pool = MagicMock()
transport.wrap_items = MagicMock()
item_groups = [[object()], [object(), object()]]
prepared = transport.prepare_for_dispatch(
[
None,
SimpleNamespace(mm_items=[]),
*(SimpleNamespace(mm_items=items) for items in item_groups),
]
)
self.assertEqual(prepared, item_groups[0] + item_groups[1])
self.assertEqual(
transport.wrap_items.call_args_list,
[call(items) for items in item_groups],
)
def test_engine_shutdown_is_idempotent(self):
from sglang.srt.entrypoints import engine as engine_module
from sglang.srt.entrypoints.engine import Engine
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
manager = object.__new__(TokenizerManager)
transport = object.__new__(CudaVmmFeatureTransport)
pool = MagicMock()
transport.pool = pool
manager.cuda_vmm_feature_transport = transport
manager._subprocess_watchdog = None
engine = object.__new__(Engine)
engine.tokenizer_manager = manager
with patch.object(
engine_module,
"kill_process_tree",
side_effect=RuntimeError("base failed"),
):
for _ in range(2):
with self.assertRaisesRegex(RuntimeError, "base failed"):
engine.shutdown()
self.assertEqual(pool.shutdown.call_count, 2)
self.assertIs(transport.pool, pool)
def test_engine_startup_failure_releases_parent_pool(self):
from sglang.srt.entrypoints import engine as engine_module
from sglang.srt.entrypoints.engine import Engine
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
manager = object.__new__(TokenizerManager)
transport = object.__new__(CudaVmmFeatureTransport)
pool = MagicMock()
transport.pool = pool
manager.cuda_vmm_feature_transport = transport
server_args = SimpleNamespace(
remote_instance_weight_loader_start_seed_via_transfer_engine=False,
reasoning_parser=None,
tool_call_parser=None,
weight_cache_mode=None,
enable_elastic_expert_backup=False,
elastic_ep_backend=None,
node_rank=0,
tokenizer_worker_num=1,
check_server_args=MagicMock(),
)
scheduler_init_result = SimpleNamespace(
all_child_pids=[],
scheduler_infos=[],
wait_for_ready=MagicMock(side_effect=RuntimeError("startup failed")),
engine_info_bootstrap_server=None,
)
with (
patch.object(engine_module, "configure_logger"),
patch.object(engine_module, "_set_envs_and_config"),
patch.object(engine_module, "load_plugins"),
patch.object(
Engine,
"_launch_scheduler_processes",
return_value=(scheduler_init_result, []),
),
patch.object(
Engine, "_launch_detokenizer_subprocesses", return_value=([], [])
),
self.assertRaisesRegex(RuntimeError, "startup failed"),
):
Engine._launch_subprocesses(
server_args=server_args,
init_tokenizer_manager_func=MagicMock(return_value=(manager, object())),
run_scheduler_process_func=MagicMock(),
run_detokenizer_process_func=MagicMock(),
port_args=SimpleNamespace(),
)
pool.shutdown.assert_called_once_with()
self.assertIs(transport.pool, pool)
def test_failed_pool_shutdown_remains_retryable(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmFeatureTransport,
)
transport = object.__new__(CudaVmmFeatureTransport)
pool = MagicMock()
pool.shutdown.side_effect = [RuntimeError("shutdown failed"), None]
transport.pool = pool
with self.assertRaisesRegex(RuntimeError, "shutdown failed"):
transport.shutdown()
self.assertIs(transport.pool, pool)
transport.shutdown()
self.assertIs(transport.pool, pool)
class TestSchedulerMmTransportBoundary(unittest.TestCase):
@staticmethod
def _prepare_scheduler(scheduler):
scheduler.session_controller = SimpleNamespace(maybe_reap=MagicMock())
scheduler._request_dispatcher = MagicMock(return_value=None)
scheduler.flush_wrapper = SimpleNamespace(check_pending=MagicMock())
scheduler.external_corpus_manager = None
def test_materializes_inputs_directly_before_base_dispatch(self):
from sglang.srt.managers import scheduler as scheduler_module
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.server_args = SimpleNamespace(
mm_feature_transport="cuda_vmm",
enable_broadcast_mm_inputs_process=True,
)
self._prepare_scheduler(scheduler)
raw_inputs = object()
materialized = object()
request = SimpleNamespace(mm_inputs=raw_inputs)
with (
patch.object(
scheduler_module, "TokenizedGenerateReqInput", SimpleNamespace
),
patch.object(
scheduler_module.MultimodalInputs,
"from_processor_output",
return_value=materialized,
) as build_inputs,
patch.object(
scheduler, "_process_and_broadcast_mm_inputs"
) as cpu_broadcast,
patch.object(
scheduler_module, "is_health_check_generate_req", return_value=False
),
):
scheduler.process_input_requests([request])
build_inputs.assert_called_once_with(raw_inputs)
self.assertIs(request.mm_inputs, materialized)
scheduler._request_dispatcher.assert_called_once_with(request)
cpu_broadcast.assert_not_called()
def test_materializes_batched_inputs_before_dispatch(self):
from sglang.srt.managers import scheduler as scheduler_module
class TokenizedRequest:
def __init__(self, mm_inputs):
self.mm_inputs = mm_inputs
class BatchRequest:
def __init__(self, batch):
self.batch = batch
def __iter__(self):
return iter(self.batch)
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
self._prepare_scheduler(scheduler)
raw_inputs = [object(), object()]
materialized = [object(), object()]
inner_requests = [TokenizedRequest(value) for value in raw_inputs]
request = BatchRequest(inner_requests)
with (
patch.object(
scheduler_module, "TokenizedGenerateReqInput", TokenizedRequest
),
patch.object(
scheduler_module, "TokenizedEmbeddingReqInput", TokenizedRequest
),
patch.object(
scheduler_module, "BatchTokenizedGenerateReqInput", BatchRequest
),
patch.object(
scheduler_module, "BatchTokenizedEmbeddingReqInput", BatchRequest
),
patch.object(
scheduler_module.MultimodalInputs,
"from_processor_output",
side_effect=materialized,
) as build_inputs,
patch.object(
scheduler_module, "is_health_check_generate_req", return_value=False
),
):
scheduler.process_input_requests([request])
self.assertEqual(
build_inputs.call_args_list,
[call(value) for value in raw_inputs],
)
self.assertEqual(
[inner.mm_inputs for inner in inner_requests],
materialized,
)
scheduler._request_dispatcher.assert_called_once_with(request)
def test_already_materialized_inputs_are_reused(self):
from sglang.srt.managers.schedule_batch import MultimodalInputs
from sglang.srt.managers.scheduler import Scheduler
scheduler = object.__new__(Scheduler)
mm_inputs = MultimodalInputs(mm_items=[])
with patch.object(
scheduler, "_process_and_broadcast_mm_inputs"
) as process_and_broadcast:
self.assertIs(scheduler._get_multimodal_inputs(mm_inputs), mm_inputs)
process_and_broadcast.assert_not_called()
class TestVmmConsumerCount(unittest.TestCase):
def test_proxy_defaults_to_one_consumer(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
proxy = object.__new__(vmm.CudaVmmTensorTransportProxy)
proxy.consumer_count = 4
self.assertEqual(proxy._resolve_consumer_count(None), 1)
self.assertEqual(proxy._resolve_consumer_count(2), 2)
def test_acknowledgement_ranges_include_cp_rank(self):
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmTensorTransportProxy,
)
proxy = object.__new__(CudaVmmTensorTransportProxy)
proxy.consumer_count = 4
with get_parallel().override(
attn_tp_size=2,
attn_tp_rank=1,
attn_cp_size=2,
attn_cp_rank=1,
):
self.assertEqual(proxy._acknowledgement_range(1), (3, 4))
self.assertEqual(proxy._acknowledgement_range(2), (2, 4))
self.assertEqual(proxy._acknowledgement_range(4), (0, 4))
if __name__ == "__main__":
unittest.main()
@@ -259,6 +259,32 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertRaisesRegex(ValueError, "single node"):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=False)
def test_cuda_vmm_rejects_non_nvidia_platforms(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
def test_cuda_vmm_rejects_rust_server(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
with (
envs.SGLANG_RUST_SERVER.override(True),
self.assertRaisesRegex(ValueError, "SGLANG_RUST_SERVER"),
):
server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True)
def test_cuda_vmm_rejects_pipeline_parallelism(self, _mock_is_cuda):
server_args = ServerArgs(
model_path="dummy", mm_feature_transport="cuda_vmm", pp_size=2
)
with self.assertRaisesRegex(ValueError, "pipeline parallelism"):
server_args._handle_multimodal_feature_transport()
class TestMambaCacheStochasticRounding(unittest.TestCase):
def test_rejects_fp32_ssm_cache(self):