From eb349efb14606d78c29bafd2cfe91764e3d65fed Mon Sep 17 00:00:00 2001 From: Yuang Chen <77919385+cccccya@users.noreply.github.com> Date: Mon, 15 Jun 2026 19:47:52 +0800 Subject: [PATCH] [EPD][BugFix] Fix encode_with_global_cache_mooncake (#28031) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 晟海 Co-authored-by: Michael Qiu --- .../srt/disaggregation/encode_receiver.py | 2 +- .../srt/disaggregation/encode_server.py | 92 ++++++++++++++----- 2 files changed, 72 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index b409638f6..bbe120c9a 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -1182,7 +1182,7 @@ class WaitingImageRDMARequest(WaitingImageRequest): self.embeddings_buffer = None self._buffer_from_pool = False self.recv_req.mm_inputs = mm_inputs - self.recv_req.input_ids = mm_inputs.input_ids + self.recv_req.input_ids = array("q", mm_inputs.input_ids) self.status = WaitingImageRequestStatus.SUCCESS self._cleanup_gpu_buffer() self.recv_socket.close() diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 89c181c12..c97b7257b 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -1077,6 +1077,7 @@ class MMEncoder: ) else: mm_hashes = hashes + str_mm_hashes = [str(h) for h in mm_hashes] # All ranks: launch background task for cache check + VIT forward. # Do NOT use run_in_executor: get_feature_fn relies on a session @@ -1088,7 +1089,7 @@ class MMEncoder: # Step 1: Rank 0 checks cache, broadcast mask if TP > 1 if self.rank == 0: exist_mask = await self.mm_global_cache.batch_is_exist( - mm_hashes + str_mm_hashes ) mask_tensor = torch.tensor( [1 if e else 0 for e in exist_mask], @@ -1122,14 +1123,13 @@ class MMEncoder: grid_thw, keep_on_gpu=True, ) - if self.rank == 0: - for i, idx in enumerate(missing_indices): - final_slices[idx] = new_slices[i] - # Step 3: Rank 0 prefetches cache-hit embeddings - prefetch_status = torch.tensor([1], dtype=torch.int32) + # Step 3: Rank 0 prefetches cache-hit embeddings and builds fallback_mask. + fallback_mask = torch.zeros(num_items, dtype=torch.int32) + cached_slices = [] + if self.rank == 0 and hit_indices: - hit_hashes = [mm_hashes[i] for i in hit_indices] + hit_hashes = [str_mm_hashes[i] for i in hit_indices] hit_tokens = [ self.get_num_tokens(grid_thw[i], modality) for i in hit_indices @@ -1146,46 +1146,83 @@ class MMEncoder: await asyncio.sleep(0.005) await asyncio.wait_for(_wait_prefetch(), timeout=60.0) + cached_slices = self.mm_global_cache.get_embeddings( hit_hashes ) for i, idx in enumerate(hit_indices): - final_slices[idx] = cached_slices[i] + if cached_slices[i] is None: + fallback_mask[idx] = 1 + num_partial_fail = int(fallback_mask.sum().item()) + if num_partial_fail > 0: + logger.warning( + f"Req {req_id}: {num_partial_fail}/{len(hit_indices)} " + f"cache-hit items failed to load (pool full), " + f"falling back to ViT" + ) except (asyncio.TimeoutError, Exception) as e: logger.error( f"Prefetch failed for {req_id}: {e}. " f"Falling back to ViT for " f"{len(hit_indices)} hit items." ) - prefetch_status[0] = 0 + for idx in hit_indices: + fallback_mask[idx] = 1 - # Broadcast prefetch result if TP > 1 + # Step 4: Broadcast fallback_mask to all ranks so they stay in sync. if self.server_args.tp_size > 1: torch.distributed.broadcast( - prefetch_status, + fallback_mask, src=0, group=self.mm_global_cache.prefetch_tp_group, ) - # Step 4: All ranks fallback VIT for failed prefetch - # (runs in event loop to preserve session context) + # Step 5: All ranks run ViT for items that need fallback recomputation. + fallback_indices = [ + i for i in range(num_items) if fallback_mask[i].item() == 1 + ] fallback_slices = None - if prefetch_status.item() == 0 and hit_indices: + if fallback_indices: + logger.info( + f"Req {req_id}: All ranks running ViT fallback " + f"for {len(fallback_indices)} items." + ) fallback_slices = self._encode_missing( mm_feature, mm_inputs, - hit_indices, + fallback_indices, modality, get_feature_fn, grid_thw, keep_on_gpu=True, ) - if self.rank == 0: + + # Step 6: Rank 0 assembles final embedding. + if self.rank == 0: + for i, idx in enumerate(missing_indices): + final_slices[idx] = new_slices[i] + + # Fill in successfully loaded cache-hit embeddings + if cached_slices: for i, idx in enumerate(hit_indices): + if cached_slices[i] is not None: + final_slices[idx] = cached_slices[i] + + # Fill in ViT fallback results for failed items + if fallback_slices is not None: + for i, idx in enumerate(fallback_indices): final_slices[idx] = fallback_slices[i] - # Step 5: Rank 0 assembles and stores result - if self.rank == 0: + # Move cached CPU slices to GPU and match model dtype + device = torch.device(f"cuda:{self.gpu_id}") + final_slices = [ + ( + s.to(device=device, dtype=self._embedding_dtype) + if s.device.type == "cpu" + else s + ) + for s in final_slices + ] mm_embedding = torch.cat(final_slices, dim=0) # Wait for any pending VIT / cat kernels to finish # before publishing to /send: mooncake transfer_sync @@ -1193,11 +1230,24 @@ class MMEncoder: # and would otherwise race with in-flight kernels. torch.cuda.current_stream(mm_embedding.device).synchronize() - # Background insert new embeddings into cache - all_new_hashes = [mm_hashes[i] for i in missing_indices] + # Release cache refs after data is copied to GPU + if cached_slices: + loaded_hashes = [ + str_mm_hashes[idx] + for idx in hit_indices + if fallback_mask[idx].item() == 0 + ] + if loaded_hashes: + self.mm_global_cache.release_embeddings(loaded_hashes) + + # Background insert: store newly computed embeddings into global cache. + # Includes both original misses and fallback-recomputed hits. + all_new_hashes = [str_mm_hashes[i] for i in missing_indices] all_new_slices = list(new_slices) if fallback_slices is not None: - all_new_hashes += [mm_hashes[i] for i in hit_indices] + all_new_hashes += [ + str_mm_hashes[i] for i in fallback_indices + ] all_new_slices += list(fallback_slices) if all_new_hashes: