[vlm] fix: contain multimodal feature transport failures (#37047)
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user