fix(glm4v): disambiguate mixed image video offsets (#37971)

Co-authored-by: duxin <xinheng.dx@alibaba-inc.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
Xin Du
2026-09-05 13:16:16 +08:00
committed by GitHub
co-authored by duxin Xiaoyu Zhang
parent 0454c074b4
commit e980c1a2f1
3 changed files with 153 additions and 7 deletions
@@ -1471,6 +1471,22 @@ class BaseMultimodalProcessor(ABC):
return list(zip(indices_start.tolist(), indices_end.tolist()))
def get_mm_item_offsets(
self,
input_ids: torch.Tensor,
mm_tokens: MultimodalSpecialTokens,
modality: Modality,
) -> List[Tuple[int, int]]:
"""Return placeholder offsets belonging to one modality.
Processors that reuse one token ID for multiple modalities can override
this method and use surrounding boundary tokens to disambiguate spans.
"""
mm_token_id = mm_tokens.get_token_id_by_modality(modality)
if mm_token_id is None:
raise ValueError(f"No token id found for modality: {modality}")
return self.get_mm_items_offset(input_ids, mm_token_id)
def collect_mm_items_from_processor_output(
self, data_dict: dict, modality: Modality = None
) -> List[MultimodalDataItem]:
@@ -1834,12 +1850,10 @@ class BaseMultimodalProcessor(ABC):
for mm_item in all_collected_items:
if mm_item.offsets is not None:
continue
mm_token_id = mm_tokens.get_token_id_by_modality(mm_item.modality)
if mm_token_id is None:
raise ValueError(f"No token id found for modality: {mm_item.modality}")
mm_item.offsets = self.get_mm_items_offset(
mm_item.offsets = self.get_mm_item_offsets(
input_ids=input_ids,
mm_token_id=mm_token_id,
mm_tokens=mm_tokens,
modality=mm_item.modality,
)
# Split bundled items into per-image/video items for better cache granularity
@@ -1,12 +1,12 @@
import asyncio
import math
from typing import List, Union
from typing import List, Tuple, Union
import numpy as np
import torch
from sglang.srt.layers.rotary_embedding import MRotaryEmbedding
from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput
from sglang.srt.managers.schedule_batch import Modality, MultimodalProcessorOutput
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration
from sglang.srt.models.glm4v_moe import Glm4vMoeForConditionalGeneration
from sglang.srt.multimodal.processors.base_processor import (
@@ -296,6 +296,37 @@ class Glm4vImageProcessor(SGLangBaseProcessor):
video_token_id=self.IM_TOKEN_ID,
).build(_processor)
def get_mm_item_offsets(
self,
input_ids: torch.Tensor,
mm_tokens: MultimodalSpecialTokens,
modality: Modality,
) -> List[Tuple[int, int]]:
"""Disambiguate image and video spans sharing ``IM_TOKEN_ID``."""
if (
modality not in (Modality.IMAGE, Modality.VIDEO)
or mm_tokens.image_token_id != mm_tokens.video_token_id
):
return super().get_mm_item_offsets(input_ids, mm_tokens, modality)
all_offsets = self.get_mm_items_offset(input_ids, self.IM_TOKEN_ID)
video_ranges = self.get_mm_items_offset_by_pair(
input_ids,
self.VIDEO_START_TOKEN_ID,
self.VIDEO_END_TOKEN_ID,
)
def is_video_offset(offset: Tuple[int, int]) -> bool:
start, end = offset
return any(
video_start <= start and end <= video_end
for video_start, video_end in video_ranges
)
if modality == Modality.VIDEO:
return [offset for offset in all_offsets if is_video_offset(offset)]
return [offset for offset in all_offsets if not is_video_offset(offset)]
def compute_mrope_positions(self, input_ids, mm_items):
image_grid_thw = None
video_grid_thw = None