[BugFix][EPD] Early-release mooncake GPU embeddings; fix gpu_id via scheduler.ps (#31591)
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user