[Fix] SGLANG_USE_CUDA_IPC_TRANSPORT=1 and SGLANG_ENABLE_MM_SPLITTING=1 do not work at the same time. (#19915)
This commit is contained in:
@@ -1509,7 +1509,7 @@ def get_new_expanded_mm_items(original_mm_items):
|
|||||||
total_feature_len = feature_len
|
total_feature_len = feature_len
|
||||||
for i in range(num_items):
|
for i in range(num_items):
|
||||||
start, end = slice_indices[i], slice_indices[i + 1]
|
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:
|
if item.feature is not None:
|
||||||
new_item.feature = _slice_value(item.feature, start, end)
|
new_item.feature = _slice_value(item.feature, start, end)
|
||||||
if item.precomputed_embeddings is not None:
|
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],
|
frame_start_indices[video_idx + 1],
|
||||||
)
|
)
|
||||||
|
|
||||||
new_item = copy.deepcopy(item)
|
new_item = copy.copy(item)
|
||||||
if item.feature is not None:
|
if item.feature is not None:
|
||||||
new_item.feature = _slice_value(item.feature, start, end)
|
new_item.feature = _slice_value(item.feature, start, end)
|
||||||
if item.precomputed_embeddings is not None:
|
if item.precomputed_embeddings is not None:
|
||||||
|
|||||||
@@ -336,6 +336,28 @@ class MultimodalDataItem:
|
|||||||
self.hash = hash((self.hash, other.hash))
|
self.hash = hash((self.hash, other.hash))
|
||||||
self.set_pad_value()
|
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
|
@dataclasses.dataclass
|
||||||
class MultimodalInputs:
|
class MultimodalInputs:
|
||||||
@@ -373,13 +395,16 @@ class MultimodalInputs:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dict(obj: dict):
|
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
|
# Check if MM splitting is enabled
|
||||||
if not envs.SGLANG_ENABLE_MM_SPLITTING.get():
|
if not envs.SGLANG_ENABLE_MM_SPLITTING.get():
|
||||||
mm_items = obj["mm_items"]
|
mm_items = original_mm_items
|
||||||
else:
|
else:
|
||||||
from sglang.srt.managers.mm_utils import get_new_expanded_mm_items
|
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.
|
# Now, `mm_items` contains one item per image.
|
||||||
mm_items = get_new_expanded_mm_items(original_mm_items)
|
mm_items = get_new_expanded_mm_items(original_mm_items)
|
||||||
|
|
||||||
@@ -1673,13 +1698,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
pixel_values = getattr(mm_item, "feature", None)
|
pixel_values = getattr(mm_item, "feature", None)
|
||||||
if isinstance(pixel_values, torch.Tensor):
|
if isinstance(pixel_values, torch.Tensor):
|
||||||
mm_item.feature = pixel_values.to(self.device, non_blocking=True)
|
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:
|
if get_global_server_args().language_only:
|
||||||
precomputed_embeddings = getattr(
|
precomputed_embeddings = getattr(
|
||||||
mm_item, "precomputed_embeddings", None
|
mm_item, "precomputed_embeddings", None
|
||||||
|
|||||||
@@ -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)
|
||||||
Reference in New Issue
Block a user