From 4ba7d5ad93a6c5efa1c44b34f53084cd81062dc7 Mon Sep 17 00:00:00 2001 From: Zheng Wengang Date: Thu, 30 Jul 2026 02:36:55 -0700 Subject: [PATCH] [BugFix][EPD] Early-release mooncake GPU embeddings; fix gpu_id via scheduler.ps (#31591) --- .../sglang/srt/disaggregation/encode_receiver.py | 12 ++++++------ python/sglang/srt/disaggregation/encode_server.py | 15 ++++++++++++--- 2 files changed, 18 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/disaggregation/encode_receiver.py b/python/sglang/srt/disaggregation/encode_receiver.py index ee45b415e..861149645 100644 --- a/python/sglang/srt/disaggregation/encode_receiver.py +++ b/python/sglang/srt/disaggregation/encode_receiver.py @@ -1053,6 +1053,8 @@ class WaitingImageRDMARequest(WaitingImageRequest): "modality": modality.name, "prefill_host": self.host_name, "embedding_port": self.embedding_port, + # Echoed via /send so encoder can release GPU embedding early. + "receive_count": self.receive_count, } ) cum_idx += 1 @@ -1528,6 +1530,7 @@ class MMReceiverBase(ABC): self.hostname = get_local_ip_auto() self.waiting_list: List[WaitingImageRequest] = [] self.scheduler = scheduler + self.gpu_id = scheduler.ps.gpu_id if scheduler is not None else 0 self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get() self.model_type = ( @@ -1554,10 +1557,9 @@ class MMReceiverBase(ABC): self.embedding_pool = None pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get() if pool_mb and pool_mb > 0 and scheduler is not None: - gpu_id = getattr(scheduler, "gpu_id", 0) try: self.embedding_pool = MooncakeEmbeddingPool( - self.embeddings_engine, gpu_id, pool_mb * 1024 * 1024 + self.embeddings_engine, self.gpu_id, pool_mb * 1024 * 1024 ) except Exception: logger.exception( @@ -2015,8 +2017,7 @@ class MMReceiverBase(ABC): f"Pre-allocating GPU buffer for mooncake RDMA: " f"req_id={req_id}, size={total_bytes} bytes" ) - gpu_id = getattr(self.scheduler, "gpu_id", 0) - embeddings = torch.empty(total_bytes, dtype=torch.uint8, device=gpu_id) + embeddings = torch.empty(total_bytes, dtype=torch.uint8, device=self.gpu_id) self.embeddings_engine.register( embeddings.data_ptr(), embeddings.nbytes, @@ -2136,13 +2137,12 @@ class MMReceiverHTTP(MMReceiverBase): # For zmq_to_scheduler and mooncake def process_waiting_requests(self, recv_reqs): if self.encoder_transfer_backend == "mooncake": - gpu_id = getattr(self.scheduler, "gpu_id", 0) return self._process_waiting_requests( recv_reqs, WaitingImageRDMARequest, embeddings_engine=self.embeddings_engine, dtype=self.dtype, - gpu_id=gpu_id, + gpu_id=self.gpu_id, embedding_pool=self.embedding_pool, ) return self._process_waiting_requests(recv_reqs, WaitingImageRequest) diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index 7dca71334..fb30b6655 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -108,6 +108,8 @@ rid_to_receive_count: Dict[str, int] = dict() rid_to_err_msg: Dict[str, str] = dict() cond_dict_lock = asyncio.Lock() rid_to_cond: Dict[str, asyncio.Condition] = {} +# mooncake: /send completions per part; release GPU embedding once receive_count reached. +mooncake_send_done_count: Dict[str, int] = dict() use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get() @@ -2024,6 +2026,7 @@ class MMEncoder: async def _cleanup_inflight_encode_state(self, req_id: str): if not hasattr(self, "_inflight_encode_events"): return + mooncake_send_done_count.pop(req_id, None) async with self._inflight_encode_lock: self._inflight_encode_events.pop(req_id, None) self._inflight_encode_meta.pop(req_id, None) @@ -4098,9 +4101,15 @@ async def handle_send_request(request: dict): buffer_address=request["buffer_address"], ) req_id = request["req_id"] - # Don't pop embedding_to_send here — other decoder TP ranks may still - # need it for their /send calls. Cleanup is handled by the scheduled - # timeout task or _cleanup_inflight_encode_state. + # Keep embedding until all ranks have /send'd; release early when receive_count is met. + expected_sends = request.get("receive_count") + if expected_sends: + done = mooncake_send_done_count.get(req_id, 0) + 1 + if done >= expected_sends: + mooncake_send_done_count.pop(req_id, None) + await encoder._cleanup_inflight_encode_state(req_id) + else: + mooncake_send_done_count[req_id] = done return ORJSONResponse(content=None)