[vlm] fix: contain multimodal feature transport failures (#37047)

This commit is contained in:
Mick
2026-09-01 13:46:38 +08:00
committed by GitHub
parent 33ed29a0ee
commit ae2bd5728b
16 changed files with 1408 additions and 96 deletions
@@ -495,6 +495,38 @@ class TestStreamOrderedMmFeaturePool(CustomTestCase):
self.assertFalse(pool._recycle_thread.is_alive())
class TestCudaIpcProcessorRollback(CustomTestCase):
def test_partial_wrap_failure_restores_items_and_cancels_proxy(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
)
from sglang.srt.multimodal.transport.cuda_ipc import (
CudaIpcTensorTransportProxy,
)
with patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()):
processor = BaseMultimodalProcessor.__new__(BaseMultimodalProcessor)
processor.use_cuda_ipc = True
processor.cudaipc_mmfeature_pool = MagicMock()
proxy = object.__new__(CudaIpcTensorTransportProxy)
processor._wrap_tensor_for_cuda_ipc = MagicMock(
side_effect=[proxy, RuntimeError("wrap failed")]
)
features = [torch.ones(2), torch.ones(3)]
items = [
MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
for feature in features
]
with self.assertRaisesRegex(RuntimeError, "wrap failed"):
processor._prepare_mm_items_for_transport(items)
processor.cudaipc_mmfeature_pool.cancel_proxy.assert_called_once_with(proxy)
self.assertIs(items[0].feature, features[0])
self.assertIs(items[1].feature, features[1])
class TestPrecomputeHashBeforeCpuTransfer(CustomTestCase):
@staticmethod
def _processor(enabled):
@@ -0,0 +1,305 @@
import unittest
from array import array
from pathlib import Path
from tempfile import TemporaryDirectory
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
import torch.distributed
import torch.multiprocessing
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.managers.io_struct import ( # noqa: E402
BatchTokenizedEmbeddingReqInput,
MMInputsProcessError,
TokenizedEmbeddingReqInput,
)
from sglang.srt.managers.mm_utils import ShmPointerMMData # noqa: E402
from sglang.srt.managers.schedule_batch import ( # noqa: E402
Modality,
MultimodalDataItem,
MultimodalProcessorOutput,
)
from sglang.srt.managers.scheduler import ( # noqa: E402
Scheduler,
_MultimodalInputProcessingError,
)
from sglang.srt.managers.scheduler_components.request_receiver import ( # noqa: E402
SchedulerRequestReceiver,
)
register_cpu_ci(est_time=3, suite="base-a-test-cpu")
class _CloneFailure:
def clone(self):
raise RuntimeError("clone failed")
class _Handle:
def __init__(self, *, fail_unlink: bool = False):
self.closed = False
self.unlinked = False
self.fail_unlink = fail_unlink
def close(self):
self.closed = True
def unlink(self):
if self.fail_unlink:
raise PermissionError("unlink denied")
self.unlinked = True
def _failed_pointer() -> ShmPointerMMData:
pointer = object.__new__(ShmPointerMMData)
pointer.shm_name = "missing-vlm-feature"
pointer.shape = torch.Size([1])
pointer.dtype = torch.float32
pointer.precomputed_hash = None
pointer._shm_handle = None
pointer.tensor = None
pointer._materialization_error = "FileNotFoundError: missing feature"
return pointer
def _successful_pointer() -> ShmPointerMMData:
pointer = object.__new__(ShmPointerMMData)
pointer.shm_name = "unused"
pointer.shape = torch.Size([1])
pointer.dtype = torch.float32
pointer.precomputed_hash = None
pointer._shm_handle = _Handle()
pointer.tensor = torch.ones(1)
pointer._materialization_error = None
return pointer
def _request(feature, rid: str = "vlm-request") -> TokenizedEmbeddingReqInput:
return TokenizedEmbeddingReqInput(
rid=rid,
input_text="",
input_ids=array("q", [1]),
mm_inputs=MultimodalProcessorOutput(
mm_items=[MultimodalDataItem(modality=Modality.IMAGE, feature=feature)]
),
token_type_ids=None,
sampling_params=MagicMock(),
)
def _receiver(tp_size: int = 1) -> SchedulerRequestReceiver:
group = SimpleNamespace(rank=0, ranks=[0], cpu_group=object())
return SchedulerRequestReceiver(
recv_from_tokenizer=None,
recv_from_rpc=None,
recv_skipper=None,
input_blocker=None,
mm_receiver=None,
ps=SimpleNamespace(
pp_rank=0,
tp_size=tp_size,
attn_tp_rank=0,
attn_cp_rank=0,
attn_tp_size=1,
attn_cp_size=1,
),
tp_group=group,
tp_cpu_group=group,
attn_tp_group=group,
attn_tp_cpu_group=group,
attn_cp_group=group,
attn_cp_cpu_group=group,
world_group=group,
server_args=SimpleNamespace(),
model_config=SimpleNamespace(is_multimodal=True),
max_recv_per_poll=-1,
stream_output=lambda *args, **kwargs: None,
get_last_batch=lambda: None,
)
def _run_consensus_rank(rank: int, world_size: int, init_file: str) -> None:
torch.distributed.init_process_group(
backend="gloo",
init_method=Path(init_file).as_uri(),
rank=rank,
world_size=world_size,
)
try:
req = _request(_failed_pointer() if rank == 1 else _successful_pointer())
parallel = SimpleNamespace(enable_dp_attention=False)
receiver = _receiver(tp_size=world_size)
object.__setattr__(receiver, "tp_cpu_group", torch.distributed.group.WORLD)
with (
patch(
"sglang.srt.managers.mm_utils._get_is_default_transport",
return_value=False,
),
patch(
"sglang.srt.managers.mm_utils.get_serving",
return_value=SimpleNamespace(skip_tokenizer_init=False),
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.get_parallel",
return_value=parallel,
),
):
receiver._finalize_shm_features([req])
if not isinstance(req.mm_inputs, MMInputsProcessError):
raise AssertionError(f"rank {rank} did not receive the VLM request error")
finally:
torch.distributed.destroy_process_group()
class TestShmPointerFailureCleanup(unittest.TestCase):
def test_clone_failure_still_unlinks_and_closes(self):
pointer = object.__new__(ShmPointerMMData)
handle = _Handle()
pointer.shm_name = "unused"
pointer._shm_handle = handle
pointer.tensor = _CloneFailure()
pointer._materialization_error = None
with self.assertRaisesRegex(RuntimeError, "clone failed"):
pointer.materialize()
self.assertTrue(handle.unlinked)
self.assertTrue(handle.closed)
self.assertIsNone(pointer._shm_handle)
self.assertIsNone(pointer.tensor)
def test_shm_open_failure_is_deferred_until_materialization(self):
pointer = object.__new__(ShmPointerMMData)
state = {
"shm_name": "missing",
"shape": torch.Size([1]),
"dtype": torch.float32,
"precomputed_hash": None,
}
with patch(
"sglang.srt.managers.mm_utils.shared_memory.SharedMemory",
side_effect=FileNotFoundError("missing"),
):
pointer.__setstate__(state)
with self.assertRaisesRegex(RuntimeError, "FileNotFoundError"):
pointer.materialize()
def test_cleanup_error_does_not_escape_the_request_boundary(self):
pointer = object.__new__(ShmPointerMMData)
handle = _Handle(fail_unlink=True)
pointer.shm_name = "unused"
pointer._shm_handle = handle
pointer.tensor = torch.ones(1)
pointer._materialization_error = None
with self.assertLogs("sglang.utils", level="WARNING"):
result = pointer.materialize()
self.assertTrue(torch.equal(result, torch.ones(1)))
self.assertTrue(handle.closed)
class TestShmRequestFailureConsensus(unittest.TestCase):
def test_real_gloo_group_propagates_one_rank_failure(self):
with TemporaryDirectory() as directory:
init_file = str(Path(directory) / "gloo-init")
torch.multiprocessing.spawn(
_run_consensus_rank,
args=(2, init_file),
nprocs=2,
join=True,
)
def test_local_materialization_failure_becomes_request_error(self):
req = _request(_failed_pointer())
parallel = SimpleNamespace(enable_dp_attention=False)
with (
patch(
"sglang.srt.managers.mm_utils._get_is_default_transport",
return_value=False,
),
patch(
"sglang.srt.managers.mm_utils.get_serving",
return_value=SimpleNamespace(skip_tokenizer_init=False),
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.get_parallel",
return_value=parallel,
),
):
_receiver()._finalize_shm_features([req])
self.assertIsInstance(req.mm_inputs, MMInputsProcessError)
with self.assertRaises(_MultimodalInputProcessingError):
Scheduler._get_multimodal_inputs(object.__new__(Scheduler), req.mm_inputs)
def test_peer_failure_rejects_the_local_request(self):
req = _request(torch.zeros(1))
parallel = SimpleNamespace(enable_dp_attention=False)
def inject_peer_failure(mask, **kwargs):
mask.fill_(1)
with (
patch(
"sglang.srt.managers.scheduler_components.request_receiver.get_parallel",
return_value=parallel,
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.has_shm_features",
return_value=True,
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.unwrap_shm_features"
),
patch("sglang.srt.managers.scheduler_components.request_receiver.barrier"),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.all_reduce",
side_effect=inject_peer_failure,
) as all_reduce,
):
_receiver(tp_size=2)._finalize_shm_features([req])
all_reduce.assert_called_once()
self.assertIsInstance(req.mm_inputs, MMInputsProcessError)
def test_batched_requests_only_reject_the_failed_item(self):
failed_req = _request(torch.zeros(1), rid="failed")
healthy_req = _request(torch.zeros(1), rid="healthy")
batch = BatchTokenizedEmbeddingReqInput(batch=[failed_req, healthy_req])
parallel = SimpleNamespace(enable_dp_attention=False)
def materialize(req):
if req.rid == "failed":
raise RuntimeError("bad shared feature")
with (
patch(
"sglang.srt.managers.scheduler_components.request_receiver.get_parallel",
return_value=parallel,
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.has_shm_features",
return_value=True,
),
patch(
"sglang.srt.managers.scheduler_components.request_receiver.unwrap_shm_features",
side_effect=materialize,
),
):
_receiver()._finalize_shm_features([batch])
self.assertIsInstance(failed_req.mm_inputs, MMInputsProcessError)
self.assertIsInstance(healthy_req.mm_inputs, MultimodalProcessorOutput)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,89 @@
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalInputs,
Req,
)
from sglang.srt.multimodal.transport.cuda_ipc import CudaIpcTensorTransportProxy
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
def _deferred_proxy():
proxy = object.__new__(CudaIpcTensorTransportProxy)
proxy.total_consumer_count = 1
proxy.acknowledge_consumption = MagicMock()
return proxy
def test_release_features_acknowledges_deferred_transport():
proxy = _deferred_proxy()
item = MultimodalDataItem(modality=Modality.IMAGE, feature=proxy)
mm_inputs = MultimodalInputs(mm_items=[item])
mm_inputs.release_features()
proxy.acknowledge_consumption.assert_called_once_with(1)
assert item.feature is None
def test_release_features_keeps_cleanup_error_request_local():
proxy = _deferred_proxy()
proxy.acknowledge_consumption.side_effect = RuntimeError("ack failed")
item = MultimodalDataItem(modality=Modality.IMAGE, feature=proxy)
mm_inputs = MultimodalInputs(mm_items=[item])
mm_inputs.release_features()
assert item.feature is None
def test_request_abort_releases_multimodal_features():
mm_inputs = MagicMock()
req = object.__new__(Req)
req.rid = "rejected-vlm-request"
req.session = None
req.multimodal_inputs = mm_inputs
req.grammar = object()
req.return_logprob = True
req.logprob_start_len = 0
with patch(
"sglang.srt.managers.schedule_batch.get_parallel",
return_value=SimpleNamespace(tp_rank=1),
):
req.set_finish_with_abort("invalid multimodal request")
mm_inputs.release_features.assert_called_once_with()
assert req.multimodal_inputs is None
def test_session_abort_preserves_shared_multimodal_features():
mm_inputs = MagicMock()
req = object.__new__(Req)
req.rid = "rejected-session-turn"
req.session = object()
req.multimodal_inputs = mm_inputs
req.grammar = object()
req.return_logprob = True
req.logprob_start_len = 0
with patch(
"sglang.srt.managers.schedule_batch.get_parallel",
return_value=SimpleNamespace(tp_rank=1),
):
req.set_finish_with_abort("invalid session turn")
mm_inputs.release_features.assert_not_called()
assert req.multimodal_inputs is None
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))
@@ -14,6 +14,12 @@ from unittest.mock import Mock, patch
import torch
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalInputs,
MultimodalProcessorOutput,
)
from sglang.srt.multimodal.transport.cuda_ipc import (
CudaIpcTensorTransportProxy,
MmItemMemoryPool,
@@ -122,6 +128,74 @@ class TestCudaIpcTransport(CustomTestCase):
producer.join(timeout=10)
self.assertEqual(producer.exitcode, 0)
def test_failed_reconstruction_releases_pooled_tensor(self):
ctx = mp.get_context("spawn")
proxy_queue = ctx.Queue()
producer_results = ctx.Queue()
consumer_done = ctx.Event()
producer = ctx.Process(
target=_produce_pooled_tensor,
args=(proxy_queue, consumer_done, producer_results),
)
producer.start()
proxy = None
producer_result = None
original_empty = torch.empty
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 IPC producer failed before sending its proxy: {payload}"
)
output_shape = proxy.proxy_state["ipc_extra"]["recons_shape"]
def fail_destination_allocation(size, *args, **kwargs):
if isinstance(size, (tuple, torch.Size)) and tuple(size) == tuple(
output_shape
):
raise RuntimeError("forced reconstruction failure")
return original_empty(size, *args, **kwargs)
item = MultimodalDataItem(
modality=Modality.IMAGE,
hash=1,
pad_value=1,
feature=proxy,
)
output = MultimodalProcessorOutput(input_ids=[1], mm_items=[item])
with (
patch(
"sglang.srt.multimodal.transport.cuda_ipc.torch.empty",
side_effect=fail_destination_allocation,
),
self.assertRaisesRegex(RuntimeError, "forced reconstruction failure"),
):
MultimodalInputs.from_processor_output(output)
torch.cuda.synchronize()
self.assertTrue(proxy._consumer_acknowledged)
finally:
del proxy
_pool_handle_cache_clear()
gc.collect()
torch.cuda.ipc_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)
self.assertEqual(producer.exitcode, 0)
def test_uncached_mapping_waits_before_proxy_release(self):
proxy = object.__new__(CudaIpcTensorTransportProxy)
proxy.proxy_state = {"ipc_extra": {"use_pool_handle_cache": False}}
@@ -137,6 +211,83 @@ class TestCudaIpcTransport(CustomTestCase):
stream.synchronize.assert_called_once_with()
self.assertIsNone(proxy._pool_storage)
def test_failed_item_batch_releases_undispatched_pool_slice(self):
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
from sglang.srt.multimodal.processors.base_processor import (
BaseMultimodalProcessor,
)
pool = MmItemMemoryPool(
memory_size=1 << 20,
recycle_interval=0.01,
base_gpu_id=0,
consumer_count=4,
)
with patch.object(BaseMultimodalProcessor, "__abstractmethods__", set()):
processor = BaseMultimodalProcessor.__new__(BaseMultimodalProcessor)
processor.use_cuda_ipc = True
processor.use_ipc_pool_handle_cache = True
processor.cudaipc_mmfeature_pool = pool
features = [
torch.ones(16, device="cuda"),
torch.empty(0, device="cuda"),
]
items = [
MultimodalDataItem(modality=Modality.IMAGE, feature=feature)
for feature in features
]
try:
with self.assertRaisesRegex(ValueError, "empty tensor"):
processor._prepare_mm_items_for_transport(items)
deadline = time.monotonic() + 5
while pool.active_lease_count and time.monotonic() < deadline:
time.sleep(0.01)
self.assertEqual(pool.active_lease_count, 0)
self.assertIs(items[0].feature, features[0])
self.assertIs(items[1].feature, features[1])
finally:
pool.shutdown()
def test_rejected_request_releases_unconsumed_pool_slice(self):
ctx = mp.get_context("spawn")
proxy_queue = ctx.Queue()
producer_results = ctx.Queue()
consumer_done = ctx.Event()
producer = ctx.Process(
target=_produce_pooled_tensor,
args=(proxy_queue, consumer_done, producer_results),
)
producer.start()
proxy = None
producer_result = None
try:
proxy, _ = proxy_queue.get(timeout=60)
item = MultimodalDataItem(modality=Modality.IMAGE, feature=proxy)
mm_inputs = MultimodalInputs(mm_items=[item])
mm_inputs.release_features()
torch.cuda.synchronize()
self.assertIsNone(item.feature)
finally:
del proxy
_pool_handle_cache_clear()
gc.collect()
torch.cuda.ipc_collect()
consumer_done.set()
producer.join(timeout=60)
try:
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)
self.assertEqual(producer.exitcode, 0)
if __name__ == "__main__":
unittest.main(verbosity=2)
@@ -11,6 +11,86 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_failed_consumer_reconstruction_releases_remaining_proxies(self):
from sglang.srt.managers.schedule_batch import (
Modality,
MultimodalDataItem,
MultimodalInputs,
MultimodalProcessorOutput,
)
from sglang.srt.multimodal.transport.cuda_ipc import (
CudaIpcTensorTransportProxy,
)
class FakeProxy(CudaIpcTensorTransportProxy):
def __init__(self, *, fail_reconstruct=False, fail_release=False):
self.fail_reconstruct = fail_reconstruct
self.fail_release = fail_release
self.released = False
def reconstruct_on_target_device(self, _device, consumer_count=1):
if self.fail_reconstruct:
raise RuntimeError("reconstruct failed")
return torch.ones(1)
def release_without_reconstruction(self, consumer_count=1):
self.released = True
if self.fail_release:
raise RuntimeError("release failed")
reconstructed = FakeProxy()
failed = FakeProxy(fail_reconstruct=True, fail_release=True)
remaining = FakeProxy()
items = [
MultimodalDataItem(
modality=Modality.IMAGE,
hash=1,
pad_value=1,
feature=reconstructed,
),
MultimodalDataItem(
modality=Modality.IMAGE,
hash=2,
pad_value=2,
feature=failed,
),
MultimodalDataItem(
modality=Modality.IMAGE,
hash=3,
pad_value=3,
feature=remaining,
),
]
output = MultimodalProcessorOutput(input_ids=[1], mm_items=items)
with (
patch(
"sglang.srt.managers.schedule_batch.torch.cuda.current_device",
return_value=0,
),
self.assertRaisesRegex(RuntimeError, "reconstruct failed"),
):
MultimodalInputs.from_processor_output(output)
self.assertIsInstance(items[0].feature, torch.Tensor)
self.assertTrue(failed.released)
self.assertTrue(remaining.released)
def test_abandoned_packed_proxy_releases_shared_owner(self):
from sglang.srt.utils.cuda_vmm_transport_utils import (
CudaVmmPackedTensorTransportProxy,
)
owner = MagicMock()
proxy = object.__new__(CudaVmmPackedTensorTransportProxy)
proxy._packed_owner = owner
proxy._consumer_acknowledged = False
proxy.release_without_reconstruction(consumer_count=2)
owner.acknowledge_consumption.assert_called_once_with(2)
self.assertTrue(proxy._consumer_acknowledged)
def test_partial_pool_release_can_be_retried(self):
from sglang.srt.utils import cuda_vmm_transport_utils as vmm
@@ -520,6 +600,58 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
scheduler.flush_wrapper = SimpleNamespace(check_pending=MagicMock())
scheduler.external_corpus_manager = None
@staticmethod
def _materialize_with_rank_errors(local_exception=None, remote_error=None):
from sglang.srt.managers import scheduler as scheduler_module
class TokenizedRequest:
def __init__(self):
self.mm_inputs = object()
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.dp_tp_cpu_group = object()
request = TokenizedRequest()
def gather_errors(errors, local_error, **_kwargs):
errors[:] = [local_error, remote_error]
materialize = MagicMock(
side_effect=local_exception,
return_value=object(),
)
with (
patch.object(
scheduler_module, "TokenizedGenerateReqInput", TokenizedRequest
),
patch.object(
scheduler_module, "TokenizedEmbeddingReqInput", TokenizedRequest
),
patch.object(
scheduler_module.MultimodalInputs,
"from_processor_output",
materialize,
),
patch.object(
scheduler_module.torch.distributed, "is_available", return_value=True
),
patch.object(
scheduler_module.torch.distributed,
"is_initialized",
return_value=True,
),
patch.object(
scheduler_module.torch.distributed, "get_world_size", return_value=2
),
patch.object(
scheduler_module.torch.distributed,
"all_gather_object",
side_effect=gather_errors,
),
):
errors = scheduler._materialize_cuda_vmm_inputs(request)
return request, errors
def test_materializes_inputs_directly_before_base_dispatch(self):
from sglang.srt.managers import scheduler as scheduler_module
@@ -628,6 +760,222 @@ class TestSchedulerMmTransportBoundary(unittest.TestCase):
process_and_broadcast.assert_not_called()
def test_broadcast_mm_inputs_sends_entry_rank_processing_error(self):
from sglang.srt.managers import scheduler as scheduler_module
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.dp_tp_group = SimpleNamespace(rank_in_group=0, first_rank=0)
scheduler.dp_tp_cpu_group = object()
with (
patch.object(
scheduler_module.MultimodalInputs,
"from_processor_output",
side_effect=ValueError("bad image"),
),
patch.object(
scheduler_module.torch.distributed, "is_available", return_value=True
),
patch.object(
scheduler_module.torch.distributed,
"is_initialized",
return_value=True,
),
patch.object(
scheduler_module.torch.distributed, "get_world_size", return_value=2
),
patch.object(
scheduler_module.torch.distributed, "broadcast_object_list"
) as broadcast,
self.assertRaisesRegex(
scheduler_module._MultimodalInputProcessingError,
"ValueError: bad image",
),
):
scheduler._process_and_broadcast_mm_inputs(object())
payload = broadcast.call_args.args[0][0]
self.assertIn("ValueError: bad image", payload.error)
def test_broadcast_mm_inputs_peer_rank_receives_processing_error(self):
from sglang.srt.managers import scheduler as scheduler_module
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.dp_tp_group = SimpleNamespace(rank_in_group=1, first_rank=0)
scheduler.dp_tp_cpu_group = object()
def receive_error(obj_list, **_kwargs):
obj_list[0] = scheduler_module._MultimodalInputBroadcast(error="bad image")
with (
patch.object(
scheduler_module.MultimodalInputs, "from_processor_output"
) as materialize,
patch.object(
scheduler_module.torch.distributed, "is_available", return_value=True
),
patch.object(
scheduler_module.torch.distributed,
"is_initialized",
return_value=True,
),
patch.object(
scheduler_module.torch.distributed, "get_world_size", return_value=2
),
patch.object(
scheduler_module.torch.distributed,
"broadcast_object_list",
side_effect=receive_error,
),
self.assertRaisesRegex(
scheduler_module._MultimodalInputProcessingError, "bad image"
),
):
scheduler._process_and_broadcast_mm_inputs(object())
materialize.assert_not_called()
def test_embedding_request_aborts_broadcast_processing_error(self):
from sglang.srt.managers import scheduler as scheduler_module
scheduler = object.__new__(scheduler_module.Scheduler)
scheduler.tokenizer = object()
scheduler._maybe_namespace_elastic_radix_cache = MagicMock()
scheduler._add_request_to_queue = MagicMock()
scheduler._get_multimodal_inputs = MagicMock(
side_effect=scheduler_module._MultimodalInputProcessingError("bad image")
)
req = MagicMock()
recv_req = SimpleNamespace(
rid="request-id",
input_text="prompt",
input_ids=[1],
sampling_params=object(),
positional_embed_overrides=None,
token_type_ids=None,
routed_dp_rank=None,
priority=None,
dimensions=None,
lora_id=None,
http_worker_ipc=None,
time_stats=None,
return_pooled_hidden_states=False,
multi_item_delimiter_indices=None,
mm_inputs=object(),
)
with patch.object(scheduler_module, "Req", return_value=req):
scheduler.handle_embedding_request(recv_req)
req.set_finish_with_abort.assert_called_once_with(
"bad image",
status_code=500,
err_type="InternalServerError",
)
scheduler._add_request_to_queue.assert_called_once_with(req)
def test_vmm_materialization_consensus_rejects_any_rank_failure(self):
cases = (
(None, "RuntimeError: remote failure", "rank 1: RuntimeError"),
(ValueError("bad proxy"), None, "rank 0: ValueError: bad proxy"),
)
for local_exception, remote_error, expected in cases:
with self.subTest(expected=expected):
request, errors = self._materialize_with_rank_errors(
local_exception, remote_error
)
self.assertIn(expected, errors[0])
self.assertIsNone(request.mm_inputs)
def test_vmm_batch_dispatches_good_and_failed_requests_individually(self):
from sglang.srt.managers import scheduler as scheduler_module
class TokenizedRequest:
pass
class EmbeddingRequest:
pass
class BatchRequest:
def __init__(self, requests):
self.requests = requests
def __iter__(self):
return iter(self.requests)
scheduler = object.__new__(scheduler_module.Scheduler)
self._publish(mm_feature_transport="cuda_vmm")
self._prepare_scheduler(scheduler)
scheduler.is_fully_idle = MagicMock(return_value=True)
scheduler.return_health_check_ipcs = []
scheduler.handle_generate_request = MagicMock()
scheduler.handle_embedding_request = MagicMock()
scheduler._materialize_cuda_vmm_inputs = MagicMock(
return_value=[None, "reconstruction failed"]
)
requests = [TokenizedRequest(), TokenizedRequest()]
batch = BatchRequest(requests)
with (
patch.object(
scheduler_module, "TokenizedGenerateReqInput", TokenizedRequest
),
patch.object(
scheduler_module, "TokenizedEmbeddingReqInput", EmbeddingRequest
),
patch.object(
scheduler_module, "BatchTokenizedGenerateReqInput", BatchRequest
),
patch.object(scheduler_module, "BatchTokenizedEmbeddingReqInput", tuple),
patch.object(
scheduler_module, "is_health_check_generate_req", return_value=False
),
):
scheduler.process_input_requests([batch])
self.assertEqual(
scheduler.handle_generate_request.call_args_list,
[
call(requests[0], mm_input_error=None),
call(requests[1], mm_input_error="reconstruction failed"),
],
)
scheduler.handle_embedding_request.assert_not_called()
scheduler._request_dispatcher.assert_not_called()
def test_vmm_materialization_abort_reports_internal_error(self):
from sglang.srt.managers import schedule_batch
req = object.__new__(schedule_batch.Req)
req.rid = "request-id"
req.multimodal_inputs = schedule_batch.MultimodalInputs(mm_items=[])
req.session = None
req.grammar = object()
req.origin_input_ids = [1, 2]
req.return_logprob = True
req.logprob_start_len = 0
req.to_finish = None
with patch.object(
schedule_batch, "get_parallel", return_value=SimpleNamespace(tp_rank=1)
):
req.set_finish_with_abort(
"reconstruction failed",
status_code=500,
err_type="InternalServerError",
)
self.assertEqual(
req.to_finish.to_json(),
{
"type": "abort",
"message": "reconstruction failed",
"status_code": 500,
"err_type": "InternalServerError",
},
)
self.assertIsNone(req.multimodal_inputs)
class TestVmmConsumerCount(unittest.TestCase):
def test_proxy_defaults_to_one_consumer(self):