Revert "improve: combine vit calls for images from different reqs from one batch (#25910)" (#26442)

Co-authored-by: Cursor Agent <cursoragent@cursor.com>
This commit is contained in:
Bingxu Chen
2026-05-29 11:13:34 -07:00
committed by GitHub
co-authored by Cursor Agent
parent ec075d8bc5
commit f113ece5cc
2 changed files with 192 additions and 173 deletions
+192 -171
View File
@@ -494,15 +494,13 @@ def _can_skip_pre_embed_feature_move(data_embedding_func: DataEmbeddingFunc) ->
def _move_items_to_device(
items: List[MultimodalDataItem], device: torch.device
) -> None:
"""Move item features to the target device (in-place, non-blocking).
Saves a CPU reference so the offload path can restore without GPU->CPU copy."""
"""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._cpu_feature = item.feature
item.feature = item.feature.to(device, non_blocking=True)
def get_chunked_embedding_legacy(
def _get_chunked_embedding_full(
data_embedding_func: DataEmbeddingFunc,
embedding_items_per_req: List[MultimodalDataItem],
items_offset: List[Tuple[int, int]],
@@ -551,85 +549,77 @@ def get_chunked_embedding_legacy(
return embedding_per_req_chunk, input_ids
def find_chunk_items_and_check_cache(
def _get_chunked_embedding_by_item(
data_embedding_func: DataEmbeddingFunc,
embedding_items_per_req: List[MultimodalDataItem],
items_offset: List[Tuple[int, int]],
chunk_start: int,
chunk_end: int,
) -> List[Tuple[MultimodalDataItem, Optional[torch.Tensor], int, int]]:
"""Return (item, cached_embedding_or_None, start, end) for items in [chunk_start, chunk_end)."""
chunk_entries = []
for item, (start, end) in zip(embedding_items_per_req, items_offset):
if end >= chunk_start and start < chunk_end:
cached = embedding_cache.get_single(item.hash)
emb = cached.embedding if cached is not None else None
chunk_entries.append((item, emb, start, end))
return chunk_entries
def assemble_chunk_embedding(
chunk_entries: List[Tuple[Any, torch.Tensor, int, int]],
chunk_start: int,
chunk_end: int,
extend_prefix_len: int,
extend_seq_len: int,
device: torch.device,
) -> Optional[torch.Tensor]:
"""
Assemble a chunk of embeddings by slicing each item's embedding
to the portion that falls within [chunk_start, chunk_end).
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).
"""
chunk_start = extend_prefix_len
chunk_end = extend_prefix_len + extend_seq_len # exclusive
if extend_seq_len <= 0:
return None
# 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))
if not overlapping:
return None
# 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))
# 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]
)
# 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)
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 _, emb, start, end in chunk_entries:
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])
if not chunk_slices:
return None
return torch.cat(chunk_slices, dim=0)
def get_chunked_prefill_embedding_legacy(
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,
max_iterations: int,
) -> tuple[torch.Tensor | None, torch.Tensor]:
"""Non-per-image path: encode each request independently."""
embedding_list = []
device = input_ids.device
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
extend_prefix_len = prefix_length[i]
extend_seq_len = extend_length[i] if i < len(extend_length) else 0
chunk_embedding, input_ids = get_chunked_embedding_legacy(
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)
if len(embedding_list) == 0:
return None, input_ids
return torch.concat(embedding_list, dim=0), input_ids
def _get_chunked_prefill_embedding(
data_embedding_func: DataEmbeddingFunc,
embedding_items: List[MultimodalDataItem],
@@ -640,99 +630,56 @@ def _get_chunked_prefill_embedding(
input_ids: torch.Tensor,
) -> tuple[torch.Tensor | None, torch.Tensor]:
"""
Chunked prefill embedding: collect cache misses across all per-image
requests, batch them into a single ViT call, then assemble per-request
chunk embeddings from the results.
Chunked prefill embedding: encode per-request items and extract the chunk.
Items are already split per-image at processor stage.
"""
embedding_list = []
device = input_ids.device
# FIXME(Xinyuan): temporary workaround for eagle3
# FIXME(yhyang201): check this
max_iterations = min(len(items_size) - 1, len(prefix_length))
per_image_process = (
len(embedding_items) > 0 and len(embedding_items[0].offsets) == 1
)
if not per_image_process:
return get_chunked_prefill_embedding_legacy(
data_embedding_func,
embedding_items,
items_size,
prefix_length,
extend_length,
items_offset_list,
input_ids,
max_iterations,
)
# collect chunk entries per request, accumulate all misses
pending_requests = []
all_miss_items = []
all_miss_token_counts = []
for i in range(max_iterations):
if items_size[i] == items_size[i + 1]:
continue
extend_seq_len = extend_length[i] if i < len(extend_length) else 0
if extend_seq_len <= 0:
continue
extend_prefix_len = prefix_length[i]
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
chunk_start = extend_prefix_len
chunk_end = extend_prefix_len + extend_seq_len
chunk_entries = find_chunk_items_and_check_cache(
embedding_items_per_req,
items_offset,
chunk_start,
chunk_end,
)
if not chunk_entries:
extend_prefix_len = prefix_length[i]
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
for item, emb, start, end in chunk_entries:
if emb is None:
all_miss_items.append(item)
all_miss_token_counts.append(end - start + 1)
# 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)
pending_requests.append((chunk_entries, chunk_start, chunk_end))
miss_embeddings = []
if all_miss_items:
if not _can_skip_pre_embed_feature_move(data_embedding_func):
_move_items_to_device(all_miss_items, device)
# vit_input_tokens = sum(
# item.feature.shape[0] for item in all_miss_items
# if isinstance(item.feature, torch.Tensor)
# )
# logger.info(f"ViT batch: {len(all_miss_items)} items, {vit_input_tokens} input patches, {sum(all_miss_token_counts)} output tokens")
all_miss_embedding = data_embedding_func(all_miss_items)
all_miss_embedding = all_miss_embedding.reshape(
-1, all_miss_embedding.shape[-1]
)
miss_embeddings = list(
torch.split(all_miss_embedding, all_miss_token_counts, dim=0)
)
for item, emb in zip(all_miss_items, miss_embeddings):
embedding_cache.set(item.hash, EmbeddingResult(embedding=emb))
# fill in miss embeddings and assemble per-request chunks
miss_iter = iter(miss_embeddings)
for chunk_entries, chunk_start, chunk_end in pending_requests:
chunk_entries = [
(item, next(miss_iter) if emb is None else emb, start, end)
for item, emb, start, end in chunk_entries
]
chunk_embedding = assemble_chunk_embedding(
chunk_entries, chunk_start, chunk_end
)
if chunk_embedding is not None:
embedding_list.append(chunk_embedding)
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)
if len(embedding_list) == 0:
return None, input_ids
@@ -1072,29 +1019,6 @@ def _embed_mm_inputs_with_split(
return input_embeds, other_info
def offload_mm_features_to_cpu(mm_inputs_list: List[MultimodalInputs]):
"""Free GPU features after embedding. CPU copies are kept for later use
(e.g. chunked prefill or recovery after retraction)."""
language_only = get_global_server_args().language_only
for mm_input in mm_inputs_list or []:
if not mm_input or not hasattr(mm_input, "mm_items"):
continue
for item in mm_input.mm_items:
if isinstance(item.feature, torch.Tensor) and item.feature.is_cuda:
if item._cpu_feature is not None:
item.feature = item._cpu_feature
else:
item.feature = item.feature.to("cpu", non_blocking=True)
if language_only:
pe = item.precomputed_embeddings
if (
isinstance(pe, torch.Tensor)
and pe.is_cuda
and not getattr(item, "_keep_device_embedding", False)
):
item.precomputed_embeddings = pe.to("cpu", non_blocking=True)
def general_mm_embed_routine(
input_ids: torch.Tensor,
forward_batch: ForwardBatch,
@@ -1107,6 +1031,18 @@ def general_mm_embed_routine(
) -> torch.Tensor:
"""
Process multimodal inputs and forward through language model.
Args:
input_ids: Input token IDs tensor
forward_batch: Batch information for model forward pass
language_model: Base language model to use
data_embedding_funcs: A dictionary mapping from modality type to the corresponding embedding function.
placeholder_tokens: Token IDs for multimodal placeholders
use_deepstack: Whether to use deepstack embeddings for each modality, default False
**kwargs: Additional arguments passed to language model
Returns:
Hidden states from language model forward pass
"""
assert hasattr(language_model, "get_input_embeddings")
embed_tokens = language_model.get_input_embeddings()
@@ -1160,9 +1096,34 @@ def general_mm_embed_routine(
# add for qwen3_vl deepstack
if use_deepstack:
kwargs["input_deepstack_embeds"] = other_info["input_deepstack_embeds"]
# Free GPU features after embedding. CPU copies are kept for
# later use (e.g. chunked prefill or recovery after retraction).
offload_mm_features_to_cpu(mm_inputs_list)
# Offload GPU features to CPU instead of discarding them to balance memory
# efficiency and data persistence.
# In chunked-prefill, a request is processed across multiple batches, and
# the original multimodal data must remain accessible until the entire
# prefill phase is complete. Since the multimodal embedding cache is
# best-effort, offloading to CPU ensures we have a reliable fallback
# if a cache miss occurs in subsequent chunks, while still freeing up
# critical GPU memory.
if mm_inputs_list:
for mm_input_obj in mm_inputs_list:
if mm_input_obj and hasattr(mm_input_obj, "mm_items"):
for mm_item in mm_input_obj.mm_items:
feature = getattr(mm_item, "feature", None)
if isinstance(feature, torch.Tensor) and feature.is_cuda:
mm_item.feature = feature.to("cpu", non_blocking=True)
if get_global_server_args().language_only:
precomputed_embeddings = getattr(
mm_item, "precomputed_embeddings", None
)
if (
isinstance(precomputed_embeddings, torch.Tensor)
and precomputed_embeddings.is_cuda
):
mm_item.precomputed_embeddings = (
precomputed_embeddings.to(
"cpu", non_blocking=True
)
)
forward_batch.mm_inputs = None
forward_batch.mm_input_embeds = input_embeds
else:
@@ -1183,6 +1144,66 @@ def general_mm_embed_routine(
return hidden_states
def get_multimodal_data_bounds(
input_ids: torch.Tensor, pad_values: List[int], token_pairs: List[Tuple[int, int]]
) -> torch.Tensor:
"""
Returns a tensor indicating the bounds of multimodal data (images, video, audio, etc.)
Returns:
[bounds_count, 2]
"""
# All the multimodal data in the batch should share the same special bound token ids.
start_tokens = {s for s, _e in token_pairs}
end_tokens = {e for _s, e in token_pairs}
assert all(isinstance(t, int) for t in start_tokens)
assert all(isinstance(t, int) for t in end_tokens)
start_cond = torch.isin(
input_ids, torch.as_tensor(start_tokens, device=input_ids.device)
)
end_cond = torch.isin(
input_ids, torch.as_tensor(end_tokens, device=input_ids.device)
)
(data_start_tokens,) = torch.where(start_cond)
(data_end_tokens,) = torch.where(end_cond)
data_start_tokens_cpu = data_start_tokens.cpu().tolist()
data_end_tokens_cpu = data_end_tokens.cpu().tolist()
# the im_start_id sometimes can be cached as prefix, but it is needed for the embedding of the multimodal data
if len(data_start_tokens_cpu) != len(data_end_tokens_cpu):
if (
len(data_start_tokens_cpu) + 1 == len(data_end_tokens_cpu)
and input_ids[0].item() in pad_values
and data_end_tokens_cpu
and data_start_tokens_cpu
and data_end_tokens_cpu[0] < data_start_tokens_cpu[0]
):
data_start_tokens_cpu.insert(0, 0)
valid_mm_data_nums = min(len(data_start_tokens_cpu), len(data_end_tokens_cpu))
if valid_mm_data_nums == 0:
return torch.zeros((0, 2), device=input_ids.device)
# Filter out pairs where start_token >= end_token
valid_pairs = []
for i in range(valid_mm_data_nums):
start_token = data_start_tokens_cpu[i]
end_token = data_end_tokens_cpu[i]
if start_token < end_token:
valid_pairs.append((start_token + 1, end_token - 1))
if not valid_pairs:
return torch.zeros((0, 2), device=input_ids.device)
# Convert valid pairs to tensor
valid_pairs_tensor = torch.as_tensor(valid_pairs, device=input_ids.device)
return valid_pairs_tensor
def data_hash(data) -> int:
hash_bytes = hashlib.sha256(data).digest()[:8]
return int.from_bytes(hash_bytes, byteorder="big", signed=False)
@@ -254,8 +254,6 @@ class MultimodalDataItem:
# the raw features returned by processor, e.g. pixel_values or audio_features
feature: Union[torch.Tensor, np.ndarray] = None
# CPU reference kept during GPU encoding, used to skip GPU->CPU copy on offload
_cpu_feature: Optional[torch.Tensor] = None
# the precomputed embeddings, passed as final encoder embeddings
# One and only one of the feature and precomputed_embeddings will be empty
precomputed_embeddings: Optional[Union[torch.Tensor, np.ndarray]] = None