[VLM] feat: support chunked vit attention (#14907)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2025-12-15 12:11:02 +08:00
committed by GitHub
co-authored by luoyuan.luo
parent 1ab9b8e0a3
commit 3912ee4991
2 changed files with 363 additions and 8 deletions
+97 -8
View File
@@ -14,6 +14,7 @@
# ==============================================================================
"""Inference-only Qwen3-VL model compatible with HuggingFace weights."""
import logging
import math
import re
from functools import lru_cache, partial
from typing import Callable, Iterable, List, Optional, Tuple, Union
@@ -53,7 +54,7 @@ from sglang.srt.models.qwen3 import Qwen3Model
from sglang.srt.models.utils import RotaryPosMixin, compute_cu_seqlens_from_grid_numpy
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix
from sglang.srt.utils import add_prefix, get_int_env_var
from sglang.srt.utils.hf_transformers_utils import get_processor
logger = logging.getLogger(__name__)
@@ -673,13 +674,101 @@ class Qwen3VLForConditionalGeneration(nn.Module):
image_grid_thw = torch.concat([item.image_grid_thw for item in items], dim=0)
assert pixel_values.dim() == 2, pixel_values.dim()
assert image_grid_thw.dim() == 2, image_grid_thw.dim()
if self.use_data_parallel:
return run_dp_sharded_mrope_vision_model(
self.visual, pixel_values, image_grid_thw.tolist(), rope_type="rope_3d"
)
else:
image_embeds = self.visual(pixel_values, grid_thw=image_grid_thw)
return image_embeds
max_patches_per_call = get_int_env_var("SGLANG_VLM_MAX_PATCHES_PER_VIT", 0)
max_images_per_call = get_int_env_var("SGLANG_VLM_MAX_IMAGES_PER_VIT", 0)
if max_patches_per_call == 0 and max_images_per_call == 0:
if self.use_data_parallel:
return run_dp_sharded_mrope_vision_model(
self.visual,
pixel_values,
image_grid_thw.tolist(),
rope_type="rope_3d",
)
else:
return self.visual(pixel_values, grid_thw=image_grid_thw)
# compute the number of patches per image and the slice positions in pixel_values
grid_thw_list = (
image_grid_thw.tolist()
) # List[List[int]], each is [T, H, W] or similar
patches_per_image = [int(math.prod(g)) for g in grid_thw_list]
num_images = len(patches_per_image)
# cumulative sum used to slice pixel_values along the image dimension
cum_patches = [0]
for p in patches_per_image:
cum_patches.append(cum_patches[-1] + p)
total_patches = cum_patches[-1]
assert pixel_values.size(0) == total_patches, (
f"pixel_values rows ({pixel_values.size(0)}) "
f"!= total patches ({total_patches})"
)
# split into chunks in image order, each chunk obeys the patch/image limits
all_chunk_embeds: List[torch.Tensor] = []
img_start = 0
while img_start < num_images:
img_end = img_start
patches_in_chunk = 0
images_in_chunk = 0
# try to pack more images into the current chunk until some limit would be exceeded
while img_end < num_images:
next_patches = patches_per_image[img_end]
# if adding this image would exceed the patch limit, stop
if (
max_patches_per_call > 0
and patches_in_chunk + next_patches > max_patches_per_call
):
break
# if adding this image would exceed the image-count limit, also stop
if (
max_images_per_call > 0
and images_in_chunk + 1 > max_images_per_call
):
break
patches_in_chunk += next_patches
images_in_chunk += 1
img_end += 1
# extreme case: the first image alone exceeds the patch limit -> at least ensure img_end > img_start
if img_end == img_start:
img_end = img_start + 1
patches_in_chunk = patches_per_image[img_start]
images_in_chunk = 1
# slice pixel_values and grid_thw according to [img_start:img_end]
patch_start = cum_patches[img_start]
patch_end = cum_patches[img_end]
pixel_chunk = pixel_values[patch_start:patch_end]
grid_chunk = image_grid_thw[img_start:img_end]
# run ViT once on this chunk without extra padding
if self.use_data_parallel:
chunk_embeds = run_dp_sharded_mrope_vision_model(
self.visual,
pixel_chunk,
grid_chunk.tolist(),
rope_type="rope_3d",
)
else:
chunk_embeds = self.visual(pixel_chunk, grid_thw=grid_chunk)
# chunk_embeds: (sum_patches_after_merge_this_chunk, hidden)
all_chunk_embeds.append(chunk_embeds)
# next batch
img_start = img_end
# concatenate back the full image embedding sequence
return torch.cat(all_chunk_embeds, dim=0)
def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
# in qwen-vl, last dim is the same