[VLM] Split Pixtral multi-image features before the CUDA IPC wrap (#35463)

This commit is contained in:
Mohammad Miadh Angkad
2026-08-22 15:17:55 +08:00
committed by GitHub
parent 6fd0384d42
commit 90354326c7
3 changed files with 247 additions and 65 deletions
@@ -1726,29 +1726,60 @@ class BaseMultimodalProcessor(ABC):
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,
)
for item in all_collected_items:
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 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)
self._precompute_hashes_before_cpu_transfer(mm_items)
return self._prepare_mm_items_for_transport(mm_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
)
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
return all_collected_items, input_ids, ret
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,
@@ -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 @@ class PixtralProcessor(BaseMultimodalProcessor):
"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 @@ class PixtralProcessor(BaseMultimodalProcessor):
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