From d6d9f12444ebb55f6d581341128c8eeb1b69a285 Mon Sep 17 00:00:00 2001 From: Mick Date: Sun, 24 May 2026 11:51:24 +0800 Subject: [PATCH] [VLM] adopt simplified get_rope_index for image-only requests (#26100) --- .../srt/multimodal/processors/qwen_vl.py | 123 +++++++++++++++++- 1 file changed, 120 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index 2023e7ba6..962cf9b43 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -2,7 +2,7 @@ import math import os import re import time -from typing import List, Union +from typing import List, Optional, Union import numpy as np import torch @@ -436,6 +436,110 @@ class QwenVLImageProcessor(SGLangBaseProcessor): mrope_position_delta = mrope_position_delta.reshape(-1, 1) return mrope_positions, mrope_position_delta + @staticmethod + def _as_grid_batch(value): + if value is None: + return None + if isinstance(value, torch.Tensor): + return value.unsqueeze(0) if value.ndim == 1 else value + tensor = torch.as_tensor(value, dtype=torch.long) + return tensor.unsqueeze(0) if tensor.ndim == 1 else tensor + + def _compute_image_only_mrope_positions_from_offsets( + self, + input_len: int, + mm_items: List[MultimodalDataItem], + dtype: torch.dtype, + device: torch.device, + ) -> Optional[tuple[torch.Tensor, torch.Tensor]]: + """instead of calling get_rope_index, build mrope position from mm_items.offsets and image_grid_thw of each image + basically a simplified version of get_rope_index for image-only reqs + """ + if self.model_type not in ( + "qwen3_vl", + "qwen3_vl_moe", + "qwen3_5", + "qwen3_5_moe", + "intern_s2_preview", + ): + return None + + image_items = [item for item in mm_items if item.is_image()] + if not image_items or len(image_items) != len(mm_items): + return None + + spatial_merge_size = self.hf_config.vision_config.spatial_merge_size + sorted_items = sorted(image_items, key=lambda item: item.offsets[0][0]) + position_segments = [] + st = 0 + next_pos = 0 + + for item in sorted_items: + if item.offsets is None or len(item.offsets) != 1: + return None + + start, end = item.offsets[0] + if start < st or end >= input_len: + return None + + text_len = start - st + if text_len > 0: + position_segments.append( + torch.arange(text_len, dtype=dtype, device=device) + .view(1, -1) + .expand(3, -1) + + next_pos + ) + next_pos += text_len + + grid = self._as_grid_batch(item.model_specific_data.get("image_grid_thw")) + if grid is None or grid.shape[0] != 1: + return None + t, h, w = [int(x) for x in grid[0].tolist()] + llm_grid_t = t + llm_grid_h = h // spatial_merge_size + llm_grid_w = w // spatial_merge_size + num_image_tokens = llm_grid_t * llm_grid_h * llm_grid_w + if num_image_tokens != end - start + 1: + return None + + t_index = ( + torch.arange(llm_grid_t, dtype=dtype, device=device) + .view(-1, 1) + .expand(llm_grid_t, llm_grid_h * llm_grid_w) + .reshape(-1) + ) + h_index = ( + torch.arange(llm_grid_h, dtype=dtype, device=device) + .view(1, -1, 1) + .expand(llm_grid_t, llm_grid_h, llm_grid_w) + .reshape(-1) + ) + w_index = ( + torch.arange(llm_grid_w, dtype=dtype, device=device) + .view(1, 1, -1) + .expand(llm_grid_t, llm_grid_h, llm_grid_w) + .reshape(-1) + ) + position_segments.append( + torch.stack([t_index, h_index, w_index]) + next_pos + ) + next_pos += max(llm_grid_t, llm_grid_h, llm_grid_w) + st = end + 1 + + if st < input_len: + text_len = input_len - st + position_segments.append( + torch.arange(text_len, dtype=dtype, device=device) + .view(1, -1) + .expand(3, -1) + + next_pos + ) + + mrope_positions = torch.cat(position_segments, dim=1).unsqueeze(1) + mrope_position_delta = (mrope_positions.max() + 1 - input_len).reshape(1, 1) + return mrope_positions, mrope_position_delta + def get_mm_data(self, prompt, embeddings, **kwargs): img_grid_thw = kwargs.get("img_grid_thw", None) video_grid_thw = kwargs.get("video_grid_thw", None) @@ -636,6 +740,18 @@ class QwenVLImageProcessor(SGLangBaseProcessor): video_grid_thw = first_video.get("video_grid_thw") mrope_result = self._get_precomputed_mrope_from_output(ret) + if mrope_result is None: + if ( + video_grid_thw is None + and second_per_grid_ts is None + and audio_feature_lengths is None + ): + mrope_result = self._compute_image_only_mrope_positions_from_offsets( + input_len=input_ids.numel(), + mm_items=mm_items, + dtype=input_ids.dtype, + device=input_ids.device, + ) if mrope_result is None: mrope_result = MRotaryEmbedding.get_rope_index( spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, @@ -659,9 +775,10 @@ class QwenVLImageProcessor(SGLangBaseProcessor): self.hf_config, "position_id_per_seconds", None ), ) + mrope_positions, mrope_position_delta = mrope_result - if mrope_positions.ndim == 3: - mrope_positions = mrope_positions.squeeze(1) + mrope_positions = mrope_positions.squeeze(1) + get_rope_index_time = time.perf_counter() logger.debug( f"[QwenVLProcessor Perf] {rid=}, "