[VLM] Enable per-image MM splitting by default and remove MULTI_IMAGES modality (#21899)

This commit is contained in:
Yuhao Yang
2026-04-03 11:04:41 +08:00
committed by GitHub
parent 8897ac58f0
commit 69e89a1fcc
12 changed files with 215 additions and 134 deletions
-3
View File
@@ -467,9 +467,6 @@ class Envs:
SGLANG_MM_FEATURE_CACHE_MB = EnvInt(4 * 1024)
SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05)
# MM splitting behavior control
SGLANG_ENABLE_MM_SPLITTING = EnvBool(False)
# Mamba
SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16")
SGLANG_MAMBA_SSM_DTYPE = EnvStr(None)
+71 -39
View File
@@ -327,46 +327,26 @@ class MultiModalityDataPaddingPatternMultimodalTokens(MultiModalityDataPaddingPa
input_ids_tensor = torch.as_tensor(input_ids)
# Check if MM splitting is enabled
if envs.SGLANG_ENABLE_MM_SPLITTING.get():
items_by_modality = defaultdict(list)
for item in mm_inputs.mm_items:
items_by_modality[item.modality].append(item)
# Replace multimodal tokens using per-item offsets
items_by_modality = defaultdict(list)
for item in mm_inputs.mm_items:
items_by_modality[item.modality].append(item)
token_id_map = {
Modality.IMAGE: mm_inputs.im_token_id,
Modality.MULTI_IMAGES: mm_inputs.im_token_id,
Modality.AUDIO: mm_inputs.audio_token_id,
Modality.VIDEO: mm_inputs.video_token_id,
}
token_id_map = {
Modality.IMAGE: mm_inputs.im_token_id,
Modality.AUDIO: mm_inputs.audio_token_id,
Modality.VIDEO: mm_inputs.video_token_id,
}
for modality, items in items_by_modality.items():
token_id = token_id_map.get(modality)
for modality, items in items_by_modality.items():
token_id = token_id_map.get(modality)
if not items or token_id is None:
continue
if not items or token_id is None:
continue
for i, item in enumerate(items):
for offset in items[i].offsets:
input_ids_tensor[offset[0] : offset[1] + 1] = item.pad_value
else:
# Create mapping of token_ids to pad_values for each modality
token_to_pad_mapping = {}
for item in mm_inputs.mm_items:
if item.is_image() and mm_inputs.im_token_id is not None:
token_to_pad_mapping[mm_inputs.im_token_id] = item.pad_value
elif item.is_audio() and mm_inputs.audio_token_id is not None:
token_to_pad_mapping[mm_inputs.audio_token_id] = item.pad_value
elif item.is_video() and mm_inputs.video_token_id is not None:
token_to_pad_mapping[mm_inputs.video_token_id] = item.pad_value
else:
raise ValueError(
f"No multimodal token id provided for {item.modality}"
)
# Apply replacements for all tokens at once
for token_id, pad_value in token_to_pad_mapping.items():
input_ids_tensor[input_ids_tensor == token_id] = pad_value
for i, item in enumerate(items):
for offset in items[i].offsets:
input_ids_tensor[offset[0] : offset[1] + 1] = item.pad_value
ret_input_ids = input_ids_tensor.tolist()
return ret_input_ids
@@ -1476,6 +1456,54 @@ def _slice_model_data(
return sliced
def _try_simple_split(item, num_items, expanded_mm_items):
"""Try to split a bundled item by matching feature dim-0 to offset count.
Returns True if split succeeded, False otherwise."""
feature = item.feature if item.feature is not None else item.precomputed_embeddings
if feature is None:
return False
if isinstance(feature, (torch.Tensor, np.ndarray)):
feature_count = feature.shape[0]
elif isinstance(feature, (list, tuple)):
feature_count = len(feature)
else:
return False
if feature_count != num_items:
return False
for i in range(num_items):
new_item = copy.copy(item)
if item.feature is not None:
if isinstance(item.feature, (list, tuple)):
new_item.feature = [item.feature[i]]
else:
new_item.feature = item.feature[i : i + 1]
if item.precomputed_embeddings is not None:
if isinstance(item.precomputed_embeddings, (list, tuple)):
new_item.precomputed_embeddings = [item.precomputed_embeddings[i]]
else:
new_item.precomputed_embeddings = item.precomputed_embeddings[i : i + 1]
new_item.offsets = [item.offsets[i]]
new_data = {}
for k, v in item.model_specific_data.items():
if isinstance(v, (list, tuple)) and len(v) == num_items:
new_data[k] = [v[i]]
elif (
isinstance(v, (torch.Tensor, np.ndarray))
and len(v.shape) > 0
and v.shape[0] == num_items
):
new_data[k] = v[i : i + 1]
else:
new_data[k] = v
new_item.model_specific_data = new_data
new_item.hash = None
expanded_mm_items.append(new_item)
return True
def get_new_expanded_mm_items(original_mm_items):
expanded_mm_items = []
for item in original_mm_items:
@@ -1488,7 +1516,9 @@ def get_new_expanded_mm_items(original_mm_items):
image_grid_thw = item.model_specific_data.get("image_grid_thw")
grid_len = _get_length(image_grid_thw)
if image_grid_thw is None or grid_len != num_items:
expanded_mm_items.append(item)
# No grid info — fall back to simple split by feature dim-0
if not _try_simple_split(item, num_items, expanded_mm_items):
expanded_mm_items.append(item)
continue
patches_per_item = []
@@ -1533,7 +1563,8 @@ def get_new_expanded_mm_items(original_mm_items):
elif item.is_video():
video_grid_thw = item.model_specific_data.get("video_grid_thw")
if video_grid_thw is None:
expanded_mm_items.append(item)
if not _try_simple_split(item, num_items, expanded_mm_items):
expanded_mm_items.append(item)
continue
# video_grid_thw shape: [num_videos, 3] where each row is [T, H, W]
@@ -1623,7 +1654,8 @@ def get_new_expanded_mm_items(original_mm_items):
new_item.hash = None
expanded_mm_items.append(new_item)
else:
expanded_mm_items.append(item)
if not _try_simple_split(item, num_items, expanded_mm_items):
expanded_mm_items.append(item)
else:
expanded_mm_items.append(item)
+7 -22
View File
@@ -198,7 +198,6 @@ class FINISH_ABORT(BaseFinishReason):
class Modality(Enum):
IMAGE = auto()
MULTI_IMAGES = auto()
VIDEO = auto()
AUDIO = auto()
@@ -225,9 +224,10 @@ class MultimodalInputFormat(Enum):
@dataclasses.dataclass
class MultimodalDataItem:
"""
One MultimodalDataItem contains all inputs for one modality.
For example, if there are 3 images and 1 audio inputs, there will be 2 MultimodalDataItem.
One for images and one for audio.
One MultimodalDataItem represents a single multimodal input (one image, one video, or one audio).
For example, if there are 3 images and 1 audio, there will be 4 MultimodalDataItems.
Each item has its own hash and pad_value, enabling per-image RadixAttention caching.
We put the common fields first and the model-specific fields in model_specific_data.
"""
@@ -305,7 +305,7 @@ class MultimodalDataItem:
return self.modality == Modality.AUDIO
def is_image(self):
return self.modality in [Modality.IMAGE, Modality.MULTI_IMAGES]
return self.modality == Modality.IMAGE
def is_video(self):
return self.modality == Modality.VIDEO
@@ -330,12 +330,6 @@ class MultimodalDataItem:
ret.validate()
return ret
def merge(self, other):
self.feature += other.feature
self.offsets += other.offsets
self.hash = hash((self.hash, other.hash))
self.set_pad_value()
def reconstruct(self):
if not isinstance(self.feature, CudaIpcTensorTransportProxy):
return
@@ -395,19 +389,10 @@ class MultimodalInputs:
@staticmethod
def from_dict(obj: dict):
original_mm_items = obj["mm_items"]
for mm_item in original_mm_items:
mm_items = obj["mm_items"]
for mm_item in mm_items:
mm_item.reconstruct()
# Check if MM splitting is enabled
if not envs.SGLANG_ENABLE_MM_SPLITTING.get():
mm_items = original_mm_items
else:
from sglang.srt.managers.mm_utils import get_new_expanded_mm_items
# Now, `mm_items` contains one item per image.
mm_items = get_new_expanded_mm_items(original_mm_items)
ret = MultimodalInputs(
mm_items=mm_items,
)
+47 -34
View File
@@ -58,6 +58,21 @@ _KNOWN_BROKEN_AUTOMODEL_ERROR = "Could not find VoxtralRealtimeTextModel"
class LlavaBaseForCausalLM(nn.Module):
@staticmethod
def _infer_image_aspect_ratio(mm_items):
"""Determine image_aspect_ratio from processor metadata or item count."""
# Check if processor stored the aspect_ratio it used
for item in mm_items:
ar = item.model_specific_data.get("image_aspect_ratio")
if ar is not None:
return ar
# Fallback: multi-image or video → pad, single image → anyres
image_items = [item for item in mm_items if item.is_image()]
has_video = any(item.is_video() for item in mm_items)
if len(image_items) > 1 or has_video:
return "pad"
return "anyres"
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
image_sizes = flatten_nested_list(
[item.image_sizes for item in image_inputs.mm_items]
@@ -66,13 +81,8 @@ class LlavaBaseForCausalLM(nn.Module):
pad_values = [item.pad_value for item in image_inputs.mm_items]
# hardcode for spatial_unpad + anyres
if any(
item.modality == Modality.MULTI_IMAGES or item.modality == Modality.VIDEO
for item in image_inputs.mm_items
):
image_aspect_ratio = "pad"
else:
image_aspect_ratio = "anyres"
# Use per-item aspect_ratio from processor if available, else infer
image_aspect_ratio = self._infer_image_aspect_ratio(image_inputs.mm_items)
offset_list = []
image_inputs.image_pad_len = []
for image_idx, image_s in enumerate(image_sizes):
@@ -168,13 +178,9 @@ class LlavaBaseForCausalLM(nn.Module):
# Embed text inputs
input_embeds = self.language_model.model.embed_tokens(input_ids)
# Got List[List[str]] extend it to List[str]
# The length of the List should be equal to batch size
modalities_list = []
# Compute max image offset per request to determine need_vision
max_image_offset = []
for im in image_inputs:
if im:
modalities_list.extend([item.modality for item in im.mm_items])
if im and im.image_offsets:
max_image_offset.append(
np.max(np.array(im.image_offsets) + np.array(im.image_pad_len))
@@ -187,6 +193,18 @@ class LlavaBaseForCausalLM(nn.Module):
if need_vision.any():
bs = forward_batch.batch_size
# Build per-image lists filtered by need_vision
modalities_list = []
aspect_ratios = [] # per-image aspect ratio
for i in range(bs):
if need_vision[i] and image_inputs[i]:
items = image_inputs[i].mm_items
ar = self._infer_image_aspect_ratio(items)
for item in items:
modalities_list.append(item.modality)
aspect_ratios.append(ar)
pixel_values = flatten_nested_list(
[
[item.feature for item in image_inputs[i].mm_items]
@@ -194,12 +212,12 @@ class LlavaBaseForCausalLM(nn.Module):
if need_vision[i]
]
)
# Per-image sizes (each entry is [(w,h)] for one image)
image_sizes = [
flatten_nested_list(
[item.image_sizes for item in image_inputs[i].mm_items]
)
item.image_sizes
for i in range(bs)
if need_vision[i]
for item in image_inputs[i].mm_items
]
########## Encode Image ########
@@ -228,18 +246,7 @@ class LlavaBaseForCausalLM(nn.Module):
new_image_features = []
height = width = self.num_patches_per_side
for image_idx, image_feature in enumerate(image_features):
if modalities_list[image_idx] == Modality.IMAGE:
image_aspect_ratio = (
self.config.image_aspect_ratio
) # single image
elif (
modalities_list[image_idx] == Modality.MULTI_IMAGES
or modalities_list[image_idx] == Modality.VIDEO
):
image_aspect_ratio = "pad" # multi image
# image_aspect_ratio = (
# "anyres" if len(image_sizes[image_idx]) == 1 else "pad"
# )
image_aspect_ratio = aspect_ratios[image_idx]
if (
image_feature.shape[0] > 1
and "anyres" in image_aspect_ratio
@@ -388,6 +395,7 @@ class LlavaBaseForCausalLM(nn.Module):
extend_start_loc_cpu = forward_batch.extend_start_loc.cpu().numpy()
extend_seq_lens = forward_batch.extend_seq_lens.cpu().numpy()
prefix_lens_cpu = forward_batch.extend_prefix_lens_cpu
# Fill in the image features using flat indexing (one pt per image)
pt = 0
for i in range(bs):
if not need_vision[i]:
@@ -396,20 +404,25 @@ class LlavaBaseForCausalLM(nn.Module):
start_idx = extend_start_loc_cpu[i]
seq_len = extend_seq_lens[i]
prefix_len = prefix_lens_cpu[i]
n_images = len(image_inputs[i].image_offsets)
for j in range(n_images):
image_offset = image_inputs[i].image_offsets[j]
# Multiple images
for image_idx, image_offset in enumerate(
image_inputs[i].image_offsets
):
if (
image_offset + image_inputs[i].image_pad_len[image_idx]
image_offset + image_inputs[i].image_pad_len[j]
<= prefix_len
):
pt += 1
continue
if image_offset >= prefix_len + seq_len:
pt += n_images - j
break
tmp_image_feature = image_features[pt][image_idx]
tmp_image_feature = image_features[pt]
# Squeeze batch dim from per-image features [1, feat, hidden]
if tmp_image_feature.ndim == 3:
tmp_image_feature = tmp_image_feature[0]
pad_len = tmp_image_feature.shape[0]
input_offset = image_offset - prefix_len
@@ -432,7 +445,7 @@ class LlavaBaseForCausalLM(nn.Module):
print(
f"{start_idx=}, {image_offset=}, {prefix_len=}, {pad_len=}"
)
pt += 1
pt += 1
return self.language_model(
input_ids, positions, forward_batch, input_embeds=input_embeds
+15 -3
View File
@@ -993,7 +993,11 @@ class MiniCPMV2_6(MiniCPMBaseModel):
slice_end_id: int = image_inputs.slice_end_id
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
# Only increment data_idx on im_start (not slice_start) so all slices
# within one image share the same pad_value for per-image caching.
pattern = MultiModalityDataPaddingPatternTokenPairs(
media_token_pairs, data_start_token_ids=[im_start_id]
)
return pattern.pad_input_tokens(input_ids, image_inputs)
@@ -1155,7 +1159,11 @@ class MiniCPMV4_0(MiniCPMBaseModel):
slice_end_id: int = image_inputs.slice_end_id
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
# Only increment data_idx on im_start (not slice_start) so all slices
# within one image share the same pad_value for per-image caching.
pattern = MultiModalityDataPaddingPatternTokenPairs(
media_token_pairs, data_start_token_ids=[im_start_id]
)
return pattern.pad_input_tokens(input_ids, image_inputs)
@@ -1321,7 +1329,11 @@ class MiniCPMV4_5(MiniCPMBaseModel):
slice_end_id: int = image_inputs.slice_end_id
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
# Only increment data_idx on im_start (not slice_start) so all slices
# within one image share the same pad_value for per-image caching.
pattern = MultiModalityDataPaddingPatternTokenPairs(
media_token_pairs, data_start_token_ids=[im_start_id]
)
return pattern.pad_input_tokens(input_ids, image_inputs)
@@ -137,7 +137,6 @@ class MultimodalSpecialTokens:
def get_token_id_by_modality(self, modality: Modality) -> Optional[int]:
return {
Modality.IMAGE: self.image_token_id,
Modality.MULTI_IMAGES: self.image_token_id,
Modality.VIDEO: self.video_token_id,
Modality.AUDIO: self.audio_token_id,
}.get(modality)
@@ -359,7 +358,7 @@ class BaseMultimodalProcessor(ABC):
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
offsets=[offset],
precomputed_embeddings=embedding_slice,
)
)
@@ -998,7 +997,8 @@ class BaseMultimodalProcessor(ABC):
self, data_dict: dict, modality: Modality = None
) -> List[MultimodalDataItem]:
"""
Create mm_items directly from processor output, with one item for each modality
Create mm_items from processor output. Initially creates one item per modality;
these are later split into per-image/video items by get_new_expanded_mm_items.
Note that the data_dict can be passed via offline engine api
"""
@@ -1141,6 +1141,11 @@ class BaseMultimodalProcessor(ABC):
mm_token_id=mm_token_id,
)
# Split bundled items into per-image/video items for better cache granularity
from sglang.srt.managers.mm_utils import get_new_expanded_mm_items
all_collected_items = get_new_expanded_mm_items(all_collected_items)
"""
solution for cuda-ipc memory-leak:
1. memory-pool: each time get a slice from memory-pool and use it as transport-data (with async lock guard)
@@ -588,11 +588,21 @@ class InternVLProcessor(BaseMultimodalProcessor):
items = []
if image_tensor is not None:
items.append(
MultimodalDataItem(
feature=image_tensor, modality=Modality.IMAGE, offsets=image_offsets
)
# Split per-image for better cache granularity
assert len(num_patches_list) == len(image_offsets), (
f"InternVL: num_patches_list ({len(num_patches_list)}) != "
f"image_offsets ({len(image_offsets)})"
)
cumulative = 0
for i, num_patches in enumerate(num_patches_list):
items.append(
MultimodalDataItem(
feature=image_tensor[cumulative : cumulative + num_patches],
modality=Modality.IMAGE,
offsets=[image_offsets[i]],
)
)
cumulative += num_patches
if video_tensor is not None:
items.append(
MultimodalDataItem(
@@ -702,11 +712,21 @@ class InternVLProcessor(BaseMultimodalProcessor):
items = []
if pixel_values is not None:
items.append(
MultimodalDataItem(
feature=pixel_values, modality=Modality.IMAGE, offsets=image_offsets
)
# Split per-image for better cache granularity
assert len(num_patches_list) == len(image_offsets), (
f"InternVL: num_patches_list ({len(num_patches_list)}) != "
f"image_offsets ({len(image_offsets)})"
)
cumulative = 0
for i, num_patches in enumerate(num_patches_list):
items.append(
MultimodalDataItem(
feature=pixel_values[cumulative : cumulative + num_patches],
modality=Modality.IMAGE,
offsets=[image_offsets[i]],
)
)
cumulative += num_patches
return {
"input_ids": input_ids,
@@ -187,34 +187,39 @@ class LlavaImageProcessor(BaseMultimodalProcessor):
pixel_values.append(pixel_v)
data_hashes.append(image_h)
image_sizes.append(image_s)
if isinstance(pixel_values[0], np.ndarray):
pixel_values = np.stack(pixel_values, axis=0)
else:
# A single image
pixel_values, image_hash, image_size = await self._process_single_image(
image_data[0], aspect_ratio, grid_pinpoints
)
pixel_values = [pixel_values]
image_sizes = [image_size]
else:
raise ValueError(f"Invalid image data: {image_data}")
modality = Modality.IMAGE
if isinstance(request_obj.modalities, list):
if request_obj.modalities[0] == "multi-images":
modality = Modality.MULTI_IMAGES
elif request_obj.modalities[0] == "video":
if request_obj.modalities[0] == "video":
modality = Modality.VIDEO
return {
"mm_items": [
# Create one item per image for better cache granularity
mm_items = []
for pixel_v, image_s in zip(pixel_values, image_sizes):
# Ensure ndim=4 so the model forward takes the correct encode branch
if isinstance(pixel_v, np.ndarray) and pixel_v.ndim == 3:
pixel_v = np.expand_dims(pixel_v, 0)
mm_items.append(
MultimodalDataItem(
feature=pixel_values,
feature=pixel_v,
model_specific_data={
"image_sizes": image_sizes,
"image_sizes": [image_s],
"image_aspect_ratio": aspect_ratio,
},
modality=modality,
)
],
)
return {
"mm_items": mm_items,
}
@@ -223,6 +223,8 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
f"{len(pixel_values)} vs. {len(tgt_sizes)}"
)
# Track slices per image (like vLLM's num_slices)
slices_per_image: List[int] = []
pixel_values_flat: List[torch.Tensor] = []
tgt_sizes_flat: List[torch.Tensor] = []
for pixel_b, tgt_b in zip(pixel_values, tgt_sizes):
@@ -231,6 +233,7 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
raise ValueError(
"Inconsistent N lengths, found: " f"{len(pixel_b)} vs {len(tgt_b)}"
)
slices_per_image.append(len(pixel_b))
for pixel_n, tgt_n in zip(pixel_b, tgt_b):
pixel_values_flat += [pixel_n]
tgt_sizes_flat += [tgt_n]
@@ -250,14 +253,23 @@ class MiniCPMMultimodalProcessor(BaseMultimodalProcessor):
image_offsets.extend(slice_offsets)
image_offsets = sorted(image_offsets)
# Create one item per image, each with its own slices and offsets
if len(pixel_values) != 0:
item = MultimodalDataItem(
feature=pixel_values,
offsets=image_offsets,
model_specific_data={"tgt_size": tgt_sizes_flat},
modality=Modality.IMAGE,
)
items += [item]
pv_idx = 0
offset_idx = 0
for num_slices in slices_per_image:
items.append(
MultimodalDataItem(
feature=pixel_values[pv_idx : pv_idx + num_slices],
offsets=image_offsets[offset_idx : offset_idx + num_slices],
model_specific_data={
"tgt_size": tgt_sizes_flat[pv_idx : pv_idx + num_slices]
},
modality=Modality.IMAGE,
)
)
pv_idx += num_slices
offset_idx += num_slices
if (
"audio_features" in res
@@ -61,7 +61,7 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor):
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
offsets=[offset],
precomputed_embeddings=embedding_slice,
)
)
@@ -469,7 +469,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
mm_items.append(
MultimodalDataItem(
modality=modality,
offsets=offset,
offsets=[offset],
precomputed_embeddings=embedding_slice,
)
)
+2 -2
View File
@@ -34,13 +34,13 @@ class TestMultimodalInputsFromDict(unittest.TestCase):
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]})
# Splitting happens at the processor layer, not in from_dict.
# from_dict just reconstructs and passes through.
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)