From 5753b82b684a46cd4c958f310f38e98dca2d655e Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Wed, 19 Aug 2026 11:28:44 +0800 Subject: [PATCH 1/2] fix(vlm): split Pixtral features before transport --- .../multimodal/processors/base_processor.py | 61 +++++++--- .../srt/multimodal/processors/pixtral.py | 110 ++++++++++-------- .../unit/multimodal/test_pixtral_processor.py | 59 ++++++++++ 3 files changed, 165 insertions(+), 65 deletions(-) create mode 100644 test/registered/unit/multimodal/test_pixtral_processor.py diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 5b097f40d148..7ed906dbcd64 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -1725,29 +1725,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, + base_output=base_output, + ) + + return all_collected_items, input_ids, ret - for item in all_collected_items: + def _finalize_mm_items( + self, + mm_items: List[MultimodalDataItem], + *, + base_output: BaseMultiModalProcessorOutput, + ) -> List[MultimodalDataItem]: + mm_items = self._postprocess_mm_items_before_transport( + mm_items, + base_output=base_output, + ) + + 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) + self._prepare_mm_items_for_transport(mm_items) + return mm_items - return all_collected_items, input_ids, ret + def _postprocess_mm_items_before_transport( + self, + mm_items: List[MultimodalDataItem], + *, + base_output: BaseMultiModalProcessorOutput, + ) -> 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] + ) -> None: + """Wrap final GPU features for dispatch to the scheduler.""" + if not self.use_cuda_ipc: + return + + # 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 + ) 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..316f0a98defb 100644 --- a/python/sglang/srt/multimodal/processors/pixtral.py +++ b/python/sglang/srt/multimodal/processors/pixtral.py @@ -6,13 +6,18 @@ _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, ) from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, + BaseMultiModalProcessorOutput, MultimodalSpecialTokens, ) @@ -77,58 +82,63 @@ 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], + *, + base_output: BaseMultiModalProcessorOutput, + ) -> List[MultimodalDataItem]: + if len(base_output.images) <= 1: + return mm_items + + 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 = old_item.model_specific_data.get("image_sizes") + image_nrows = self._get_image_nrows(base_output.images) + + 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 + model_specific_data = {} + if old_image_sizes is not None: + model_specific_data["image_sizes"] = old_image_sizes[ + image_idx : image_idx + 1 + ] + split_items.append( + MultimodalDataItem( + modality=Modality.IMAGE, + feature=old_feature[image_idx : image_idx + 1], + offsets=item_offsets, + model_specific_data=model_specific_data, + ) + ) + return split_items + + def _get_image_nrows(self, images) -> List[int]: + effective_patch = self.patch_size * self._spatial_merge_size + 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), + (effective_patch, effective_patch), + ) + 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..3c7acf8ee04f --- /dev/null +++ b/test/registered/unit/multimodal/test_pixtral_processor.py @@ -0,0 +1,59 @@ +"""Regression tests for Pixtral multimodal item processing.""" + +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock + +import torch + +from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +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 test_multi_image_features_are_split_before_transport(self): + """CUDA IPC dispatch must receive per-image tensors, not a bundled proxy.""" + processor = object.__new__(PixtralProcessor) + processor.use_cuda_ipc = True + processor._get_image_nrows = MagicMock(return_value=[2, 3]) + proxies = [object(), object()] + processor._wrap_tensor_for_cuda_ipc = MagicMock(side_effect=proxies) + processor._precompute_hashes_before_cpu_transfer = MagicMock() + + 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, 1, 2, 3, 4], + model_specific_data={"image_sizes": image_sizes}, + ) + base_output = SimpleNamespace(images=[object(), object()]) + + items = processor._finalize_mm_items( + [bundled_item], + base_output=base_output, + ) + + self.assertEqual(len(items), 2) + self.assertEqual([item.offsets for item in items], [[0, 1], [2, 3, 4]]) + 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) + + +if __name__ == "__main__": + unittest.main() From 7629f3c58d459bbbd2671ec869740df844b8b429 Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Wed, 19 Aug 2026 13:27:45 +0800 Subject: [PATCH 2/2] fix(vlm): harden Pixtral multi-image splitting --- .../multimodal/processors/base_processor.py | 16 ++-- .../srt/multimodal/processors/pixtral.py | 60 +++++++++----- .../unit/multimodal/test_pixtral_processor.py | 80 ++++++++++++++++--- 3 files changed, 119 insertions(+), 37 deletions(-) diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 7ed906dbcd64..9b89be354321 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -1727,7 +1727,7 @@ def process_and_combine_mm_data( all_collected_items = get_new_expanded_mm_items(all_collected_items) all_collected_items = self._finalize_mm_items( all_collected_items, - base_output=base_output, + images=base_output.images, ) return all_collected_items, input_ids, ret @@ -1736,11 +1736,11 @@ def _finalize_mm_items( self, mm_items: List[MultimodalDataItem], *, - base_output: BaseMultiModalProcessorOutput, + images: Optional[List[Any]], ) -> List[MultimodalDataItem]: mm_items = self._postprocess_mm_items_before_transport( mm_items, - base_output=base_output, + images=images, ) for item in mm_items: @@ -1751,24 +1751,23 @@ def _finalize_mm_items( item.set_pad_value() self._precompute_hashes_before_cpu_transfer(mm_items) - self._prepare_mm_items_for_transport(mm_items) - return mm_items + return self._prepare_mm_items_for_transport(mm_items) def _postprocess_mm_items_before_transport( self, mm_items: List[MultimodalDataItem], *, - base_output: BaseMultiModalProcessorOutput, + 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] - ) -> None: + ) -> List[MultimodalDataItem]: """Wrap final GPU features for dispatch to the scheduler.""" if not self.use_cuda_ipc: - return + return mm_items # Pool misses fall back to plain CPU tensors. The scheduler copies out # and releases each successful pool slice. @@ -1779,6 +1778,7 @@ def _prepare_mm_items_for_transport( 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 316f0a98defb..968510263641 100644 --- a/python/sglang/srt/multimodal/processors/pixtral.py +++ b/python/sglang/srt/multimodal/processors/pixtral.py @@ -1,5 +1,6 @@ +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 ( @@ -17,7 +18,6 @@ ) from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor, - BaseMultiModalProcessorOutput, MultimodalSpecialTokens, ) @@ -46,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: @@ -96,39 +97,62 @@ def _postprocess_mm_items_before_transport( self, mm_items: List[MultimodalDataItem], *, - base_output: BaseMultiModalProcessorOutput, + images: Optional[List[Any]], ) -> List[MultimodalDataItem]: - if len(base_output.images) <= 1: + if not images or len(images) <= 1: return mm_items - old_item = next(item for item in mm_items if item.modality == Modality.IMAGE) + 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(base_output.images) + 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 - model_specific_data = {} + 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: - model_specific_data["image_sizes"] = old_image_sizes[ + new_item.model_specific_data["image_sizes"] = old_image_sizes[ image_idx : image_idx + 1 ] - split_items.append( - MultimodalDataItem( - modality=Modality.IMAGE, - feature=old_feature[image_idx : image_idx + 1], - offsets=item_offsets, - model_specific_data=model_specific_data, - ) + 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[int]: - effective_patch = self.patch_size * self._spatial_merge_size + def _get_image_nrows(self, images: List[Any]) -> List[int]: image_nrows = [] for image in images: width, height = image.size @@ -138,7 +162,7 @@ def _get_image_nrows(self, images) -> List[int]: height = int(math.floor(height / ratio)) num_rows, _ = _get_pixtral_hf_num_image_tokens( (height, width), - (effective_patch, effective_patch), + (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 index 3c7acf8ee04f..b45607a08b82 100644 --- a/test/registered/unit/multimodal/test_pixtral_processor.py +++ b/test/registered/unit/multimodal/test_pixtral_processor.py @@ -1,12 +1,17 @@ """Regression tests for Pixtral multimodal item processing.""" import unittest -from types import SimpleNamespace from unittest.mock import MagicMock import torch +from PIL import Image -from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem +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 @@ -15,32 +20,48 @@ class TestPixtralProcessor(CustomTestCase): - def test_multi_image_features_are_split_before_transport(self): - """CUDA IPC dispatch must receive per-image tensors, not a bundled proxy.""" + def _make_processor(self): processor = object.__new__(PixtralProcessor) processor.use_cuda_ipc = True - processor._get_image_nrows = MagicMock(return_value=[2, 3]) + 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) - processor._precompute_hashes_before_cpu_transfer = MagicMock() 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, 1, 2, 3, 4], - model_specific_data={"image_sizes": image_sizes}, + 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"}, ) - base_output = SimpleNamespace(images=[object(), object()]) + images = [Image.new("RGB", (28, 56)), Image.new("RGB", (28, 84))] items = processor._finalize_mm_items( [bundled_item], - base_output=base_output, + images=images, ) self.assertEqual(len(items), 2) - self.assertEqual([item.offsets for item in items], [[0, 1], [2, 3, 4]]) + 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]) ) @@ -54,6 +75,43 @@ def test_multi_image_features_are_split_before_transport(self): 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()