[VLM] adopt simplified get_rope_index for image-only requests (#26100)

This commit is contained in:
Mick
2026-05-24 11:51:24 +08:00
committed by GitHub
parent af8f66940e
commit d6d9f12444
@@ -2,7 +2,7 @@ import math
import os import os
import re import re
import time import time
from typing import List, Union from typing import List, Optional, Union
import numpy as np import numpy as np
import torch import torch
@@ -436,6 +436,110 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
mrope_position_delta = mrope_position_delta.reshape(-1, 1) mrope_position_delta = mrope_position_delta.reshape(-1, 1)
return mrope_positions, mrope_position_delta 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): def get_mm_data(self, prompt, embeddings, **kwargs):
img_grid_thw = kwargs.get("img_grid_thw", None) img_grid_thw = kwargs.get("img_grid_thw", None)
video_grid_thw = kwargs.get("video_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") video_grid_thw = first_video.get("video_grid_thw")
mrope_result = self._get_precomputed_mrope_from_output(ret) 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: if mrope_result is None:
mrope_result = MRotaryEmbedding.get_rope_index( mrope_result = MRotaryEmbedding.get_rope_index(
spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, 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 self.hf_config, "position_id_per_seconds", None
), ),
) )
mrope_positions, mrope_position_delta = mrope_result 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() get_rope_index_time = time.perf_counter()
logger.debug( logger.debug(
f"[QwenVLProcessor Perf] {rid=}, " f"[QwenVLProcessor Perf] {rid=}, "