[VLM] fix: fix only the grids from last split mm item is collected for qwen-vl (#26094)

This commit is contained in:
Mick
2026-05-25 09:09:46 +08:00
committed by GitHub
parent ed179bf9b2
commit 72c1582d4e
@@ -374,13 +374,12 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
return input_ids, offsets, modality_list
def compute_mrope_positions(self, input_ids, mm_items):
image_grid_thw = None
video_grid_thw = None
for item in mm_items:
if "image_grid_thw" in item.model_specific_data:
image_grid_thw = item.model_specific_data["image_grid_thw"]
if "video_grid_thw" in item.model_specific_data:
video_grid_thw = item.model_specific_data["video_grid_thw"]
image_grid_thw = self._concat_mm_item_grid(
mm_items, "image_grid_thw", Modality.IMAGE
)
video_grid_thw = self._concat_mm_item_grid(
mm_items, "video_grid_thw", Modality.VIDEO
)
input_ids_tensor = torch.tensor(input_ids, dtype=torch.long).unsqueeze(0)
mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index(
@@ -529,6 +528,32 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
mrope_position_delta = (mrope_positions.max() + 1 - input_len).reshape(1, 1)
return mrope_positions, mrope_position_delta
@classmethod
def _concat_mm_item_grid(cls, mm_items: list[MultimodalDataItem], key, modality):
grids = []
for item in mm_items:
if not item.is_modality(modality):
continue
grid = cls._as_grid_batch(item.model_specific_data.get(key))
if grid is not None:
grids.append(grid)
if not grids:
return None
if len(grids) == 1:
return grids[0]
return torch.cat(grids, dim=0)
@classmethod
def _get_grid_from_output_or_items(
cls, ret, mm_items, key, modality, input_data=None
):
grid = cls._get_processor_output_value(ret, key)
if grid is None:
grid = cls._concat_mm_item_grid(mm_items, key, modality)
if grid is None and input_data and isinstance(input_data[0], dict):
grid = input_data[0].get(key)
return grid
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)
@@ -646,7 +671,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
rid = getattr(request_obj, "rid", "anonymous_rid")
video_metadata = None
if base_output.videos:
if base_output.videos and not isinstance(base_output.videos[0], dict):
videos_processed = [
await preprocess_video(video, video_config=self.video_config)
for video in base_output.videos
@@ -683,9 +708,11 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
audio_item.feature_attention_mask, dim=1
)
second_per_grid_ts = getattr(ret, "second_per_grid_ts", None)
second_per_grid_ts = self._get_processor_output_value(ret, "second_per_grid_ts")
if second_per_grid_ts is None:
second_per_grid_ts = getattr(ret, "video_second_per_grid", None)
second_per_grid_ts = self._get_processor_output_value(
ret, "video_second_per_grid"
)
process_time = time.perf_counter()
@@ -712,15 +739,16 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
else:
padded_input_ids = list(padded_input_ids)
image_grid_thw = self._get_processor_output_value(ret, "image_grid_thw")
if image_grid_thw is None and image_data and isinstance(image_data[0], dict):
image_grid_thw = image_data[0].get("image_grid_thw")
video_grid_thw = self._get_processor_output_value(ret, "video_grid_thw")
if video_grid_thw is None and request_obj.video_data:
first_video = request_obj.video_data[0]
if isinstance(first_video, dict):
video_grid_thw = first_video.get("video_grid_thw")
image_grid_thw = self._get_grid_from_output_or_items(
ret, mm_items, "image_grid_thw", Modality.IMAGE, image_data
)
video_grid_thw = self._get_grid_from_output_or_items(
ret,
mm_items,
"video_grid_thw",
Modality.VIDEO,
request_obj.video_data,
)
mrope_result = self._get_precomputed_mrope_from_output(ret)
if mrope_result is None: