From fa6f4dfb353c8e0ccf1488b77ebee4c07f32f0d0 Mon Sep 17 00:00:00 2001 From: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Date: Fri, 22 May 2026 14:57:41 +0800 Subject: [PATCH] improve: combine vit calls for images from different reqs from one batch (#25910) Co-authored-by: Yaochen Han <814073252@qq.com> --- python/sglang/srt/managers/mm_utils.py | 358 +++++++++---------- python/sglang/srt/managers/schedule_batch.py | 2 + 2 files changed, 168 insertions(+), 192 deletions(-) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 1c7d3afe1..50a641ac6 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -462,13 +462,15 @@ DataEmbeddingFunc = Callable[ def _move_items_to_device( items: List[MultimodalDataItem], device: torch.device ) -> None: - """Move item features to the target device (in-place, non-blocking).""" + """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.""" 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_full( +def get_chunked_embedding_legacy( data_embedding_func: DataEmbeddingFunc, embedding_items_per_req: List[MultimodalDataItem], items_offset: List[Tuple[int, int]], @@ -516,77 +518,85 @@ def _get_chunked_embedding_full( return embedding_per_req_chunk, input_ids -def _get_chunked_embedding_by_item( - data_embedding_func: DataEmbeddingFunc, +def find_chunk_items_and_check_cache( embedding_items_per_req: List[MultimodalDataItem], items_offset: List[Tuple[int, int]], - extend_prefix_len: int, - extend_seq_len: int, - device: torch.device, + 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, ) -> Optional[torch.Tensor]: """ - 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). + Assemble a chunk of embeddings by slicing each item's embedding + to the portion that falls within [chunk_start, chunk_end). """ - 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 idx, _, start, end in overlapping: - emb = cached_embeddings[idx] # shape: (end - start + 1, hidden) + for _, emb, start, end in chunk_entries: 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], @@ -597,56 +607,98 @@ def _get_chunked_prefill_embedding( input_ids: torch.Tensor, ) -> tuple[torch.Tensor | None, torch.Tensor]: """ - Chunked prefill embedding: encode per-request items and extract the chunk. - Items are already split per-image at processor stage. + 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. """ 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 - 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): + 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: 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) + 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) - 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) + pending_requests.append((chunk_entries, chunk_start, chunk_end)) + + miss_embeddings = [] + if all_miss_items: + _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 len(embedding_list) == 0: return None, input_ids @@ -987,6 +1039,25 @@ 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: + item.precomputed_embeddings = pe.to("cpu", non_blocking=True) + + def general_mm_embed_routine( input_ids: torch.Tensor, forward_batch: ForwardBatch, @@ -999,18 +1070,6 @@ 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() @@ -1064,34 +1123,9 @@ def general_mm_embed_routine( # add for qwen3_vl deepstack if use_deepstack: kwargs["input_deepstack_embeds"] = other_info["input_deepstack_embeds"] - # 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 - ) - ) + # 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) forward_batch.mm_inputs = None forward_batch.mm_input_embeds = input_embeds else: @@ -1112,66 +1146,6 @@ 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) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 70886301f..9d266a221 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -252,6 +252,8 @@ 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