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
@@ -0,0 +1,101 @@
import pytest
import torch
from sglang.srt.managers.schedule_batch import Modality
from sglang.srt.multimodal.processors.base_processor import MultimodalSpecialTokens
from sglang.srt.multimodal.processors.glm4v import Glm4vImageProcessor
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _processor():
processor = Glm4vImageProcessor.__new__(Glm4vImageProcessor)
processor.IM_TOKEN_ID = 99
processor.VIDEO_START_TOKEN_ID = 101
processor.VIDEO_END_TOKEN_ID = 102
return processor
def test_glm4v_partitions_shared_image_and_video_token_offsets():
processor = _processor()
mm_tokens = MultimodalSpecialTokens(image_token_id=99, video_token_id=99)
input_ids = torch.tensor(
[
1,
99,
99, # image 1
2,
99,
99,
99, # image 2
3,
101, # begin video
4,
99,
99, # frame 1
5,
99,
99,
99, # frame 2
6,
102, # end video
7,
]
)
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.IMAGE) == [
(1, 2),
(4, 6),
]
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.VIDEO) == [
(10, 11),
(13, 15),
]
def test_glm4v_keeps_interleaved_media_offsets_in_their_modalities():
processor = _processor()
mm_tokens = MultimodalSpecialTokens(image_token_id=99, video_token_id=99)
input_ids = torch.tensor(
[
101,
99,
99,
102, # video 1
8,
99, # image
9,
101,
99,
10,
99,
102, # video 2
]
)
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.IMAGE) == [
(5, 5)
]
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.VIDEO) == [
(1, 2),
(8, 8),
(10, 10),
]
def test_glm4v_uses_default_offsets_when_token_ids_are_distinct():
processor = _processor()
mm_tokens = MultimodalSpecialTokens(image_token_id=99, video_token_id=100)
input_ids = torch.tensor([99, 99, 1, 100, 100])
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.IMAGE) == [
(0, 1)
]
assert processor.get_mm_item_offsets(input_ids, mm_tokens, Modality.VIDEO) == [
(3, 4)
]
if __name__ == "__main__":
raise SystemExit(pytest.main([__file__, "-v"]))