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

This commit is contained in:
Mick
2026-09-01 13:46:38 +08:00
committed by GitHub
parent 33ed29a0ee
commit ae2bd5728b
16 changed files with 1408 additions and 96 deletions
+6
View File
@@ -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]
+98 -22
View File
@@ -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)
+52 -13
View File
@@ -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:
+166 -25
View File
@@ -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.
@@ -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 = [
@@ -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(
@@ -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,
)
@@ -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(),
@@ -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,
@@ -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():
@@ -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
):
@@ -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):