diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 6c790543a..77f26c959 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1509,7 +1509,7 @@ def get_new_expanded_mm_items(original_mm_items): total_feature_len = feature_len for i in range(num_items): start, end = slice_indices[i], slice_indices[i + 1] - new_item = copy.deepcopy(item) + new_item = copy.copy(item) if item.feature is not None: new_item.feature = _slice_value(item.feature, start, end) if item.precomputed_embeddings is not None: @@ -1600,7 +1600,7 @@ def get_new_expanded_mm_items(original_mm_items): frame_start_indices[video_idx + 1], ) - new_item = copy.deepcopy(item) + new_item = copy.copy(item) if item.feature is not None: new_item.feature = _slice_value(item.feature, start, end) if item.precomputed_embeddings is not None: diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 6a019496a..9cd6e4b89 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -336,6 +336,28 @@ class MultimodalDataItem: self.hash = hash((self.hash, other.hash)) self.set_pad_value() + def reconstruct(self): + if not isinstance(self.feature, CudaIpcTensorTransportProxy): + return + + reconstruct_device = torch.cuda.current_device() + if isinstance(self.feature, CudaIpcTensorTransportProxy): + self.feature = self.feature.reconstruct_on_target_device(reconstruct_device) + if isinstance(self.precomputed_embeddings, CudaIpcTensorTransportProxy): + self.precomputed_embeddings = ( + self.precomputed_embeddings.reconstruct_on_target_device( + reconstruct_device + ) + ) + for extra_key in self.model_specific_data: + if isinstance( + self.model_specific_data[extra_key], CudaIpcTensorTransportProxy + ): + extra_data = self.model_specific_data[ + extra_key + ].reconstruct_on_target_device(reconstruct_device) + self.model_specific_data[extra_key] = extra_data + @dataclasses.dataclass class MultimodalInputs: @@ -373,13 +395,16 @@ class MultimodalInputs: @staticmethod def from_dict(obj: dict): + original_mm_items = obj["mm_items"] + for mm_item in original_mm_items: + mm_item.reconstruct() + # Check if MM splitting is enabled if not envs.SGLANG_ENABLE_MM_SPLITTING.get(): - mm_items = obj["mm_items"] + mm_items = original_mm_items else: from sglang.srt.managers.mm_utils import get_new_expanded_mm_items - original_mm_items = obj["mm_items"] # Now, `mm_items` contains one item per image. mm_items = get_new_expanded_mm_items(original_mm_items) @@ -1673,13 +1698,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): pixel_values = getattr(mm_item, "feature", None) if isinstance(pixel_values, torch.Tensor): mm_item.feature = pixel_values.to(self.device, non_blocking=True) - elif isinstance(pixel_values, CudaIpcTensorTransportProxy): - mm_item.feature = pixel_values.reconstruct_on_target_device( - torch.cuda.current_device() - ) - # The reference by CudaIpcTensorTransportProxy was cut off, - # proactively delete to avoid slow gc. - del pixel_values if get_global_server_args().language_only: precomputed_embeddings = getattr( mm_item, "precomputed_embeddings", None diff --git a/python/sglang/test/test_mm_utils.py b/python/sglang/test/test_mm_utils.py new file mode 100644 index 000000000..1d8585417 --- /dev/null +++ b/python/sglang/test/test_mm_utils.py @@ -0,0 +1,50 @@ +import unittest +from unittest.mock import Mock, patch + +import torch + +from sglang.srt.managers import mm_utils, schedule_batch +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputs, +) + + +def _make_proxy_with_reconstruct_result(tensor: torch.Tensor): + proxy = mm_utils.CudaIpcTensorTransportProxy.__new__( + mm_utils.CudaIpcTensorTransportProxy + ) + proxy.reconstruct_on_target_device = Mock(return_value=tensor) + return proxy + + +class TestMultimodalInputsFromDict(unittest.TestCase): + def test_materialize_proxy(self): + feature_tensor = torch.tensor([[7.0], [8.0]], dtype=torch.float32) + proxy_feature = _make_proxy_with_reconstruct_result(feature_tensor) + mm_item = MultimodalDataItem( + modality=Modality.IMAGE, + offsets=[(0, 1), (1, 2)], + feature=proxy_feature, + model_specific_data={"image_grid_thw": [[1, 1, 1], [1, 1, 1]]}, + ) + + with patch.object( + schedule_batch.torch.cuda, "is_available", return_value=True + ), patch.object( + schedule_batch.torch.cuda, "current_device", return_value=0 + ), patch.object( + schedule_batch.envs.SGLANG_ENABLE_MM_SPLITTING, "get", return_value=False + ), patch.object( + schedule_batch.envs.SGLANG_MM_BUFFER_SIZE_MB, "get", return_value=0 + ): + mm_inputs = MultimodalInputs.from_dict({"mm_items": [mm_item]}) + + self.assertEqual(len(mm_inputs.mm_items), 1) + self.assertTrue(torch.equal(mm_inputs.mm_items[0].feature, feature_tensor)) + proxy_feature.reconstruct_on_target_device.assert_called_once_with(0) + + +if __name__ == "__main__": + unittest.main(verbosity=2)