[BugFix][EPD] Early-release mooncake GPU embeddings; fix gpu_id via scheduler.ps (#31591)

This commit is contained in:
Zheng Wengang
2026-07-30 17:36:55 +08:00
committed by GitHub
parent 6ab3231b97
commit 4ba7d5ad93
2 changed files with 18 additions and 9 deletions
@@ -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)
@@ -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)