diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 76c463777..47f570821 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -350,6 +350,9 @@ 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) + # Release & Resume Memory SGLANG_MEMORY_SAVER_CUDA_GRAPH = EnvBool(False) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 5b59dce76..fea6bee0c 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -2,9 +2,11 @@ Multi-modality utils """ +import copy import hashlib import pickle from abc import abstractmethod +from collections import defaultdict from typing import Any, Callable, Dict, List, Literal, Optional, Tuple import numpy as np @@ -322,22 +324,46 @@ class MultiModalityDataPaddingPatternMultimodalTokens(MultiModalityDataPaddingPa input_ids_tensor = torch.as_tensor(input_ids) - # Create mapping of token_ids to pad_values for each modality - token_to_pad_mapping = {} + # 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) - 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}") + 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, + } - # 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 modality, items in items_by_modality.items(): + token_id = token_id_map.get(modality) + + 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 ret_input_ids = input_ids_tensor.tolist() return ret_input_ids @@ -1238,3 +1264,204 @@ def extend_mrope_positions_for_retracted_request( # Concatenate to the original mrope_positions return torch.cat([mrope_positions, output_positions], dim=1) + + +def _get_length(value): + if value is None: + return None + if isinstance(value, torch.Tensor): + return value.shape[0] if value.ndim > 0 else None + if isinstance(value, np.ndarray): + return value.shape[0] if value.ndim > 0 else None + if isinstance(value, (list, tuple)): + return len(value) + return None + + +def _slice_value(value, start, end): + if isinstance(value, torch.Tensor): + return value[start:end] + if isinstance(value, np.ndarray): + return value[start:end] + if isinstance(value, list): + return value[start:end] + if isinstance(value, tuple): + return value[start:end] + try: + return value[start:end] + except Exception: + return value + + +def _slice_model_data( + data: dict, + index: int, + start: int, + end: int, + num_items: int, + total_feature_len: Optional[int], +): + sliced = {} + for key, value in data.items(): + length = _get_length(value) + if length == num_items: + sliced[key] = _slice_value(value, index, index + 1) + elif total_feature_len is not None and length == total_feature_len: + sliced[key] = _slice_value(value, start, end) + else: + sliced[key] = value + return sliced + + +def get_new_expanded_mm_items(original_mm_items): + expanded_mm_items = [] + for item in original_mm_items: + is_bundled = item.offsets is not None and len(item.offsets) > 1 + + if is_bundled: + num_items = len(item.offsets) + + if item.is_image(): + 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) + continue + + patches_per_item = [] + for grid in image_grid_thw: + grid_tensor = torch.as_tensor(grid, dtype=torch.long) + patches_per_item.append(int(torch.prod(grid_tensor).item())) + + cumulative = torch.cumsum( + torch.tensor(patches_per_item, dtype=torch.long), dim=0 + ) + slice_indices = [0] + cumulative.tolist() + + feature_len = _get_length(item.feature) + if feature_len is None: + feature_len = _get_length(item.precomputed_embeddings) + if feature_len is None or slice_indices[-1] != feature_len: + expanded_mm_items.append(item) + continue + + total_feature_len = feature_len + for i in range(num_items): + start, end = slice_indices[i], slice_indices[i + 1] + new_item = copy.deepcopy(item) + if item.feature is not None: + new_item.feature = _slice_value(item.feature, start, end) + if item.precomputed_embeddings is not None: + new_item.precomputed_embeddings = _slice_value( + item.precomputed_embeddings, start, end + ) + new_item.offsets = [item.offsets[i]] + new_item.model_specific_data = _slice_model_data( + item.model_specific_data, + index=i, + start=start, + end=end, + num_items=num_items, + total_feature_len=total_feature_len, + ) + new_item.hash = None + expanded_mm_items.append(new_item) + + 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) + continue + + # video_grid_thw shape: [num_videos, 3] where each row is [T, H, W] + # When T > 1, item.offsets contains frames (num_items = total frames) + # grid_len = num_videos, num_items = sum(T for each video) = total frames + grid_len = _get_length(video_grid_thw) + num_videos = grid_len + + # Calculate total frames and frames per video + frames_per_video = [] + total_frames = 0 + for i in range(num_videos): + grid = video_grid_thw[i] + if isinstance(grid, torch.Tensor): + T = int(grid[0].item()) # T is the first element [T, H, W] + else: + grid_tensor = torch.as_tensor(grid, dtype=torch.long) + T = int(grid_tensor[0].item()) + frames_per_video.append(T) + total_frames += T + + # num_items should equal total_frames when T > 1 + if num_items != total_frames: + expanded_mm_items.append(item) + continue + + # Calculate patches per video: T * H * W for each video + patches_per_video = [] + for i in range(num_videos): + grid = video_grid_thw[i] + if isinstance(grid, torch.Tensor): + patches_per_video.append(int(torch.prod(grid).item())) + else: + grid_tensor = torch.as_tensor(grid, dtype=torch.long) + patches_per_video.append(int(torch.prod(grid_tensor).item())) + + # Calculate cumulative patches to get slice indices for each video + cumulative = torch.cumsum( + torch.tensor(patches_per_video, dtype=torch.long), dim=0 + ) + slice_indices = [0] + cumulative.tolist() + + feature_len = _get_length(item.feature) + if feature_len is None: + feature_len = _get_length(item.precomputed_embeddings) + if feature_len is None or slice_indices[-1] != feature_len: + expanded_mm_items.append(item) + continue + + total_feature_len = feature_len + # Group frames by video: calculate frame indices for each video + frame_start_indices = [0] + for i in range(num_videos): + frame_start_indices.append( + frame_start_indices[-1] + frames_per_video[i] + ) + + # Expand each video into a separate item + for video_idx in range(num_videos): + start, end = ( + slice_indices[video_idx], + slice_indices[video_idx + 1], + ) + frame_start, frame_end = ( + frame_start_indices[video_idx], + frame_start_indices[video_idx + 1], + ) + + new_item = copy.deepcopy(item) + if item.feature is not None: + new_item.feature = _slice_value(item.feature, start, end) + if item.precomputed_embeddings is not None: + new_item.precomputed_embeddings = _slice_value( + item.precomputed_embeddings, start, end + ) + # Group offsets for this video (all frames of this video) + new_item.offsets = item.offsets[frame_start:frame_end] + # For video_grid_thw, slice the corresponding row [T, H, W] for this video + new_item.model_specific_data = _slice_model_data( + item.model_specific_data, + index=video_idx, + start=start, + end=end, + num_items=num_videos, + total_feature_len=total_feature_len, + ) + new_item.hash = None + expanded_mm_items.append(new_item) + else: + expanded_mm_items.append(item) + + else: + expanded_mm_items.append(item) + return expanded_mm_items diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 07a6455a4..4b41ad9a7 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -334,8 +334,18 @@ class MultimodalInputs: @staticmethod def from_dict(obj: dict): + # Check if MM splitting is enabled + if not envs.SGLANG_ENABLE_MM_SPLITTING.get(): + mm_items = obj["mm_items"] + else: + 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. + mm_items = get_new_expanded_mm_items(original_mm_items) + ret = MultimodalInputs( - mm_items=obj["mm_items"], + mm_items=mm_items, ) assert isinstance(ret.mm_items, list)