From ae2bd5728b79cdbce5efbc592743c457dde2db62 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 1 Sep 2026 13:46:38 +0800 Subject: [PATCH] [vlm] fix: contain multimodal feature transport failures (#37047) --- python/sglang/srt/managers/io_struct.py | 6 + python/sglang/srt/managers/mm_utils.py | 120 ++++-- python/sglang/srt/managers/schedule_batch.py | 65 +++- python/sglang/srt/managers/scheduler.py | 191 ++++++++-- .../scheduler_components/request_receiver.py | 81 +++- .../multimodal/processors/base_processor.py | 37 +- .../srt/multimodal/processors/kimi_k3.py | 4 +- .../srt/multimodal/processors/moss_vl.py | 9 +- .../srt/multimodal/transport/cuda_ipc.py | 15 + .../srt/multimodal/transport/memory_pool.py | 44 +++ .../srt/utils/cuda_vmm_transport_utils.py | 7 + .../unit/managers/test_mm_process_config.py | 32 ++ .../managers/test_mm_shm_error_consensus.py | 305 +++++++++++++++ .../managers/test_multimodal_abort_cleanup.py | 89 +++++ .../multimodal/test_cuda_ipc_transport.py | 151 ++++++++ .../multimodal/test_gpu_feature_transport.py | 348 ++++++++++++++++++ 16 files changed, 1408 insertions(+), 96 deletions(-) create mode 100644 test/registered/unit/managers/test_mm_shm_error_consensus.py create mode 100644 test/registered/unit/managers/test_multimodal_abort_cleanup.py diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 42e37971c..d29ff305b 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -105,6 +105,12 @@ class BaseBatchReq(msgspec.Struct, tag=True, kw_only=True, array_like=True): return msgspec_struct_pydantic_core_schema(cls, handler) +class MMInputsProcessError(msgspec.Struct, frozen=True): + """Request-local multimodal input failure produced after tokenizer fanout.""" + + message: str + + class BeamSearchOutput(BaseBatchReq, kw_only=True): sequences: List[BeamSearchSequence] diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 3a957ae7f..4a9d7a29e 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -36,6 +36,7 @@ from sglang.srt.managers.schedule_batch import ( CudaIpcTensorTransportProxy, Modality, MultimodalInputs, + MultimodalProcessorOutput, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.multimodal.transport import ( @@ -1285,6 +1286,9 @@ class ShmPointerMMData: """ def __init__(self, tensor: torch.Tensor, precomputed_hash: Optional[int] = None): + self._shm_handle = None + self.tensor = None + self._materialization_error = None if not tensor.is_cpu: tensor = tensor.cpu() if not tensor.is_contiguous(): @@ -1311,7 +1315,6 @@ class ShmPointerMMData: raise self.shm_name = shm.name shm.close() - self._shm_handle = None def __getstate__(self): return { @@ -1326,27 +1329,78 @@ class ShmPointerMMData: self.shape = state["shape"] self.dtype = state["dtype"] self.precomputed_hash = state.get("precomputed_hash") - self._shm_handle = shared_memory.SharedMemory(name=self.shm_name) - # Zero-copy view into shared memory (no clone, no unlink) - self.tensor = torch.frombuffer(self._shm_handle.buf, dtype=self.dtype).reshape( - self.shape - ) + self._shm_handle = None + self.tensor = None + self._materialization_error = None + + # keep deserialization infallible so all TP ranks finish the broadcast + handle = None + tensor = None + try: + handle = shared_memory.SharedMemory(name=self.shm_name) + tensor = torch.frombuffer(handle.buf, dtype=self.dtype) + self.tensor = tensor.reshape(self.shape) + self._shm_handle = handle + except Exception as error: + tensor = None + if handle is not None: + try: + handle.close() + except Exception: + logger.warning( + "Failed to close a malformed multimodal SHM handle", + exc_info=True, + ) + self._materialization_error = f"{type(error).__name__}: {error}" def materialize(self) -> torch.Tensor: """Clone tensor from shm to owned memory, then release shm handle.""" - tensor = self.tensor.clone() - if self._shm_handle is not None: - self._shm_handle.close() + try: + if self._materialization_error is not None: + raise RuntimeError(self._materialization_error) + return self.tensor.clone() + finally: + self.close_and_unlink() + + def close_and_unlink(self) -> None: + """Release this rank's view and unlink the shared feature segment.""" + handle = self._shm_handle + self._shm_handle = None + self.tensor = None + if handle is None: try: - self._shm_handle.unlink() + handle = shared_memory.SharedMemory(name=self.shm_name) except FileNotFoundError: - pass # Another rank already unlinked - self._shm_handle = None - return tensor + return + except OSError: + logger.warning( + "Failed to reopen a multimodal SHM segment for cleanup", + exc_info=True, + ) + return + try: + try: + handle.unlink() + except FileNotFoundError: + pass + except OSError: + logger.warning( + "Failed to unlink a multimodal SHM segment", + exc_info=True, + ) + finally: + try: + handle.close() + except Exception: + logger.warning( + "Failed to close a multimodal SHM handle", + exc_info=True, + ) def __del__(self): # Only close; never unlink. Unlinking is materialize()'s job. - if getattr(self, "_shm_handle", None) is not None: + if self._shm_handle is not None: + self.tensor = None self._shm_handle.close() self._shm_handle = None @@ -1427,10 +1481,9 @@ def has_shm_features(recv_reqs): if isinstance(req, BaseBatchReq): if has_shm_features(req.batch): return True - elif ( - isinstance(req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)) - and req.mm_inputs - ): + elif isinstance( + req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput) + ) and isinstance(req.mm_inputs, (MultimodalProcessorOutput, MultimodalInputs)): for item in req.mm_inputs.mm_items: if _feature_has_shm(item.feature): return True @@ -1439,6 +1492,30 @@ def has_shm_features(recv_reqs): return False +def _discard_tensor_or_list(value) -> None: + if isinstance(value, ShmPointerMMData): + value.close_and_unlink() + elif isinstance(value, (list, tuple)): + for tensor in value: + if isinstance(tensor, ShmPointerMMData): + tensor.close_and_unlink() + + +def discard_shm_features(obj) -> None: + """Release SHM features that will not be consumed by this request.""" + if isinstance(obj, BaseBatchReq): + for sub_obj in obj.batch: + discard_shm_features(sub_obj) + return + if not isinstance(obj, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)): + return + if not isinstance(obj.mm_inputs, (MultimodalProcessorOutput, MultimodalInputs)): + return + for item in obj.mm_inputs.mm_items: + _discard_tensor_or_list(item.feature) + _discard_tensor_or_list(item.precomputed_embeddings) + + def _unwrap_tensor_or_list(value): """Restore ShmPointerMMData wrappers back into standard torch.Tensors.""" if isinstance(value, ShmPointerMMData): @@ -1464,10 +1541,9 @@ def unwrap_shm_features(obj): unwrap_shm_features(sub_obj) return obj # Handle single requests - if ( - isinstance(obj, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)) - and obj.mm_inputs - ): + if isinstance( + obj, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput) + ) and isinstance(obj.mm_inputs, (MultimodalProcessorOutput, MultimodalInputs)): for item in obj.mm_inputs.mm_items: if item.feature is not None: item.feature = _unwrap_tensor_or_list(item.feature) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index bb9038d8c..0552fbde2 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -509,6 +509,22 @@ class MultimodalDataItem(msgspec.Struct, kw_only=True, dict=True, array_like=Tru ) self.feature.acknowledge_consumption(consumer_count) + def release_transport_proxies(self, consumer_count: int = 1) -> None: + """Best-effort release of proxies left by an abandoned request.""" + values = [self.feature, self.precomputed_embeddings] + values.extend(self.model_specific_data.values()) + for value in values: + if not isinstance(value, CudaIpcTensorTransportProxy): + continue + count = self._resolve_transport_consumer_count(value, consumer_count) + try: + value.release_without_reconstruction(count) + except Exception: + logger.warning( + "Failed to release an abandoned multimodal transport proxy", + exc_info=True, + ) + @staticmethod def _resolve_transport_consumer_count(proxy, requested_count: int) -> int: """Clamp a group acknowledgement to the proxy's actual consumer set.""" @@ -643,7 +659,18 @@ class MultimodalInputs: def release_features(self): """Release feature tensors to free GPU memory.""" for item in self.mm_items: - item.feature = None + try: + # A request can be rejected before a deferred GPU feature is + # reconstructed. Acknowledge that transport lease before the + # proxy is dropped so the tokenizer pool can reuse its slice. + item.acknowledge_deferred_cuda_ipc_feature() + except Exception: + logger.warning( + "Failed to release an unused multimodal feature transport", + exc_info=True, + ) + finally: + item.feature = None @staticmethod def from_processor_output(obj: MultimodalProcessorOutput): @@ -653,14 +680,19 @@ class MultimodalInputs: # try reconstructing from cuda-ipc reconstruct_device = None - for mm_item in mm_items: - if ( - mm_item.has_cuda_ipc_proxy() - and not mm_item.can_defer_cuda_ipc_feature_reconstruction() - ): - if reconstruct_device is None: - reconstruct_device = torch.cuda.current_device() - mm_item.reconstruct(reconstruct_device) + try: + for mm_item in mm_items: + if ( + mm_item.has_cuda_ipc_proxy() + and not mm_item.can_defer_cuda_ipc_feature_reconstruction() + ): + if reconstruct_device is None: + reconstruct_device = torch.cuda.current_device() + mm_item.reconstruct(reconstruct_device) + except BaseException: + for mm_item in mm_items: + mm_item.release_transport_proxies() + raise if envs.SGLANG_MM_BUFFER_SIZE_MB.get() > 0: # Multi-modal feature hashing optimization: @@ -1892,9 +1924,18 @@ class Req(ReqDllmMixin): logger.info(f"{prefix}: {self.time_stats.convert_to_duration()}") self.has_log_time_stats = True - def set_finish_with_abort(self, error_msg: str): + def set_finish_with_abort( + self, + error_msg: str, + status_code: int = HTTPStatus.BAD_REQUEST, + err_type: str = "BadRequestError", + ): if get_parallel().tp_rank == 0: logger.error(f"{error_msg}, {self.rid=}") + # Session requests share historical multimodal inputs with their prior + # request. The session owns and releases those features when it closes. + if self.multimodal_inputs is not None and self.session is None: + self.multimodal_inputs.release_features() self.multimodal_inputs = None self.grammar = None self.origin_input_ids = array( @@ -1902,9 +1943,7 @@ class Req(ReqDllmMixin): ) # set it to one token to skip the long prefill self.return_logprob = False self.logprob_start_len = -1 - self.to_finish = FINISH_ABORT( - error_msg, HTTPStatus.BAD_REQUEST, "BadRequestError" - ) + self.to_finish = FINISH_ABORT(error_msg, status_code, err_type) def update_reasoning_tokens(self, token_id, think_end_ids): if self._is_reasoning_over: diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index c4b1eab24..bccbf46ad 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -151,6 +151,7 @@ from sglang.srt.managers.io_struct import ( LoadLoRAAdapterFromTensorsReqOutput, LoadLoRAAdapterReqInput, LoadLoRAAdapterReqOutput, + MMInputsProcessError, OpenSessionReqInput, PauseGenerationReqInput, ProfileReq, @@ -373,6 +374,16 @@ STEP_MAX_US = 2_000_000 LOAD_STALL_REFRESH_S = 0.05 +@dataclasses.dataclass(frozen=True) +class _MultimodalInputBroadcast: + inputs: Optional[MultimodalInputs] = None + error: Optional[str] = None + + +class _MultimodalInputProcessingError(RuntimeError): + pass + + def _accumulate_decode_moment( totals: list[float], batch_size: int, @@ -1955,11 +1966,12 @@ class Scheduler( def process_input_requests(self, recv_reqs: List): now = time.monotonic() self.session_controller.maybe_reap(now) - if get_mm().mm_feature_transport == "cuda_vmm": - for recv_req in recv_reqs: - self._materialize_cuda_vmm_inputs(recv_req) for recv_req in recv_reqs: + vmm_errors = None + if get_mm().mm_feature_transport == "cuda_vmm": + vmm_errors = self._materialize_cuda_vmm_inputs(recv_req) + # Skip health check when server is busy — ongoing requests already carry health info. if is_health_check_generate_req(recv_req) and not self.is_fully_idle( for_health_check=True @@ -1969,6 +1981,10 @@ class Scheduler( ) continue + if vmm_errors is not None and any(vmm_errors): + self._dispatch_tokenized_mm_requests(recv_req, vmm_errors) + continue + output = self._request_dispatcher(recv_req) if output is not None: if self.rust_server is not None: @@ -1986,26 +2002,87 @@ class Scheduler( if self.external_corpus_manager is not None: self.external_corpus_manager.check_pending_load() - def _materialize_cuda_vmm_inputs(self, recv_req): - """Release VMM slices before request handling can reject the request.""" + @staticmethod + def _tokenized_requests(recv_req): if isinstance( recv_req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput) ): - tokenized_reqs = (recv_req,) - elif isinstance( + return (recv_req,) + if isinstance( recv_req, (BatchTokenizedGenerateReqInput, BatchTokenizedEmbeddingReqInput), ): - tokenized_reqs = recv_req - else: - return + return tuple(recv_req) + return () + def _gather_vmm_materialization_errors( + self, local_error: Optional[str] + ) -> List[Optional[str]]: + if not ( + torch.distributed.is_available() + and torch.distributed.is_initialized() + and self.dp_tp_cpu_group is not None + ): + return [local_error] + + world_size = torch.distributed.get_world_size(group=self.dp_tp_cpu_group) + errors = [None] * world_size + torch.distributed.all_gather_object( + errors, + local_error, + group=self.dp_tp_cpu_group, + ) + return errors + + def _materialize_cuda_vmm_inputs(self, recv_req) -> Optional[List[Optional[str]]]: + """Materialize each request and agree on failures across TP ranks.""" + tokenized_reqs = self._tokenized_requests(recv_req) + if not tokenized_reqs: + return None + + request_errors = [] for tokenized_req in tokenized_reqs: - if tokenized_req.mm_inputs is not None and not isinstance( - tokenized_req.mm_inputs, MultimodalInputs - ): - tokenized_req.mm_inputs = MultimodalInputs.from_processor_output( - tokenized_req.mm_inputs + local_error = None + try: + if tokenized_req.mm_inputs is not None and not isinstance( + tokenized_req.mm_inputs, MultimodalInputs + ): + tokenized_req.mm_inputs = MultimodalInputs.from_processor_output( + tokenized_req.mm_inputs + ) + except Exception as error: + local_error = f"{type(error).__name__}: {error}" + + rank_errors = self._gather_vmm_materialization_errors(local_error) + failed_ranks = [ + rank for rank, error in enumerate(rank_errors) if error is not None + ] + if failed_ranks: + details = "; ".join( + f"rank {rank}: {rank_errors[rank]}" for rank in failed_ranks + ) + error_msg = f"Multimodal feature reconstruction failed ({details})" + logger.error(error_msg) + tokenized_req.mm_inputs = None + request_errors.append(error_msg) + else: + request_errors.append(None) + return request_errors + + def _dispatch_tokenized_mm_requests( + self, recv_req, errors: List[Optional[str]] + ) -> None: + tokenized_reqs = self._tokenized_requests(recv_req) + if len(tokenized_reqs) != len(errors): + raise RuntimeError("VMM materialization results do not match requests") + for tokenized_req, error in zip(tokenized_reqs, errors, strict=True): + if isinstance(tokenized_req, TokenizedGenerateReqInput): + self.handle_generate_request(tokenized_req, mm_input_error=error) + elif isinstance(tokenized_req, TokenizedEmbeddingReqInput): + self.handle_embedding_request(tokenized_req, mm_input_error=error) + else: + raise TypeError( + f"Unsupported tokenized request type: {type(tokenized_req).__name__}" ) def init_profiler(self) -> None: @@ -2346,6 +2423,11 @@ class Scheduler( Returns: MultimodalInputs | None + + Raises: + _MultimodalInputProcessingError: The entry rank could not build the + request's multimodal inputs. The same error is broadcast to all + ranks before it is raised. """ if raw_mm_inputs is None: return None @@ -2371,18 +2453,29 @@ class Scheduler( # Since the Scheduler is single-threaded, any large CPU cost will impact # handling of other messages. For example, CPU hits 99.9% can significantly # increase the CUDA kernel launch time. + result = None if self.dp_tp_group.rank_in_group == 0: - # Only the entry rank materializes once from dict. - image_inputs = MultimodalInputs.from_processor_output(raw_mm_inputs) - # Broadcast to other TP ranks (use src=0 within the group). + try: + result = _MultimodalInputBroadcast( + inputs=MultimodalInputs.from_processor_output(raw_mm_inputs) + ) + except Exception as error: + result = _MultimodalInputBroadcast( + error=( + "Multimodal input processing failed on the TP entry rank: " + f"{type(error).__name__}: {error}" + ) + ) + + # Broadcast either the prepared inputs or the request-local error. if group_world_size > 1: - obj_list = [image_inputs] + obj_list = [result] torch.distributed.broadcast_object_list( obj_list, src=self.dp_tp_group.first_rank, group=self.dp_tp_cpu_group, ) - image_inputs = obj_list[0] + result = obj_list[0] else: # Non-entry ranks: receive if group size > 1; otherwise materialize locally. if group_world_size > 1: @@ -2392,13 +2485,19 @@ class Scheduler( src=self.dp_tp_group.first_rank, group=self.dp_tp_cpu_group, ) - image_inputs = obj_list[0] + result = obj_list[0] else: - image_inputs = MultimodalInputs.from_processor_output(raw_mm_inputs) + result = _MultimodalInputBroadcast( + inputs=MultimodalInputs.from_processor_output(raw_mm_inputs) + ) - return image_inputs + if result.error is not None: + raise _MultimodalInputProcessingError(result.error) + return result.inputs def _get_multimodal_inputs(self, mm_inputs): + if isinstance(mm_inputs, MMInputsProcessError): + raise _MultimodalInputProcessingError(mm_inputs.message) if isinstance(mm_inputs, MultimodalInputs): return mm_inputs @@ -2487,6 +2586,8 @@ class Scheduler( def handle_generate_request( self, recv_req: TokenizedGenerateReqInput, + *, + mm_input_error: Optional[str] = None, ): # Route: normal request / session request / session-not-found session_id = ( @@ -2635,6 +2736,16 @@ class Scheduler( self._maybe_namespace_elastic_radix_cache(req) + if mm_input_error is not None: + req.set_finish_with_abort( + mm_input_error, + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + err_type="InternalServerError", + ) + self.init_req_max_new_tokens(req) + self._add_request_to_queue(req) + return + if self.spec_algorithm.is_dflash_family(): error_msg = validate_dflash_request(req, self.enable_overlap) if error_msg is not None: @@ -2694,7 +2805,17 @@ class Scheduler( # Handle multimodal inputs if recv_req.mm_inputs is not None: - image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs) + try: + image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs) + except _MultimodalInputProcessingError as error: + req.set_finish_with_abort( + str(error), + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + err_type="InternalServerError", + ) + self.init_req_max_new_tokens(req) + self._add_request_to_queue(req) + return SessionController.adjust_mm_offsets(recv_req, req, image_inputs) @@ -3025,6 +3146,8 @@ class Scheduler( def handle_embedding_request( self, recv_req: TokenizedEmbeddingReqInput, + *, + mm_input_error: Optional[str] = None, ): req = Req( recv_req.rid, @@ -3045,9 +3168,27 @@ class Scheduler( req.tokenizer = self.tokenizer self._maybe_namespace_elastic_radix_cache(req) + if mm_input_error is not None: + req.set_finish_with_abort( + mm_input_error, + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + err_type="InternalServerError", + ) + self._add_request_to_queue(req) + return + # Handle multimodal inputs if recv_req.mm_inputs is not None: - image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs) + try: + image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs) + except _MultimodalInputProcessingError as error: + req.set_finish_with_abort( + str(error), + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + err_type="InternalServerError", + ) + self._add_request_to_queue(req) + return # Expand a single image token into multiple dummy tokens for receiving image embeddings # The `pad_input_ids_func` is model-specific and may be None for # embedding models or models not requiring special padding. diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index cbb6ebb18..a18dac607 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging from dataclasses import dataclass from http import HTTPStatus from typing import ( @@ -11,19 +12,22 @@ from typing import ( Union, ) +import torch import zmq -from torch.distributed import barrier +from torch.distributed import ReduceOp, all_reduce, barrier from sglang.srt.disaggregation.utils import prepare_abort from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ( BatchTokenizedEmbeddingReqInput, BatchTokenizedGenerateReqInput, + MMInputsProcessError, TokenizedEmbeddingReqInput, TokenizedGenerateReqInput, sock_recv, ) from sglang.srt.managers.mm_utils import ( + discard_shm_features, has_shm_features, unwrap_shm_features, ) @@ -44,6 +48,8 @@ if TYPE_CHECKING: ScriptedTokenizerRecvProxy, ) +logger = logging.getLogger(__name__) + @dataclass(kw_only=True, slots=True, frozen=True) class SchedulerRequestReceiver: @@ -249,24 +255,63 @@ class SchedulerRequestReceiver: return recv_reqs def _finalize_shm_features(self, recv_reqs: Optional[List]) -> None: - # Unwrap shared memory features AFTER all broadcasts complete, - # so that ShmPointerMMData metadata (not full tensor data) is what - # gets serialized during broadcast_pyobj. - if recv_reqs: - if self.model_config.is_multimodal and has_shm_features(recv_reqs): - # The broadcast source returns with its original objects while - # peer ranks may still be unpickling ShmPointerMMData - # (-> shm_open). Synchronize the same CPU groups that carried - # SHM-backed work requests before materialize() unlinks them. - if get_parallel().enable_dp_attention: - if self.ps.attn_tp_size > 1: - barrier(group=self.attn_tp_cpu_group) - if self.ps.attn_cp_size > 1: - barrier(group=self.attn_cp_cpu_group) - elif self.ps.tp_size > 1: - barrier(group=self.tp_cpu_group) - for req in recv_reqs: + """Materialize SHM features or mark the request failed on every rank.""" + if not recv_reqs or not self.model_config.is_multimodal: + return + + tokenized_reqs = [] + for req in recv_reqs: + if isinstance(req, (TokenizedGenerateReqInput, TokenizedEmbeddingReqInput)): + tokenized_reqs.append(req) + elif isinstance( + req, + (BatchTokenizedGenerateReqInput, BatchTokenizedEmbeddingReqInput), + ): + tokenized_reqs.extend(req.batch) + if not tokenized_reqs or not has_shm_features(tokenized_reqs): + return + + # 1. wait until every rank has opened the shared feature segments + parallel = get_parallel() + if parallel.enable_dp_attention: + if self.ps.attn_tp_size > 1: + barrier(group=self.attn_tp_cpu_group) + if self.ps.attn_cp_size > 1: + barrier(group=self.attn_cp_cpu_group) + elif self.ps.tp_size > 1: + barrier(group=self.tp_cpu_group) + + # 2. materialize independently so one bad VLM request does not stop the loop + failed = torch.zeros(len(tokenized_reqs), dtype=torch.int32) + for index, req in enumerate(tokenized_reqs): + if not has_shm_features([req]): + continue + try: unwrap_shm_features(req) + except Exception: + logger.exception( + "Failed to materialize shared-memory multimodal features for rid=%s", + req.rid, + ) + discard_shm_features(req) + failed[index] = 1 + + # 3. all ranks reject the same requests before entering model collectives + if parallel.enable_dp_attention: + if self.ps.attn_tp_size > 1: + all_reduce(failed, op=ReduceOp.MAX, group=self.attn_tp_cpu_group) + if self.ps.attn_cp_size > 1: + all_reduce(failed, op=ReduceOp.MAX, group=self.attn_cp_cpu_group) + elif self.ps.tp_size > 1: + all_reduce(failed, op=ReduceOp.MAX, group=self.tp_cpu_group) + + error = MMInputsProcessError( + "Failed to materialize shared-memory multimodal features on a scheduler rank." + ) + for index, req in enumerate(tokenized_reqs): + if failed[index].item(): + discard_shm_features(req) + req.mm_inputs = error def _split_work_and_control_reqs(self, recv_reqs: List): work_reqs = [ diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index db147375d..919bd1b3c 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -38,6 +38,7 @@ from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecuto from sglang.srt.multimodal.transport.cuda_ipc import ( MM_FEATURE_CACHE_SIZE, MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL, + CudaIpcTensorTransportProxy, MmItemMemoryPool, get_mm_feature_pool_size_per_worker, ) @@ -1875,19 +1876,41 @@ class BaseMultimodalProcessor(ABC): def _prepare_mm_items_for_transport( self, mm_items: List[MultimodalDataItem] ) -> List[MultimodalDataItem]: - """Wrap final GPU features for dispatch to the scheduler.""" + """Wrap final GPU features, rolling back every lease if one wrap fails.""" if not self.use_cuda_ipc: return mm_items # Pool misses fall back to plain CPU tensors. The scheduler copies out # and releases each successful pool slice. - for item in mm_items: - if isinstance(item.feature, torch.Tensor): - item.feature = self._wrap_tensor_for_cuda_ipc(item.feature) - if isinstance(item.precomputed_embeddings, torch.Tensor): - item.precomputed_embeddings = self._wrap_tensor_for_cuda_ipc( - item.precomputed_embeddings + updates = [] + try: + for item in mm_items: + fields = ( + ("feature", item.feature), + ("precomputed_embeddings", item.precomputed_embeddings), ) + for field, tensor in fields: + if not isinstance(tensor, torch.Tensor): + continue + wrapped = self._wrap_tensor_for_cuda_ipc(tensor) + setattr(item, field, wrapped) + updates.append((item, field, tensor, wrapped)) + except BaseException as error: + rollback_errors = [] + for item, field, tensor, wrapped in reversed(updates): + try: + if isinstance(wrapped, CudaIpcTensorTransportProxy): + self.cudaipc_mmfeature_pool.cancel_proxy(wrapped) + except BaseException as rollback_error: + rollback_errors.append(rollback_error) + finally: + setattr(item, field, tensor) + if rollback_errors: + error.add_note( + f"{len(rollback_errors)} CUDA IPC rollback operation(s) also failed" + ) + raise error from rollback_errors[0] + raise return mm_items async def process_and_combine_mm_data_async( diff --git a/python/sglang/srt/multimodal/processors/kimi_k3.py b/python/sglang/srt/multimodal/processors/kimi_k3.py index 9d3edf3a1..1dfa42408 100644 --- a/python/sglang/srt/multimodal/processors/kimi_k3.py +++ b/python/sglang/srt/multimodal/processors/kimi_k3.py @@ -645,8 +645,6 @@ class KimiK3ImageProcessor( model_specific_data=model_specific_data, ) item.set_hash(artifact.feature_hash) - if self.use_cuda_ipc and isinstance(item.feature, torch.Tensor): - item.feature = self._wrap_tensor_for_cuda_ipc(item.feature) if self.keep_mm_features_on_device and item.feature is not None: item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = ( True @@ -655,7 +653,7 @@ class KimiK3ImageProcessor( return MultimodalProcessorOutput( input_ids=input_ids.tolist(), - mm_items=items, + mm_items=self._prepare_mm_items_for_transport(items), im_token_id=self.mm_tokens.image_token_id, ) diff --git a/python/sglang/srt/multimodal/processors/moss_vl.py b/python/sglang/srt/multimodal/processors/moss_vl.py index 7da38c58b..50ea9eaab 100644 --- a/python/sglang/srt/multimodal/processors/moss_vl.py +++ b/python/sglang/srt/multimodal/processors/moss_vl.py @@ -582,14 +582,7 @@ class MossVLImageProcessor(SGLangBaseProcessor): if mm_items and vision_token_info: mm_items[0].set("vision_token_info", vision_token_info[0]) - if self.use_cuda_ipc: - for item in mm_items: - if isinstance(item.feature, torch.Tensor): - item.feature = self._wrap_tensor_for_cuda_ipc(item.feature) - if isinstance(item.precomputed_embeddings, torch.Tensor): - item.precomputed_embeddings = self._wrap_tensor_for_cuda_ipc( - item.precomputed_embeddings - ) + mm_items = self._prepare_mm_items_for_transport(mm_items) return MultimodalProcessorOutput( input_ids=input_ids.tolist(), diff --git a/python/sglang/srt/multimodal/transport/cuda_ipc.py b/python/sglang/srt/multimodal/transport/cuda_ipc.py index d1d30af04..1f8f14e2a 100644 --- a/python/sglang/srt/multimodal/transport/cuda_ipc.py +++ b/python/sglang/srt/multimodal/transport/cuda_ipc.py @@ -150,6 +150,17 @@ class MmItemMemoryPool: use_pool_handle_cache=use_pool_handle_cache, ) + def cancel_proxy(self, proxy: "CudaIpcTensorTransportProxy") -> None: + """Return a published slice when its request was never dispatched.""" + ipc_extra = proxy.proxy_state["ipc_extra"] + if tuple(ipc_extra["pool_handle"]) != tuple(self._pool_ipc_handle): + raise RuntimeError("CUDA IPC proxy does not belong to this pool") + self._pool.cancel_lease( + ready_byte_offset=proxy.ready_byte_offset, + ack_byte_offset=proxy.ack_byte_offset, + generation=proxy.generation, + ) + def _warn_pool_full_once(self, nbytes: int): if self._pool_full_warned: return @@ -310,6 +321,10 @@ class CudaIpcTensorTransportProxy(StreamOrderedPoolConsumerMixin): ) self._retain_storage_until_stream_completes(storage, device_id) + def release_without_reconstruction(self, consumer_count: int = 1) -> None: + """Release a pool slice when its request abandons this proxy.""" + self.acknowledge_consumption(consumer_count) + def reconstruct_on_target_device( self, rebuild_device_idx, diff --git a/python/sglang/srt/multimodal/transport/memory_pool.py b/python/sglang/srt/multimodal/transport/memory_pool.py index 1ba0547c0..a594c1c3e 100644 --- a/python/sglang/srt/multimodal/transport/memory_pool.py +++ b/python/sglang/srt/multimodal/transport/memory_pool.py @@ -362,6 +362,50 @@ class StreamOrderedMmFeaturePool: raise return lease, destination + def cancel_lease( + self, + *, + ready_byte_offset: int, + ack_byte_offset: int, + generation: int, + ) -> None: + """Acknowledge every consumer for a lease that was not dispatched.""" + slot_stride = self.control_words_per_slot * CONTROL_WORD_BYTES + if ( + ready_byte_offset % slot_stride != 0 + or ack_byte_offset != ready_byte_offset + CONTROL_WORD_BYTES + ): + raise RuntimeError(f"Invalid {self.transport_name} pool lease offsets") + slot = ready_byte_offset // slot_stride + + with self._lock: + lease = self._occupied.get(slot) + if ( + lease is None + or lease.generation != generation + or lease.ready_byte_offset != ready_byte_offset + or lease.ack_byte_offset != ack_byte_offset + ): + raise RuntimeError( + f"Cannot cancel inactive {self.transport_name} pool lease " + f"(slot={slot}, generation={generation})" + ) + + with torch.cuda.device(self.device_id): + stream_wait_value32( + self.device_id, + self.base_address + ready_byte_offset, + generation, + self.transport_name, + ) + for rank in range(self.consumer_count): + stream_write_value32( + self.device_id, + self.base_address + ack_byte_offset + rank * CONTROL_WORD_BYTES, + generation, + self.transport_name, + ) + def shutdown(self) -> None: self._recycler_stop_event.set() if self._recycle_thread.is_alive(): diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index 3216a06a1..b81805de4 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -904,6 +904,13 @@ class CudaVmmPackedTensorTransportProxy(CudaVmmTensorTransportProxy): "Packed CUDA VMM features must be reconstructed before release" ) + def release_without_reconstruction(self, consumer_count: int | None = None) -> None: + """Release the shared packed allocation when its request is abandoned.""" + if self._consumer_acknowledged: + return + self._packed_owner.acknowledge_consumption(consumer_count) + self._consumer_acknowledged = True + def reconstruct_on_target_device( self, rebuild_device_idx, consumer_count: int | None = None ): diff --git a/test/registered/unit/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index c30126232..5f52cdd65 100644 --- a/test/registered/unit/managers/test_mm_process_config.py +++ b/test/registered/unit/managers/test_mm_process_config.py @@ -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): diff --git a/test/registered/unit/managers/test_mm_shm_error_consensus.py b/test/registered/unit/managers/test_mm_shm_error_consensus.py new file mode 100644 index 000000000..595799de2 --- /dev/null +++ b/test/registered/unit/managers/test_mm_shm_error_consensus.py @@ -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() diff --git a/test/registered/unit/managers/test_multimodal_abort_cleanup.py b/test/registered/unit/managers/test_multimodal_abort_cleanup.py new file mode 100644 index 000000000..39888e8f0 --- /dev/null +++ b/test/registered/unit/managers/test_multimodal_abort_cleanup.py @@ -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"])) diff --git a/test/registered/unit/multimodal/test_cuda_ipc_transport.py b/test/registered/unit/multimodal/test_cuda_ipc_transport.py index 9d64ac8ab..ffd516ff2 100644 --- a/test/registered/unit/multimodal/test_cuda_ipc_transport.py +++ b/test/registered/unit/multimodal/test_cuda_ipc_transport.py @@ -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) diff --git a/test/registered/unit/multimodal/test_gpu_feature_transport.py b/test/registered/unit/multimodal/test_gpu_feature_transport.py index 14424291a..6504ba075 100644 --- a/test/registered/unit/multimodal/test_gpu_feature_transport.py +++ b/test/registered/unit/multimodal/test_gpu_feature_transport.py @@ -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):