diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index a7343c9659d8..90ffc00ac270 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -1726,29 +1726,60 @@ def process_and_combine_mm_data( from sglang.srt.managers.mm_utils import get_new_expanded_mm_items all_collected_items = get_new_expanded_mm_items(all_collected_items) + all_collected_items = self._finalize_mm_items( + all_collected_items, + images=base_output.images, + ) + + return all_collected_items, input_ids, ret + + def _finalize_mm_items( + self, + mm_items: List[MultimodalDataItem], + *, + images: Optional[List[Any]], + ) -> List[MultimodalDataItem]: + mm_items = self._postprocess_mm_items_before_transport( + mm_items, + images=images, + ) - for item in all_collected_items: + for item in mm_items: if item.format in ( MultimodalInputFormat.PROCESSOR_OUTPUT, MultimodalInputFormat.PRECOMPUTED_EMBEDDING, ): item.set_pad_value() - self._precompute_hashes_before_cpu_transfer(all_collected_items) - - # Wrap GPU features in the bounded IPC pool; pool misses fall back to a - # plain CPU tensor. The scheduler copies out and releases each slice. - if self.use_cuda_ipc: - # post-process, prepare for cuda-ipc transfer - for item in all_collected_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 - ) + self._precompute_hashes_before_cpu_transfer(mm_items) + return self._prepare_mm_items_for_transport(mm_items) - return all_collected_items, input_ids, ret + def _postprocess_mm_items_before_transport( + self, + mm_items: List[MultimodalDataItem], + *, + images: Optional[List[Any]], + ) -> List[MultimodalDataItem]: + """Apply model-specific item reshaping while features are still tensors.""" + return mm_items + + def _prepare_mm_items_for_transport( + self, mm_items: List[MultimodalDataItem] + ) -> List[MultimodalDataItem]: + """Wrap final GPU features for dispatch to the scheduler.""" + 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 + ) + return mm_items async def process_and_combine_mm_data_async( self, diff --git a/python/sglang/srt/multimodal/processors/pixtral.py b/python/sglang/srt/multimodal/processors/pixtral.py index a9ee65e81158..968510263641 100644 --- a/python/sglang/srt/multimodal/processors/pixtral.py +++ b/python/sglang/srt/multimodal/processors/pixtral.py @@ -1,12 +1,17 @@ +import copy import math -from typing import List, Union +from typing import Any, List, Optional, Union from transformers import PreTrainedTokenizerBase from transformers.models.pixtral.image_processing_pixtral import ( _num_image_tokens as _get_pixtral_hf_num_image_tokens, ) -from sglang.srt.managers.schedule_batch import Modality, MultimodalProcessorOutput +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalProcessorOutput, +) from sglang.srt.models.pixtral import ( PixtralForConditionalGeneration, PixtralVisionModel, @@ -41,6 +46,7 @@ def __init__(self, hf_config, server_args, _processor, *args, **kwargs): "spatial_merge_size", getattr(hf_config, "spatial_merge_size", 1), ) + self._effective_patch_size = self.patch_size * self._spatial_merge_size self._processor.patch_size = self.patch_size if self._spatial_merge_size > 1: @@ -77,58 +83,86 @@ async def process_mm_data_async( image_data=image_data, return_text=True, ) - if mm_data.images: - effective_patch = self.patch_size * self._spatial_merge_size - image_nrows = [] - for img in mm_data.images: - w, h = img.size - ratio = max(w / self.image_size, h / self.image_size) - if ratio > 1: - w = int(math.floor(w / ratio)) - h = int(math.floor(h / ratio)) - nrows, _ = _get_pixtral_hf_num_image_tokens( - (h, w), (effective_patch, effective_patch) - ) - image_nrows.append(nrows) - - mm_items, input_ids, _ = self.process_and_combine_mm_data( - mm_data, self.mm_tokens - ) - - # For multi-image: split single IMAGE mm_item into per-image items - if len(mm_data.images) > 1: - from sglang.srt.managers.schedule_batch import MultimodalDataItem - - old_item = next( - item for item in mm_items if item.modality == Modality.IMAGE - ) - all_offsets = old_item.offsets - old_feature = old_item.feature - old_image_sizes = getattr(old_item, "image_sizes", None) - - mm_items = [ - item for item in mm_items if item.modality != Modality.IMAGE - ] - offset_idx = 0 - for i, img in enumerate(mm_data.images): - nr = image_nrows[i] - item_offsets = all_offsets[offset_idx : offset_idx + nr] - offset_idx += nr - new_item = MultimodalDataItem(modality=Modality.IMAGE) - new_item.feature = old_feature[i : i + 1] - new_item.offsets = item_offsets - if old_image_sizes is not None: - new_item.model_specific_data["image_sizes"] = old_image_sizes[ - i : i + 1 - ] - mm_items.append(new_item) - else: - mm_items, input_ids, _ = self.process_and_combine_mm_data( - mm_data, self.mm_tokens - ) + mm_items, input_ids, _ = self.process_and_combine_mm_data( + mm_data, self.mm_tokens + ) return MultimodalProcessorOutput( mm_items=mm_items, input_ids=input_ids.tolist(), im_token_id=self.IM_TOKEN_ID, ) + + def _postprocess_mm_items_before_transport( + self, + mm_items: List[MultimodalDataItem], + *, + images: Optional[List[Any]], + ) -> List[MultimodalDataItem]: + if not images or len(images) <= 1: + return mm_items + + image_items = [item for item in mm_items if item.modality == Modality.IMAGE] + if len(image_items) == len(images): + return mm_items + if len(image_items) != 1: + raise ValueError( + "Pixtral multi-image processing expected one bundled IMAGE item or " + f"{len(images)} split items, but found {len(image_items)}" + ) + + old_item = image_items[0] + all_offsets = old_item.offsets + old_feature = old_item.feature + old_image_sizes = old_item.model_specific_data.get("image_sizes") + image_nrows = self._get_image_nrows(images) + if old_feature is None or len(old_feature) != len(image_nrows): + raise ValueError( + "Pixtral multi-image feature count does not match the number of " + f"images: features={0 if old_feature is None else len(old_feature)}, " + f"images={len(image_nrows)}" + ) + if all_offsets is None or sum(image_nrows) != len(all_offsets): + raise ValueError( + "Pixtral image patch rows do not match the computed offsets: " + f"rows={sum(image_nrows)}, " + f"offsets={0 if all_offsets is None else len(all_offsets)}" + ) + + split_items = [item for item in mm_items if item.modality != Modality.IMAGE] + offset_idx = 0 + for image_idx, num_rows in enumerate(image_nrows): + item_offsets = all_offsets[offset_idx : offset_idx + num_rows] + offset_idx += num_rows + new_item = copy.copy(old_item) + new_item.feature = old_feature[image_idx : image_idx + 1] + new_item.offsets = item_offsets + new_item.model_specific_data = copy.copy(old_item.model_specific_data) + if old_image_sizes is not None: + new_item.model_specific_data["image_sizes"] = old_image_sizes[ + image_idx : image_idx + 1 + ] + new_item.hash = None + new_item.pad_value = None + split_items.append(new_item) + if offset_idx != len(all_offsets): + raise ValueError( + "Pixtral multi-image split did not consume every offset: " + f"consumed={offset_idx}, offsets={len(all_offsets)}" + ) + return split_items + + def _get_image_nrows(self, images: List[Any]) -> List[int]: + image_nrows = [] + for image in images: + width, height = image.size + ratio = max(width / self.image_size, height / self.image_size) + if ratio > 1: + width = int(math.floor(width / ratio)) + height = int(math.floor(height / ratio)) + num_rows, _ = _get_pixtral_hf_num_image_tokens( + (height, width), + (self._effective_patch_size, self._effective_patch_size), + ) + image_nrows.append(num_rows) + return image_nrows diff --git a/test/registered/unit/multimodal/test_pixtral_processor.py b/test/registered/unit/multimodal/test_pixtral_processor.py new file mode 100644 index 000000000000..b45607a08b82 --- /dev/null +++ b/test/registered/unit/multimodal/test_pixtral_processor.py @@ -0,0 +1,117 @@ +"""Regression tests for Pixtral multimodal item processing.""" + +import unittest +from unittest.mock import MagicMock + +import torch +from PIL import Image + +from sglang.srt.managers.mm_utils import get_new_expanded_mm_items +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputFormat, +) +from sglang.srt.multimodal.processors.pixtral import PixtralProcessor +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestPixtralProcessor(CustomTestCase): + def _make_processor(self): + processor = object.__new__(PixtralProcessor) + processor.use_cuda_ipc = True + processor.image_size = 1024 + processor._effective_patch_size = 28 + processor._precompute_hashes_before_cpu_transfer = MagicMock() + return processor + + def test_multi_image_features_are_split_before_transport(self): + """CUDA IPC dispatch must receive per-image tensors, not a bundled proxy.""" + processor = self._make_processor() + proxies = [object(), object()] + processor._wrap_tensor_for_cuda_ipc = MagicMock(side_effect=proxies) + + feature = torch.arange(8).reshape(2, 4) + image_sizes = torch.tensor([[10, 20], [30, 40]]) + bundled_item = MultimodalDataItem( + modality=Modality.IMAGE, + feature=feature, + offsets=[(0, 0), (1, 1), (2, 2), (3, 3), (4, 4)], + format=MultimodalInputFormat.PROCESSOR_OUTPUT, + model_specific_data={"image_sizes": image_sizes, "extra_key": "keep"}, + ) + images = [Image.new("RGB", (28, 56)), Image.new("RGB", (28, 84))] + + items = processor._finalize_mm_items( + [bundled_item], + images=images, + ) + + self.assertEqual(len(items), 2) + self.assertEqual( + [item.offsets for item in items], + [[(0, 0), (1, 1)], [(2, 2), (3, 3), (4, 4)]], + ) + self.assertTrue( + all(item.format == MultimodalInputFormat.PROCESSOR_OUTPUT for item in items) + ) + self.assertTrue( + all(item.model_specific_data["extra_key"] == "keep" for item in items) + ) + self.assertTrue(all(item.pad_value is not None for item in items)) + self.assertTrue( + torch.equal(items[0].model_specific_data["image_sizes"], image_sizes[:1]) + ) + self.assertTrue( + torch.equal(items[1].model_specific_data["image_sizes"], image_sizes[1:]) + ) + wrapped_features = [ + call.args[0] for call in processor._wrap_tensor_for_cuda_ipc.call_args_list + ] + self.assertTrue(torch.equal(wrapped_features[0], feature[:1])) + self.assertTrue(torch.equal(wrapped_features[1], feature[1:])) + self.assertEqual([item.feature for item in items], proxies) + + def test_already_split_one_row_images_are_preserved(self): + """Generic per-image splits must not be collapsed and re-sliced by Pixtral.""" + processor = self._make_processor() + processor._wrap_tensor_for_cuda_ipc = MagicMock( + side_effect=[object(), object()] + ) + bundled_item = MultimodalDataItem( + modality=Modality.IMAGE, + feature=torch.arange(8).reshape(2, 4), + offsets=[(0, 0), (2, 2)], + ) + items = get_new_expanded_mm_items([bundled_item]) + images = [Image.new("RGB", (28, 28)), Image.new("RGB", (56, 28))] + + items = processor._finalize_mm_items(items, images=images) + + self.assertEqual(len(items), 2) + self.assertEqual([item.offsets for item in items], [[(0, 0)], [(2, 2)]]) + wrapped_features = [ + call.args[0] for call in processor._wrap_tensor_for_cuda_ipc.call_args_list + ] + self.assertTrue(torch.equal(wrapped_features[0], bundled_item.feature[:1])) + self.assertTrue(torch.equal(wrapped_features[1], bundled_item.feature[1:])) + + def test_mismatched_patch_rows_fail_loudly(self): + """Derived row counts cannot silently leave image placeholders unassigned.""" + processor = self._make_processor() + item = MultimodalDataItem( + modality=Modality.IMAGE, + feature=torch.arange(8).reshape(2, 4), + offsets=[(0, 0), (1, 1)], + ) + images = [Image.new("RGB", (28, 56)), Image.new("RGB", (28, 84))] + + with self.assertRaisesRegex(ValueError, "patch rows"): + processor._finalize_mm_items([item], images=images) + + +if __name__ == "__main__": + unittest.main()