[VLM] Chunk-aware ViT encoding with per-image cache and lazy device transfer (#22038)

This commit is contained in:
Yuhao Yang
2026-04-04 16:55:17 +08:00
committed by GitHub
parent b5e8c4b9e3
commit 34d5765e2f
7 changed files with 171 additions and 414 deletions
+121 -260
View File
@@ -44,9 +44,6 @@ TensorTransportMode = Literal["cuda_ipc", "auto", "default"]
_GPU_FEATURE_BUFFER: Optional[torch.Tensor] = None _GPU_FEATURE_BUFFER: Optional[torch.Tensor] = None
_BUFFER_OFFSET = 0 _BUFFER_OFFSET = 0
_EXTRA_PRE_TOKENS = 0 # pre chunk extra token (0 for the moment)
_EXTRA_POST_TOKENS = 0 # post chunk extra token (0 for the moment)
_is_default_tensor_transport = None _is_default_tensor_transport = None
@@ -462,109 +459,41 @@ DataEmbeddingFunc = Callable[
] ]
def get_embedding_items_per_chunk_with_extra_padding( def _move_items_to_device(
embedding_items_per_req: List["MultimodalDataItem"], items: List[MultimodalDataItem], device: torch.device
) -> None:
"""Move item features to the target device (in-place, non-blocking)."""
for item in items:
if isinstance(item.feature, torch.Tensor) and item.feature.device != device:
item.feature = item.feature.to(device, non_blocking=True)
def _get_chunked_embedding_full(
data_embedding_func: DataEmbeddingFunc,
embedding_items_per_req: List[MultimodalDataItem],
items_offset: List[Tuple[int, int]],
extend_prefix_len: int, extend_prefix_len: int,
extend_seq_len: int, extend_seq_len: int,
items_offset: List[Tuple[int, int]],
) -> List["MultimodalDataItem"]:
"""
From all multimodal items of a request, select the subset that is "relevant to
this prefill chunk", and allow a small amount of extra padding on both sides
of the chunk boundary (for easier caching or cross-chunk reuse).
Assumptions:
- len(embedding_items_per_req) == len(items_offset)
- items_offset[j] = (start, end), meaning the multimodal tokens of the j-th
item correspond to [start, end) (left-closed, right-open) in the entire
token sequence
- The item order in embedding_items_per_req is one-to-one aligned with
items_offset
Args:
embedding_items_per_req: all items of this modality under the current
request (e.g. each frame in a 500-frame video)
extend_prefix_len: number of tokens already prefilled before the current
chunk
extend_seq_len: number of tokens in the current chunk
items_offset: (start, end) position of each item in the whole sentence
Returns:
The subset of items to feed into ViT for this chunk (preserving the
original order)
"""
assert len(embedding_items_per_req) == len(
items_offset
), f"items_per_req({len(embedding_items_per_req)}) vs items_offset({len(items_offset)}) mismatch"
if extend_seq_len <= 0:
return []
# Current chunk's token range
chunk_start = extend_prefix_len
chunk_end = extend_prefix_len + extend_seq_len
# Current chunk's token range with extra padding
window_start = max(0, chunk_start - _EXTRA_PRE_TOKENS)
window_end = chunk_end + _EXTRA_POST_TOKENS
selected_items: List["MultimodalDataItem"] = []
for item, (start, end) in zip(embedding_items_per_req, items_offset):
if start >= end:
continue
# Check whether this item has overlap with [window_start, window_end)
# If has overlap, add the item into selected_item.
if end > window_start and start < window_end:
selected_items.append(item)
return selected_items
# TODO: To be obsoleted.
def _get_chunked_prefill_embedding(
data_embedding_func: DataEmbeddingFunc,
embedding_items: List[MultimodalDataItem],
items_size: List[int],
prefix_length: List[int],
extend_length: List[int],
items_offset_list: List[List[Tuple[int, int]]],
input_ids: torch.Tensor, input_ids: torch.Tensor,
) -> tuple[torch.Tensor | None, torch.Tensor]: device: torch.device,
# Calculate embedding for each request, try to get it from cache to avoid repeated calculation ) -> Tuple[Optional[torch.Tensor], torch.Tensor]:
embedding_list = [] """
# FIXME(Xinyuan): temporary workaround for eagle3, which may have len(items_size) > len(prefix_length) Fallback: encode all items at once, cache combined result, extract chunk.
max_iterations = min(len(items_size) - 1, len(prefix_length)) Used for non-bundled items or EVS results.
for i in range(max_iterations): """
if items_size[i] == items_size[i + 1]:
continue
embedding_items_per_req = embedding_items[items_size[i] : items_size[i + 1]]
items_offset = items_offset_list[i]
assert items_offset is not None, items_offset
# if all items has been prefixed, we do not need to calculate embedding
if all([offset_end < prefix_length[i] for _, offset_end in items_offset]):
continue
item_hashes = [item.hash for item in embedding_items_per_req] item_hashes = [item.hash for item in embedding_items_per_req]
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes) embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
embedding_per_req = embedding_cache.get(item_hashes) embedding_per_req = embedding_cache.get(item_hashes)
if embedding_per_req is None: if embedding_per_req is None:
_move_items_to_device(embedding_items_per_req, device)
embedding = data_embedding_func(embedding_items_per_req) embedding = data_embedding_func(embedding_items_per_req)
embedding_per_req = ( embedding_per_req = (
EmbeddingResult(embedding=embedding) EmbeddingResult(embedding=embedding)
if isinstance(embedding, torch.Tensor) if isinstance(embedding, torch.Tensor)
else embedding else embedding
) )
if not embedding_cache.set(embedding_items_hash, embedding_per_req): embedding_cache.set(embedding_items_hash, embedding_per_req)
print_warning_once(
"Multimodal embedding cache is full. This typically occurs when a single "
"embedding exceeds the cache size limit. Consider increasing the "
"`SGLANG_VLM_CACHE_SIZE_MB` environment variable or reducing the input "
"embedding size."
)
extend_prefix_len = prefix_length[i]
extend_seq_len = extend_length[i] if i < len(extend_length) else 0
if isinstance(embedding_per_req, EVSEmbeddingResult): if isinstance(embedding_per_req, EVSEmbeddingResult):
item = embedding_items_per_req[0] item = embedding_items_per_req[0]
@@ -584,147 +513,96 @@ def _get_chunked_prefill_embedding(
extend_seq_len=extend_seq_len, extend_seq_len=extend_seq_len,
items_offset=items_offset, items_offset=items_offset,
) )
embedding_list.append(embedding_per_req_chunk) return embedding_per_req_chunk, input_ids
if len(embedding_list) == 0:
return None, input_ids
return torch.concat(embedding_list, dim=0), input_ids
def get_embedding_chunk_remove_extra_padding( def _get_chunked_embedding_by_item(
embedding: torch.Tensor, data_embedding_func: DataEmbeddingFunc,
embedding_items_per_req: List[MultimodalDataItem],
items_offset: List[Tuple[int, int]],
extend_prefix_len: int, extend_prefix_len: int,
extend_seq_len: int, extend_seq_len: int,
items_offset: List[Tuple[int, int]], device: torch.device,
) -> Tuple[Optional[torch.Tensor], int, int]: ) -> Optional[torch.Tensor]:
""" """
From the embedding computed on "items related to this chunk + extra padding", Per-image chunk-aware encoding: only encode images overlapping with the
trim out the token embeddings that are not needed for the current chunk, and current chunk, cache each image individually.
keep only those mm tokens covered by Items must already be split per-image (each item has exactly one offset).
[extend_prefix_len, extend_prefix_len + extend_seq_len).
Assumptions:
- Each (start, end) in items_offset represents an item's multimodal token
interval [start, end) in the whole token sequence, and their order is
consistent with the order of items in `embedding`.
- The layout of `embedding`: each selected item is concatenated in order,
and item j occupies seg_len_j = end_j - start_j rows.
Args:
embedding: output of data_embedding_func(embedding_items_per_chunk),
shape = (T_total, D)
extend_prefix_len: number of tokens before the chunk (prefix_len)
extend_seq_len: number of tokens in this chunk (chunk_len)
items_offset: list of (start, end) for all items of the current request
Returns:
- trimmed_embedding: embedding that contains only the mm tokens needed
by this chunk, concatenated in token order
- num_tokens_before: number of mm tokens "before the chunk" that are
trimmed off (optional info, not used by the current caller)
- num_tokens_after: number of mm tokens "after the chunk" that are
trimmed off (optional info, not used by the current caller)
""" """
if embedding is None or embedding.numel() == 0:
return None, 0, 0
chunk_start = extend_prefix_len chunk_start = extend_prefix_len
chunk_end = extend_prefix_len + extend_seq_len chunk_end = extend_prefix_len + extend_seq_len # exclusive
if extend_seq_len <= 0 or chunk_start >= chunk_end: if extend_seq_len <= 0:
return None, 0, 0 return None
# The window with extra padding # 1. Find items overlapping with current chunk
window_start = max(0, chunk_start - _EXTRA_PRE_TOKENS) # offsets are (start, end) inclusive on both ends
window_end = chunk_end + _EXTRA_POST_TOKENS overlapping = []
for idx, (item, offset) in enumerate(zip(embedding_items_per_req, items_offset)):
start, end = offset
if end >= chunk_start and start < chunk_end:
overlapping.append((idx, item, start, end))
# Iterate item_offset to choose item. if not overlapping:
# We need to forward an embedding_idx to locate the item start-end position in embedding. return None
embedding_idx = 0
kept_slices: List[torch.Tensor] = []
num_tokens_before = 0 # 2. Check per-image cache for each overlapping item
num_tokens_after = 0 cached_embeddings = {} # idx -> tensor
miss_items = [] # (idx, item, start, end)
for start, end in items_offset: for idx, item, start, end in overlapping:
if start >= end: cached = embedding_cache.get_single(item.hash)
continue if cached is not None:
cached_embeddings[idx] = cached.embedding
seg_len = end - start
# Check whether this item has been chosen into embedding_items_per_chunk or not.
selected = end > window_start and start < window_end
if not selected:
# Not in embedding_items_per_chunk, not forward embedding_idx.
continue
# embedding has the whole item
# embedding[embedding_idx : embedding_idx + seg_len]
# Calculate the overlap range between item and the current chunk
overlap_start = max(start, chunk_start)
overlap_end = min(end, chunk_end)
if overlap_start < overlap_end:
# The item has a portion mm tokens in the current chunk
# The offset inside item
local_start = overlap_start - start
local_end = overlap_end - start
# The embedding index
slice_start = embedding_idx + local_start
slice_end = embedding_idx + local_end
kept_slices.append(embedding[slice_start:slice_end])
# Stats the token number before and after this chunk
num_tokens_before += max(0, local_start)
num_tokens_after += max(0, seg_len - local_end)
else: else:
# Although item is chosen into embedding_items_per_chunk as extra padding, miss_items.append((idx, item, start, end))
# Its mm tokens has no overlap with chunk, so don't count into the current
# chunk's embedding.
if end <= chunk_start:
num_tokens_before += seg_len
elif start >= chunk_end:
num_tokens_after += seg_len
# No matter whether this item has overlap with chunk, once it's selected, it # 3. Batch encode all cache-miss items in one ViT call
# counts seg_len in embedding, so embedding_idx has to forward. if miss_items:
embedding_idx += seg_len miss_item_list = [item for _, item, _, _ in miss_items]
_move_items_to_device(miss_item_list, device)
all_miss_embedding = data_embedding_func(miss_item_list)
all_miss_embedding = all_miss_embedding.reshape(
-1, all_miss_embedding.shape[-1]
)
if not kept_slices: # Split output by per-item token count
# No mm tokens in this chunk token_counts = [end - start + 1 for _, _, start, end in miss_items]
return None, num_tokens_before, num_tokens_after split_embeddings = torch.split(all_miss_embedding, token_counts, dim=0)
trimmed_embedding = torch.cat(kept_slices, dim=0) for (idx, item, _, _), emb in zip(miss_items, split_embeddings):
return trimmed_embedding, num_tokens_before, num_tokens_after cached_embeddings[idx] = emb
emb_result = EmbeddingResult(embedding=emb)
embedding_cache.set(item.hash, emb_result)
# 4. Assemble chunk: for each overlapping item, extract the overlap slice
chunk_slices = []
for idx, _, start, end in overlapping:
emb = cached_embeddings[idx] # shape: (end - start + 1, hidden)
overlap_start = max(start, chunk_start)
overlap_end = min(end, chunk_end - 1) # inclusive
local_start = overlap_start - start
local_end = overlap_end - start + 1 # exclusive for slicing
chunk_slices.append(emb[local_start:local_end])
return torch.cat(chunk_slices, dim=0)
# This function is for chunked prefill vit for multiple items in the next feature. def _get_chunked_prefill_embedding(
def _get_chunked_prefill_embedding_for_chunked_items( data_embedding_func: DataEmbeddingFunc,
data_embedding_func: Callable[[List["MultimodalDataItem"]], torch.Tensor], embedding_items: List[MultimodalDataItem],
embedding_items: List["MultimodalDataItem"],
items_size: List[int], items_size: List[int],
prefix_length: List[int], prefix_length: List[int],
extend_length: List[int], extend_length: List[int],
items_offset_list: List[List[Tuple[int, int]]], items_offset_list: List[List[Tuple[int, int]]],
) -> Optional[torch.Tensor]: input_ids: torch.Tensor,
) -> tuple[torch.Tensor | None, torch.Tensor]:
""" """
Multi-modal embedding computation for chunked prefill. Chunked prefill embedding: encode per-request items and extract the chunk.
Items are already split per-image at processor stage.
For each request:
1. Use items_size to split embedding_items into per-request sublists embedding_items_per_req;
2. Use get_embedding_items_per_chunk_with_extra_padding to select the subset of items related to this chunk;
3. Call data_embedding_func (ViT) on this subset to obtain embedding_per_chunk;
4. Concatenate embedding_per_req_chunk for all requests in order.
In this way, the ViT for each request only processes the frames / images related to the current chunk,
avoiding OOM caused by processing all the frames at once.
""" """
# Calculate embedding for each request, try to get it from cache to avoid repeated calculation
embedding_list = [] embedding_list = []
# FIXME(Xinyuan): temporary workaround for eagle3, which may have len(items_size) > len(prefix_length) device = input_ids.device
# FIXME(Xinyuan): temporary workaround for eagle3
max_iterations = min(len(items_size) - 1, len(prefix_length)) max_iterations = min(len(items_size) - 1, len(prefix_length))
for i in range(max_iterations): for i in range(max_iterations):
@@ -734,62 +612,45 @@ def _get_chunked_prefill_embedding_for_chunked_items(
items_offset = items_offset_list[i] items_offset = items_offset_list[i]
assert items_offset is not None, items_offset assert items_offset is not None, items_offset
# if all items has been prefixed, we do not need to calculate embedding extend_prefix_len = prefix_length[i]
if all([offset_end < prefix_length[i] for _, offset_end in items_offset]): extend_seq_len = extend_length[i] if i < len(extend_length) else 0
# Skip if all items already prefilled
if all(offset_end < prefix_length[i] for _, offset_end in items_offset):
continue continue
# 1) Pick up items related with this chunk # Use per-image path when all items have exactly one offset (already
embedding_items_per_chunk = get_embedding_items_per_chunk_with_extra_padding( # split per-image) — this avoids encoding images not in this chunk.
# Fall back to combined path for non-split items or EVS.
is_per_image = all(len(item.offsets) == 1 for item in embedding_items_per_req)
if is_per_image:
chunk_embedding = _get_chunked_embedding_by_item(
data_embedding_func,
embedding_items_per_req, embedding_items_per_req,
extend_prefix_len=prefix_length[i], items_offset,
extend_seq_len=extend_length[i] if i < len(extend_length) else 0, extend_prefix_len,
items_offset=items_offset, extend_seq_len,
) device,
if not embedding_items_per_chunk:
continue
# 2) construct cache key
# embedding_items_hash = MultiModalStaticCache.combine_hashes(
# embedding_items_per_chunk
# )
item_hashes = [item.hash for item in embedding_items_per_chunk]
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
embedding_per_chunk = embedding_cache.get(embedding_items_hash)
if embedding_per_chunk is None:
# ViT forward for items related with per chunk
embedding_per_chunk = data_embedding_func(embedding_items_per_chunk)
embedding_for_cache = embedding_per_chunk.detach().cpu()
if not embedding_cache.set(embedding_items_hash, embedding_for_cache):
print(
"[WARN] Multimodal embedding cache is full. "
"Consider increasing `SGLANG_VLM_CACHE_SIZE_MB` or reducing "
"video frame count / resolution for a single request."
) )
if chunk_embedding is not None:
embedding_list.append(chunk_embedding)
else: else:
target_device = embedding_items_per_req[0].feature.device chunk_embedding, input_ids = _get_chunked_embedding_full(
if embedding_per_chunk.device != target_device: data_embedding_func,
embedding_per_chunk = embedding_per_chunk.to(target_device) embedding_items_per_req,
items_offset,
extend_prefix_len,
extend_seq_len,
input_ids,
device,
)
if chunk_embedding is not None:
embedding_list.append(chunk_embedding)
# 3) remove extra padding from embedding_per_chunk, only keep current chunk part if len(embedding_list) == 0:
# We probably don't need this part. return None, input_ids
# embedding_per_req_chunk, _, _ = get_embedding_chunk_remove_extra_padding( return torch.concat(embedding_list, dim=0), input_ids
# embedding=embedding_per_chunk,
# extend_prefix_len=prefix_len,
# extend_seq_len=chunk_len,
# items_offset=items_offset,
# )
if embedding_per_chunk is not None and embedding_per_chunk.numel() > 0:
embedding_list.append(embedding_per_chunk)
if not embedding_list:
return None
# concat all the request's chunk embedding in token
return torch.cat(embedding_list, dim=0)
def _get_multimodal_mask( def _get_multimodal_mask(
@@ -1730,21 +1730,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
if input_embeds if input_embeds
else None else None
) )
for mm_input in multimodal_inputs:
if mm_input is None:
continue
for mm_item in mm_input.mm_items:
pixel_values = getattr(mm_item, "feature", None)
if isinstance(pixel_values, torch.Tensor):
mm_item.feature = pixel_values.to(self.device, non_blocking=True)
if get_global_server_args().language_only:
precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None
)
if isinstance(precomputed_embeddings, torch.Tensor):
mm_item.precomputed_embeddings = precomputed_embeddings.to(
self.device, non_blocking=True
)
self.multimodal_inputs = multimodal_inputs self.multimodal_inputs = multimodal_inputs
self.token_type_ids = token_type_ids_tensor self.token_type_ids = token_type_ids_tensor
self.seq_lens_sum = sum(seq_lens) self.seq_lens_sum = sum(seq_lens)
@@ -120,6 +120,13 @@ class MultiModalStaticCache(MultimodalCache):
self.current_size += data_size self.current_size += data_size
return True return True
def get_single(self, mm_hash: int) -> Optional[EmbeddingResult]:
"""Get a single cached embedding by its hash (no combine_hashes)."""
embedding = self.mm_cache.get(mm_hash)
if embedding is not None:
self.mm_cache.move_to_end(mm_hash)
return embedding
def has(self, mm_hash: int) -> bool: def has(self, mm_hash: int) -> bool:
return mm_hash in self.mm_cache return mm_hash in self.mm_cache
+1 -3
View File
@@ -270,9 +270,7 @@ class DeepseekVL2ForCausalLM(nn.Module):
for item in items: for item in items:
assert item.feature.dim() == 4 assert item.feature.dim() == 4
image_feature = self.vision.forward_features( image_feature = self.vision.forward_features(
item.feature.type(next(self.vision.parameters()).dtype).to( item.feature.type(next(self.vision.parameters()).dtype)
device=next(self.vision.parameters()).device
)
) )
images_embeds = self.projector(image_feature) images_embeds = self.projector(image_feature)
_, hw, n_dim = images_embeds.shape _, hw, n_dim = images_embeds.shape
+1 -1
View File
@@ -440,7 +440,7 @@ class Phi4MMForCausalLM(nn.Module):
self.embed_tokens_extend( self.embed_tokens_extend(
# item.feature: (num_audios_in_a_sequence, T, D) # item.feature: (num_audios_in_a_sequence, T, D)
# item.audio_attention_mask: (num_audios_in_a_sequence, T, D) BoolTensor or None # item.audio_attention_mask: (num_audios_in_a_sequence, T, D) BoolTensor or None
audio_features=item.feature.to(device).type(dtype), audio_features=item.feature.type(dtype),
audio_attention_mask=( audio_attention_mask=(
item.audio_attention_mask.to(device) item.audio_attention_mask.to(device)
if hasattr(item, "audio_attention_mask") if hasattr(item, "audio_attention_mask")
+1 -95
View File
@@ -15,7 +15,6 @@
"""Inference-only Qwen3-VL model compatible with HuggingFace weights.""" """Inference-only Qwen3-VL model compatible with HuggingFace weights."""
import logging import logging
import math
import re import re
from collections import defaultdict from collections import defaultdict
from functools import lru_cache, partial from functools import lru_cache, partial
@@ -73,7 +72,7 @@ from sglang.srt.models.utils import (
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_int_env_var, is_npu, round_up from sglang.srt.utils import add_prefix, is_npu, round_up
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
_is_npu = is_npu() _is_npu = is_npu()
@@ -1167,10 +1166,6 @@ class Qwen3VLForConditionalGeneration(nn.Module):
assert pixel_values.dim() == 2, pixel_values.dim() assert pixel_values.dim() == 2, pixel_values.dim()
assert image_grid_thw.dim() == 2, image_grid_thw.dim() assert image_grid_thw.dim() == 2, image_grid_thw.dim()
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: if self.use_data_parallel:
return run_dp_sharded_mrope_vision_model( return run_dp_sharded_mrope_vision_model(
self.visual, self.visual,
@@ -1181,100 +1176,11 @@ class Qwen3VLForConditionalGeneration(nn.Module):
else: else:
return self.visual(pixel_values, grid_thw=image_grid_thw) 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: def get_video_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
for item in items:
item.feature = item.feature.to(self.visual.device)
# in qwen-vl, last dim is the same # in qwen-vl, last dim is the same
pixel_values = torch.cat([item.feature for item in items], dim=0).type( pixel_values = torch.cat([item.feature for item in items], dim=0).type(
self.visual.dtype self.visual.dtype
) )
# Memory optimization for item.feature:
# 1. item.feature is released when request finished
# 2. High concurrency may cause device OOM due to delayed release
# 3. Fix: Offload item.feature to CPU, move to device only when needed
for item in items:
item.feature = item.feature.to("cpu")
video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0) video_grid_thw = torch.concat([item.video_grid_thw for item in items], dim=0)
assert pixel_values.dim() == 2, pixel_values.dim() assert pixel_values.dim() == 2, pixel_values.dim()
assert video_grid_thw.dim() == 2, video_grid_thw.dim() assert video_grid_thw.dim() == 2, video_grid_thw.dim()
+1 -1
View File
@@ -484,7 +484,7 @@ class StepVLForConditionalGeneration(nn.Module):
assert len(items) == 1 # We only have images. assert len(items) == 1 # We only have images.
item = items[0] item = items[0]
pixel_values = item.feature.type(self.vision_model.dtype).to(self.device) pixel_values = item.feature.type(self.vision_model.dtype)
num_patches = item.model_specific_data.get("num_patches") num_patches = item.model_specific_data.get("num_patches")
patch_pixel_values = item.model_specific_data.get("patch_pixel_values", None) patch_pixel_values = item.model_specific_data.get("patch_pixel_values", None)
if patch_pixel_values is not None: if patch_pixel_values is not None: