[VLM] reuse pretokenized ids from preprocessed input for qwen-vl (#26116)

This commit is contained in:
Mick
2026-05-23 16:01:04 +08:00
committed by GitHub
parent 8b9fb13c4a
commit 19b60a4f9e
2 changed files with 110 additions and 24 deletions
@@ -381,6 +381,7 @@ class MultimodalProcessorOutput:
mm_items: List[MultimodalDataItem] mm_items: List[MultimodalDataItem]
input_ids: Optional[List[int]] = None input_ids: Optional[List[int]] = None
padded_input_ids: Optional[List[int]] = None
# image # image
im_token_id: Optional[int] = None im_token_id: Optional[int] = None
@@ -414,6 +415,7 @@ class MultimodalProcessorOutput:
return MultimodalProcessorOutput( return MultimodalProcessorOutput(
mm_items=d["mm_items"], mm_items=d["mm_items"],
input_ids=d.get("input_ids"), input_ids=d.get("input_ids"),
padded_input_ids=d.get("padded_input_ids"),
im_token_id=d.get("im_token_id"), im_token_id=d.get("im_token_id"),
im_start_id=d.get("im_start_id"), im_start_id=d.get("im_start_id"),
im_end_id=d.get("im_end_id"), im_end_id=d.get("im_end_id"),
@@ -430,6 +432,26 @@ class MultimodalProcessorOutput:
visible_frame_counts=d.get("visible_frame_counts"), visible_frame_counts=d.get("visible_frame_counts"),
) )
@staticmethod
def build_padded_input_ids(input_ids, mm_items: List[MultimodalDataItem]):
"""pad the input_ids with mm_items if it's not already padded"""
if input_ids is None or not mm_items:
return None
for item in mm_items:
if item.pad_value is None or item.offsets is None:
return None
if isinstance(input_ids, torch.Tensor):
padded_input_ids = input_ids.flatten().tolist()
else:
padded_input_ids = list(input_ids)
for item in mm_items:
for start, end in item.offsets:
padded_input_ids[start : end + 1] = [item.pad_value] * (end - start + 1)
return padded_input_ids
@dataclasses.dataclass @dataclasses.dataclass
class MultimodalInputs: class MultimodalInputs:
@@ -437,6 +459,7 @@ class MultimodalInputs:
# items of data # items of data
mm_items: List[MultimodalDataItem] mm_items: List[MultimodalDataItem]
padded_input_ids: Optional[List[int]] = None
image_pad_len: Optional[list] = None image_pad_len: Optional[list] = None
num_image_tokens: Optional[int] = None num_image_tokens: Optional[int] = None
@@ -478,6 +501,7 @@ class MultimodalInputs:
ret = MultimodalInputs( ret = MultimodalInputs(
mm_items=mm_items, mm_items=mm_items,
padded_input_ids=obj.padded_input_ids,
) )
assert isinstance(ret.mm_items, list) assert isinstance(ret.mm_items, list)
@@ -400,6 +400,42 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
) )
return mrope_positions.squeeze(1), mrope_position_delta return mrope_positions.squeeze(1), mrope_position_delta
# TODO: consider moving it to SGLangBaseProcessor
@staticmethod
def _get_processor_output_value(ret, key):
"""get value with key from returned value of processor"""
if ret is None:
return None
if hasattr(ret, "get"):
value = ret.get(key)
if value is not None:
return value
return getattr(ret, key, None)
def _get_precomputed_mrope_from_output(self, ret):
"""get the precomputed mrope from processor output"""
mrope_positions = self._get_processor_output_value(ret, "mrope_positions")
mrope_position_delta = self._get_processor_output_value(
ret, "mrope_position_delta"
)
if mrope_positions is None or mrope_position_delta is None:
return None
mrope_positions = torch.as_tensor(mrope_positions)
if mrope_positions.ndim == 3:
if mrope_positions.shape[1] != 1:
return None
mrope_positions = mrope_positions.squeeze(1)
if mrope_positions.ndim != 2 or mrope_positions.shape[0] != 3:
return None
mrope_position_delta = torch.as_tensor(mrope_position_delta)
if mrope_position_delta.ndim == 0:
mrope_position_delta = mrope_position_delta.reshape(1, 1)
elif mrope_position_delta.ndim == 1:
mrope_position_delta = mrope_position_delta.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)
@@ -561,6 +597,27 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
process_time = time.perf_counter() process_time = time.perf_counter()
input_ids = input_ids.flatten() input_ids = input_ids.flatten()
base_input_ids = getattr(base_output, "input_ids", None)
if (
isinstance(base_input_ids, list)
and len(base_input_ids) == input_ids.numel()
):
# reuse preprocess input if it already carries list of input_ids
input_ids_list = base_input_ids
else:
input_ids_list = input_ids.tolist()
# look for if padded_input_ids already exists before computing
padded_input_ids = self._get_processor_output_value(ret, "padded_input_ids")
if padded_input_ids is None:
padded_input_ids = MultimodalProcessorOutput.build_padded_input_ids(
input_ids_list, mm_items
)
elif isinstance(padded_input_ids, torch.Tensor):
# reuse existing padded_input_ids
padded_input_ids = padded_input_ids.flatten().tolist()
else:
padded_input_ids = list(padded_input_ids)
image_grid_thw = None image_grid_thw = None
if hasattr(ret, "image_grid_thw"): if hasattr(ret, "image_grid_thw"):
@@ -578,29 +635,33 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
if isinstance(first_video, dict): if isinstance(first_video, dict):
video_grid_thw = first_video.get("video_grid_thw") video_grid_thw = first_video.get("video_grid_thw")
mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index( mrope_result = self._get_precomputed_mrope_from_output(ret)
spatial_merge_size=self.hf_config.vision_config.spatial_merge_size, if mrope_result is None:
image_token_id=self.mm_tokens.image_token_id, mrope_result = MRotaryEmbedding.get_rope_index(
video_token_id=self.mm_tokens.video_token_id, spatial_merge_size=self.hf_config.vision_config.spatial_merge_size,
vision_start_token_id=self.vision_start_token_id, image_token_id=self.mm_tokens.image_token_id,
model_type=self.model_type, video_token_id=self.mm_tokens.video_token_id,
tokens_per_second=getattr( vision_start_token_id=self.vision_start_token_id,
self.hf_config.vision_config, "tokens_per_second", None model_type=self.model_type,
), tokens_per_second=getattr(
# use the expanded token ids self.hf_config.vision_config, "tokens_per_second", None
input_ids=input_ids.unsqueeze(0), ),
image_grid_thw=getattr(ret, "image_grid_thw", None), # use the expanded token ids
video_grid_thw=getattr(ret, "video_grid_thw", None), input_ids=input_ids.unsqueeze(0),
second_per_grid_ts=second_per_grid_ts, image_grid_thw=image_grid_thw,
use_audio_in_video=False, video_grid_thw=video_grid_thw,
audio_seqlens=audio_feature_lengths, second_per_grid_ts=second_per_grid_ts,
audio_token_id=getattr(self.hf_config, "audio_token_id", None), use_audio_in_video=False,
audio_start_token_id=self.audio_start_token_id, audio_seqlens=audio_feature_lengths,
position_id_per_seconds=getattr( audio_token_id=getattr(self.hf_config, "audio_token_id", None),
self.hf_config, "position_id_per_seconds", None audio_start_token_id=self.audio_start_token_id,
), position_id_per_seconds=getattr(
) self.hf_config, "position_id_per_seconds", None
mrope_positions = mrope_positions.squeeze(1) ),
)
mrope_positions, mrope_position_delta = mrope_result
if mrope_positions.ndim == 3:
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=}, "
@@ -612,7 +673,8 @@ class QwenVLImageProcessor(SGLangBaseProcessor):
) )
return MultimodalProcessorOutput( return MultimodalProcessorOutput(
input_ids=input_ids.tolist(), input_ids=input_ids_list,
padded_input_ids=padded_input_ids,
mm_items=mm_items, mm_items=mm_items,
im_start_id=self.vision_start_token_id, im_start_id=self.vision_start_token_id,
im_end_id=self.vision_end_token_id, im_end_id=self.vision_end_token_id,