diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 747b70382..7d375459a 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index f743aaf9c..f80d4a9f0 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -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) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 0c178204a..7cfa7d50a 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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, ) diff --git a/python/sglang/srt/models/llava.py b/python/sglang/srt/models/llava.py index 2d1e69dbc..e07ca7418 100644 --- a/python/sglang/srt/models/llava.py +++ b/python/sglang/srt/models/llava.py @@ -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 diff --git a/python/sglang/srt/models/minicpmv.py b/python/sglang/srt/models/minicpmv.py index cd4489152..588c356a4 100644 --- a/python/sglang/srt/models/minicpmv.py +++ b/python/sglang/srt/models/minicpmv.py @@ -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) diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index ea17322b5..773ce8620 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -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) diff --git a/python/sglang/srt/multimodal/processors/internvl.py b/python/sglang/srt/multimodal/processors/internvl.py index 90624f86f..c95d495ad 100644 --- a/python/sglang/srt/multimodal/processors/internvl.py +++ b/python/sglang/srt/multimodal/processors/internvl.py @@ -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, diff --git a/python/sglang/srt/multimodal/processors/llava.py b/python/sglang/srt/multimodal/processors/llava.py index bdc28f1d9..8729e8547 100644 --- a/python/sglang/srt/multimodal/processors/llava.py +++ b/python/sglang/srt/multimodal/processors/llava.py @@ -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, } diff --git a/python/sglang/srt/multimodal/processors/minicpm.py b/python/sglang/srt/multimodal/processors/minicpm.py index bad2cbe3d..613079e04 100644 --- a/python/sglang/srt/multimodal/processors/minicpm.py +++ b/python/sglang/srt/multimodal/processors/minicpm.py @@ -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 diff --git a/python/sglang/srt/multimodal/processors/qwen_audio.py b/python/sglang/srt/multimodal/processors/qwen_audio.py index 817b88050..90c2ffd45 100644 --- a/python/sglang/srt/multimodal/processors/qwen_audio.py +++ b/python/sglang/srt/multimodal/processors/qwen_audio.py @@ -61,7 +61,7 @@ class Qwen2AudioMultimodalProcessor(BaseMultimodalProcessor): mm_items.append( MultimodalDataItem( modality=modality, - offsets=offset, + offsets=[offset], precomputed_embeddings=embedding_slice, ) ) diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index d76616596..95c7cd21a 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -469,7 +469,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): mm_items.append( MultimodalDataItem( modality=modality, - offsets=offset, + offsets=[offset], precomputed_embeddings=embedding_slice, ) ) diff --git a/python/sglang/test/test_mm_utils.py b/python/sglang/test/test_mm_utils.py index 1d8585417..bc8fc63de 100644 --- a/python/sglang/test/test_mm_utils.py +++ b/python/sglang/test/test_mm_utils.py @@ -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)