diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 1810d9cd0..a3d18e063 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -899,18 +899,18 @@ def embed_mm_inputs( other_info["input_deepstack_embeds"] = input_deepstack_embeds # 4. scatter embeddings into input embedding + # masked_scatter_ avoids the cudaStreamSynchronize that torch.where triggers. + def _scatter(dest, mask, src): + dest.masked_scatter_(mask.expand_as(dest), src.to(dest.device, dest.dtype)) + for i, modality, embedding, mask in zip( range(len(embeddings)), modalities, embeddings, masks ): if embedding is None or mask is None: continue - # in-place update - indices = torch.where(mask.squeeze(dim=-1))[0] - input_embeds[indices] = embedding.to(input_embeds.device, input_embeds.dtype) + _scatter(input_embeds, mask, embedding) if use_deepstack.get(modality, None): - input_deepstack_embeds[indices] = deepstack_embeddings[i].to( - input_embeds.device, input_embeds.dtype - ) + _scatter(input_deepstack_embeds, mask, deepstack_embeddings[i]) return input_embeds, other_info @@ -1243,7 +1243,9 @@ def tensor_hash(tensor_list) -> int: def hash_feature(f): if isinstance(f, list): - if isinstance(f[0], torch.Tensor): + if len(f) > 0 and isinstance(f[0], ShmPointerMMData): + return tensor_hash([x.tensor for x in f]) + if len(f) > 0 and isinstance(f[0], torch.Tensor): return tensor_hash(f) return data_hash(tuple(flatten_nested_list(f))) elif isinstance(f, np.ndarray): @@ -1257,6 +1259,10 @@ def hash_feature(f): elif isinstance(f, CudaIpcTensorTransportProxy): reconstruct_t = f.reconstruct_on_target_device(torch.cuda.current_device()) return tensor_hash([reconstruct_t]) + elif isinstance(f, ShmPointerMMData): + if f.precomputed_hash is not None: + return f.precomputed_hash + return tensor_hash([f.tensor]) return data_hash(f) @@ -1555,13 +1561,14 @@ class ShmPointerMMData: This acts as a "pointer" to the tensor data across process boundaries. """ - def __init__(self, tensor: torch.Tensor): + def __init__(self, tensor: torch.Tensor, precomputed_hash: Optional[int] = None): if not tensor.is_cpu: tensor = tensor.cpu() if not tensor.is_contiguous(): tensor = tensor.contiguous() self.shape = tensor.shape self.dtype = tensor.dtype + self.precomputed_hash = precomputed_hash nbytes = tensor.numel() * tensor.element_size() shm = shared_memory.SharedMemory(create=True, size=nbytes) try: @@ -1580,12 +1587,14 @@ class ShmPointerMMData: "shm_name": self.shm_name, "shape": self.shape, "dtype": self.dtype, + "precomputed_hash": self.precomputed_hash, } def __setstate__(self, state): self.shm_name = state["shm_name"] 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( @@ -1624,10 +1633,15 @@ def _get_is_default_transport(): return _is_default_tensor_transport -def _wrap_tensor_or_list(value): - """Wrap a CPU tensor (or list of CPU tensors) in ShmPointerMMData.""" +def _wrap_tensor_or_list(value, precomputed_hash: Optional[int] = None): + """Wrap a CPU tensor (or list of CPU tensors) in ShmPointerMMData. + + ``precomputed_hash`` is only forwarded for the single-tensor case. + For list features the item-level hash covers all elements jointly, + so per-element hashes are not applicable. + """ if isinstance(value, torch.Tensor) and value.is_cpu: - return ShmPointerMMData(value) + return ShmPointerMMData(value, precomputed_hash=precomputed_hash) elif isinstance(value, (list, tuple)): wrapped = [ (ShmPointerMMData(t) if isinstance(t, torch.Tensor) and t.is_cpu else t) @@ -1646,14 +1660,17 @@ def wrap_shm_features(obj): if hasattr(obj, "mm_inputs") and obj.mm_inputs: for item in obj.mm_inputs.mm_items: + item_hash = getattr(item, "hash", None) if hasattr(item, "feature") and item.feature is not None: - item.feature = _wrap_tensor_or_list(item.feature) + item.feature = _wrap_tensor_or_list( + item.feature, precomputed_hash=item_hash + ) if ( hasattr(item, "precomputed_embeddings") and item.precomputed_embeddings is not None ): item.precomputed_embeddings = _wrap_tensor_or_list( - item.precomputed_embeddings + item.precomputed_embeddings, precomputed_hash=item_hash ) return obj diff --git a/test/registered/vlm/test_vision_openai_server_a.py b/test/registered/vlm/test_vision_openai_server_a.py index 68b1c79f2..95814220e 100644 --- a/test/registered/vlm/test_vision_openai_server_a.py +++ b/test/registered/vlm/test_vision_openai_server_a.py @@ -146,6 +146,7 @@ class TestKimiVLServer(ImageOpenAITestMixin): extra_args = [ "--context-length=8192", "--dtype=bfloat16", + "--mem-fraction-static=0.40", ] def test_video_images_chat_completion(self):