[VLM] Chunk-aware ViT encoding with per-image cache and lazy device transfer (#22038)
This commit is contained in:
@@ -44,9 +44,6 @@ TensorTransportMode = Literal["cuda_ipc", "auto", "default"]
|
||||
_GPU_FEATURE_BUFFER: Optional[torch.Tensor] = None
|
||||
_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
|
||||
|
||||
|
||||
@@ -462,67 +459,134 @@ DataEmbeddingFunc = Callable[
|
||||
]
|
||||
|
||||
|
||||
def get_embedding_items_per_chunk_with_extra_padding(
|
||||
embedding_items_per_req: List["MultimodalDataItem"],
|
||||
def _move_items_to_device(
|
||||
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_seq_len: int,
|
||||
input_ids: torch.Tensor,
|
||||
device: torch.device,
|
||||
) -> Tuple[Optional[torch.Tensor], torch.Tensor]:
|
||||
"""
|
||||
Fallback: encode all items at once, cache combined result, extract chunk.
|
||||
Used for non-bundled items or EVS results.
|
||||
"""
|
||||
item_hashes = [item.hash for item in embedding_items_per_req]
|
||||
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
|
||||
embedding_per_req = embedding_cache.get(item_hashes)
|
||||
|
||||
if embedding_per_req is None:
|
||||
_move_items_to_device(embedding_items_per_req, device)
|
||||
embedding = data_embedding_func(embedding_items_per_req)
|
||||
embedding_per_req = (
|
||||
EmbeddingResult(embedding=embedding)
|
||||
if isinstance(embedding, torch.Tensor)
|
||||
else embedding
|
||||
)
|
||||
embedding_cache.set(embedding_items_hash, embedding_per_req)
|
||||
|
||||
if isinstance(embedding_per_req, EVSEmbeddingResult):
|
||||
item = embedding_items_per_req[0]
|
||||
input_ids, items_offset = (
|
||||
embedding_per_req.redistribute_pruned_frames_placeholders(
|
||||
input_ids,
|
||||
items_offset,
|
||||
item=item,
|
||||
extend_prefix_len=extend_prefix_len,
|
||||
extend_seq_len=extend_seq_len,
|
||||
)
|
||||
)
|
||||
|
||||
embedding_per_req_chunk, _, _ = get_embedding_chunk(
|
||||
embedding=embedding_per_req.embedding,
|
||||
extend_prefix_len=extend_prefix_len,
|
||||
extend_seq_len=extend_seq_len,
|
||||
items_offset=items_offset,
|
||||
)
|
||||
return embedding_per_req_chunk, input_ids
|
||||
|
||||
|
||||
def _get_chunked_embedding_by_item(
|
||||
data_embedding_func: DataEmbeddingFunc,
|
||||
embedding_items_per_req: List[MultimodalDataItem],
|
||||
items_offset: List[Tuple[int, int]],
|
||||
) -> List["MultimodalDataItem"]:
|
||||
extend_prefix_len: int,
|
||||
extend_seq_len: int,
|
||||
device: torch.device,
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
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)
|
||||
Per-image chunk-aware encoding: only encode images overlapping with the
|
||||
current chunk, cache each image individually.
|
||||
Items must already be split per-image (each item has exactly one offset).
|
||||
"""
|
||||
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"
|
||||
chunk_start = extend_prefix_len
|
||||
chunk_end = extend_prefix_len + extend_seq_len # exclusive
|
||||
|
||||
if extend_seq_len <= 0:
|
||||
return []
|
||||
return None
|
||||
|
||||
# Current chunk's token range
|
||||
chunk_start = extend_prefix_len
|
||||
chunk_end = extend_prefix_len + extend_seq_len
|
||||
# 1. Find items overlapping with current chunk
|
||||
# offsets are (start, end) inclusive on both ends
|
||||
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))
|
||||
|
||||
# Current chunk's token range with extra padding
|
||||
window_start = max(0, chunk_start - _EXTRA_PRE_TOKENS)
|
||||
window_end = chunk_end + _EXTRA_POST_TOKENS
|
||||
if not overlapping:
|
||||
return None
|
||||
|
||||
selected_items: List["MultimodalDataItem"] = []
|
||||
# 2. Check per-image cache for each overlapping item
|
||||
cached_embeddings = {} # idx -> tensor
|
||||
miss_items = [] # (idx, item, start, end)
|
||||
for idx, item, start, end in overlapping:
|
||||
cached = embedding_cache.get_single(item.hash)
|
||||
if cached is not None:
|
||||
cached_embeddings[idx] = cached.embedding
|
||||
else:
|
||||
miss_items.append((idx, item, start, end))
|
||||
|
||||
for item, (start, end) in zip(embedding_items_per_req, items_offset):
|
||||
if start >= end:
|
||||
continue
|
||||
# 3. Batch encode all cache-miss items in one ViT call
|
||||
if miss_items:
|
||||
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]
|
||||
)
|
||||
|
||||
# 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)
|
||||
# Split output by per-item token count
|
||||
token_counts = [end - start + 1 for _, _, start, end in miss_items]
|
||||
split_embeddings = torch.split(all_miss_embedding, token_counts, dim=0)
|
||||
|
||||
return selected_items
|
||||
for (idx, item, _, _), emb in zip(miss_items, split_embeddings):
|
||||
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)
|
||||
|
||||
|
||||
# TODO: To be obsoleted.
|
||||
def _get_chunked_prefill_embedding(
|
||||
data_embedding_func: DataEmbeddingFunc,
|
||||
embedding_items: List[MultimodalDataItem],
|
||||
@@ -532,266 +596,63 @@ def _get_chunked_prefill_embedding(
|
||||
items_offset_list: List[List[Tuple[int, int]]],
|
||||
input_ids: torch.Tensor,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor]:
|
||||
# Calculate embedding for each request, try to get it from cache to avoid repeated calculation
|
||||
"""
|
||||
Chunked prefill embedding: encode per-request items and extract the chunk.
|
||||
Items are already split per-image at processor stage.
|
||||
"""
|
||||
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))
|
||||
|
||||
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]
|
||||
embedding_items_hash = MultiModalStaticCache.combine_hashes(item_hashes)
|
||||
embedding_per_req = embedding_cache.get(item_hashes)
|
||||
if embedding_per_req is None:
|
||||
embedding = data_embedding_func(embedding_items_per_req)
|
||||
embedding_per_req = (
|
||||
EmbeddingResult(embedding=embedding)
|
||||
if isinstance(embedding, torch.Tensor)
|
||||
else embedding
|
||||
)
|
||||
if not 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):
|
||||
item = embedding_items_per_req[0]
|
||||
input_ids, items_offset = (
|
||||
embedding_per_req.redistribute_pruned_frames_placeholders(
|
||||
input_ids,
|
||||
items_offset,
|
||||
item=item,
|
||||
extend_prefix_len=extend_prefix_len,
|
||||
extend_seq_len=extend_seq_len,
|
||||
)
|
||||
# Skip if all items already prefilled
|
||||
if all(offset_end < prefix_length[i] for _, offset_end in items_offset):
|
||||
continue
|
||||
|
||||
# Use per-image path when all items have exactly one offset (already
|
||||
# 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,
|
||||
items_offset,
|
||||
extend_prefix_len,
|
||||
extend_seq_len,
|
||||
device,
|
||||
)
|
||||
if chunk_embedding is not None:
|
||||
embedding_list.append(chunk_embedding)
|
||||
else:
|
||||
chunk_embedding, input_ids = _get_chunked_embedding_full(
|
||||
data_embedding_func,
|
||||
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)
|
||||
|
||||
embedding_per_req_chunk, _, _ = get_embedding_chunk(
|
||||
embedding=embedding_per_req.embedding,
|
||||
extend_prefix_len=extend_prefix_len,
|
||||
extend_seq_len=extend_seq_len,
|
||||
items_offset=items_offset,
|
||||
)
|
||||
embedding_list.append(embedding_per_req_chunk)
|
||||
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(
|
||||
embedding: torch.Tensor,
|
||||
extend_prefix_len: int,
|
||||
extend_seq_len: int,
|
||||
items_offset: List[Tuple[int, int]],
|
||||
) -> Tuple[Optional[torch.Tensor], int, int]:
|
||||
"""
|
||||
From the embedding computed on "items related to this chunk + extra padding",
|
||||
trim out the token embeddings that are not needed for the current chunk, and
|
||||
keep only those mm tokens covered by
|
||||
[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_end = extend_prefix_len + extend_seq_len
|
||||
|
||||
if extend_seq_len <= 0 or chunk_start >= chunk_end:
|
||||
return None, 0, 0
|
||||
|
||||
# The window with extra padding
|
||||
window_start = max(0, chunk_start - _EXTRA_PRE_TOKENS)
|
||||
window_end = chunk_end + _EXTRA_POST_TOKENS
|
||||
|
||||
# Iterate item_offset to choose item.
|
||||
# We need to forward an embedding_idx to locate the item start-end position in embedding.
|
||||
embedding_idx = 0
|
||||
kept_slices: List[torch.Tensor] = []
|
||||
|
||||
num_tokens_before = 0
|
||||
num_tokens_after = 0
|
||||
|
||||
for start, end in items_offset:
|
||||
if start >= end:
|
||||
continue
|
||||
|
||||
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:
|
||||
# Although item is chosen into embedding_items_per_chunk as extra padding,
|
||||
# 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
|
||||
# counts seg_len in embedding, so embedding_idx has to forward.
|
||||
embedding_idx += seg_len
|
||||
|
||||
if not kept_slices:
|
||||
# No mm tokens in this chunk
|
||||
return None, num_tokens_before, num_tokens_after
|
||||
|
||||
trimmed_embedding = torch.cat(kept_slices, dim=0)
|
||||
return trimmed_embedding, num_tokens_before, num_tokens_after
|
||||
|
||||
|
||||
# This function is for chunked prefill vit for multiple items in the next feature.
|
||||
def _get_chunked_prefill_embedding_for_chunked_items(
|
||||
data_embedding_func: Callable[[List["MultimodalDataItem"]], torch.Tensor],
|
||||
embedding_items: List["MultimodalDataItem"],
|
||||
items_size: List[int],
|
||||
prefix_length: List[int],
|
||||
extend_length: List[int],
|
||||
items_offset_list: List[List[Tuple[int, int]]],
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
Multi-modal embedding computation for chunked prefill.
|
||||
|
||||
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 = []
|
||||
# FIXME(Xinyuan): temporary workaround for eagle3, which may have len(items_size) > len(prefix_length)
|
||||
max_iterations = min(len(items_size) - 1, len(prefix_length))
|
||||
|
||||
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
|
||||
|
||||
# 1) Pick up items related with this chunk
|
||||
embedding_items_per_chunk = get_embedding_items_per_chunk_with_extra_padding(
|
||||
embedding_items_per_req,
|
||||
extend_prefix_len=prefix_length[i],
|
||||
extend_seq_len=extend_length[i] if i < len(extend_length) else 0,
|
||||
items_offset=items_offset,
|
||||
)
|
||||
|
||||
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."
|
||||
)
|
||||
else:
|
||||
target_device = embedding_items_per_req[0].feature.device
|
||||
if embedding_per_chunk.device != target_device:
|
||||
embedding_per_chunk = embedding_per_chunk.to(target_device)
|
||||
|
||||
# 3) remove extra padding from embedding_per_chunk, only keep current chunk part
|
||||
# We probably don't need this part.
|
||||
# embedding_per_req_chunk, _, _ = get_embedding_chunk_remove_extra_padding(
|
||||
# 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(
|
||||
input_ids: torch.Tensor, placeholder_tensor: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
|
||||
@@ -1730,21 +1730,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
if input_embeds
|
||||
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.token_type_ids = token_type_ids_tensor
|
||||
self.seq_lens_sum = sum(seq_lens)
|
||||
|
||||
@@ -120,6 +120,13 @@ class MultiModalStaticCache(MultimodalCache):
|
||||
self.current_size += data_size
|
||||
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:
|
||||
return mm_hash in self.mm_cache
|
||||
|
||||
|
||||
@@ -270,9 +270,7 @@ class DeepseekVL2ForCausalLM(nn.Module):
|
||||
for item in items:
|
||||
assert item.feature.dim() == 4
|
||||
image_feature = self.vision.forward_features(
|
||||
item.feature.type(next(self.vision.parameters()).dtype).to(
|
||||
device=next(self.vision.parameters()).device
|
||||
)
|
||||
item.feature.type(next(self.vision.parameters()).dtype)
|
||||
)
|
||||
images_embeds = self.projector(image_feature)
|
||||
_, hw, n_dim = images_embeds.shape
|
||||
|
||||
@@ -440,7 +440,7 @@ class Phi4MMForCausalLM(nn.Module):
|
||||
self.embed_tokens_extend(
|
||||
# item.feature: (num_audios_in_a_sequence, T, D)
|
||||
# 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=(
|
||||
item.audio_attention_mask.to(device)
|
||||
if hasattr(item, "audio_attention_mask")
|
||||
|
||||
@@ -15,7 +15,6 @@
|
||||
"""Inference-only Qwen3-VL model compatible with HuggingFace weights."""
|
||||
|
||||
import logging
|
||||
import math
|
||||
import re
|
||||
from collections import defaultdict
|
||||
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.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||
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
|
||||
|
||||
_is_npu = is_npu()
|
||||
@@ -1167,114 +1166,21 @@ class Qwen3VLForConditionalGeneration(nn.Module):
|
||||
assert pixel_values.dim() == 2, pixel_values.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:
|
||||
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)
|
||||
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)
|
||||
|
||||
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
|
||||
pixel_values = torch.cat([item.feature for item in items], dim=0).type(
|
||||
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)
|
||||
assert pixel_values.dim() == 2, pixel_values.dim()
|
||||
assert video_grid_thw.dim() == 2, video_grid_thw.dim()
|
||||
|
||||
@@ -484,7 +484,7 @@ class StepVLForConditionalGeneration(nn.Module):
|
||||
assert len(items) == 1 # We only have images.
|
||||
|
||||
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")
|
||||
patch_pixel_values = item.model_specific_data.get("patch_pixel_values", None)
|
||||
if patch_pixel_values is not None:
|
||||
|
||||
Reference in New Issue
Block a user