[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,
|
"modality": modality.name,
|
||||||
"prefill_host": self.host_name,
|
"prefill_host": self.host_name,
|
||||||
"embedding_port": self.embedding_port,
|
"embedding_port": self.embedding_port,
|
||||||
|
# Echoed via /send so encoder can release GPU embedding early.
|
||||||
|
"receive_count": self.receive_count,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
cum_idx += 1
|
cum_idx += 1
|
||||||
@@ -1528,6 +1530,7 @@ class MMReceiverBase(ABC):
|
|||||||
self.hostname = get_local_ip_auto()
|
self.hostname = get_local_ip_auto()
|
||||||
self.waiting_list: List[WaitingImageRequest] = []
|
self.waiting_list: List[WaitingImageRequest] = []
|
||||||
self.scheduler = scheduler
|
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.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
|
||||||
|
|
||||||
self.model_type = (
|
self.model_type = (
|
||||||
@@ -1554,10 +1557,9 @@ class MMReceiverBase(ABC):
|
|||||||
self.embedding_pool = None
|
self.embedding_pool = None
|
||||||
pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get()
|
pool_mb = envs.SGLANG_EMBEDDING_POOL_SIZE_MB.get()
|
||||||
if pool_mb and pool_mb > 0 and scheduler is not None:
|
if pool_mb and pool_mb > 0 and scheduler is not None:
|
||||||
gpu_id = getattr(scheduler, "gpu_id", 0)
|
|
||||||
try:
|
try:
|
||||||
self.embedding_pool = MooncakeEmbeddingPool(
|
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:
|
except Exception:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
@@ -2015,8 +2017,7 @@ class MMReceiverBase(ABC):
|
|||||||
f"Pre-allocating GPU buffer for mooncake RDMA: "
|
f"Pre-allocating GPU buffer for mooncake RDMA: "
|
||||||
f"req_id={req_id}, size={total_bytes} bytes"
|
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=self.gpu_id)
|
||||||
embeddings = torch.empty(total_bytes, dtype=torch.uint8, device=gpu_id)
|
|
||||||
self.embeddings_engine.register(
|
self.embeddings_engine.register(
|
||||||
embeddings.data_ptr(),
|
embeddings.data_ptr(),
|
||||||
embeddings.nbytes,
|
embeddings.nbytes,
|
||||||
@@ -2136,13 +2137,12 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
# For zmq_to_scheduler and mooncake
|
# For zmq_to_scheduler and mooncake
|
||||||
def process_waiting_requests(self, recv_reqs):
|
def process_waiting_requests(self, recv_reqs):
|
||||||
if self.encoder_transfer_backend == "mooncake":
|
if self.encoder_transfer_backend == "mooncake":
|
||||||
gpu_id = getattr(self.scheduler, "gpu_id", 0)
|
|
||||||
return self._process_waiting_requests(
|
return self._process_waiting_requests(
|
||||||
recv_reqs,
|
recv_reqs,
|
||||||
WaitingImageRDMARequest,
|
WaitingImageRDMARequest,
|
||||||
embeddings_engine=self.embeddings_engine,
|
embeddings_engine=self.embeddings_engine,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
gpu_id=gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
embedding_pool=self.embedding_pool,
|
embedding_pool=self.embedding_pool,
|
||||||
)
|
)
|
||||||
return self._process_waiting_requests(recv_reqs, WaitingImageRequest)
|
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()
|
rid_to_err_msg: Dict[str, str] = dict()
|
||||||
cond_dict_lock = asyncio.Lock()
|
cond_dict_lock = asyncio.Lock()
|
||||||
rid_to_cond: Dict[str, asyncio.Condition] = {}
|
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()
|
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):
|
async def _cleanup_inflight_encode_state(self, req_id: str):
|
||||||
if not hasattr(self, "_inflight_encode_events"):
|
if not hasattr(self, "_inflight_encode_events"):
|
||||||
return
|
return
|
||||||
|
mooncake_send_done_count.pop(req_id, None)
|
||||||
async with self._inflight_encode_lock:
|
async with self._inflight_encode_lock:
|
||||||
self._inflight_encode_events.pop(req_id, None)
|
self._inflight_encode_events.pop(req_id, None)
|
||||||
self._inflight_encode_meta.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"],
|
buffer_address=request["buffer_address"],
|
||||||
)
|
)
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
# Don't pop embedding_to_send here — other decoder TP ranks may still
|
# Keep embedding until all ranks have /send'd; release early when receive_count is met.
|
||||||
# need it for their /send calls. Cleanup is handled by the scheduled
|
expected_sends = request.get("receive_count")
|
||||||
# timeout task or _cleanup_inflight_encode_state.
|
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)
|
return ORJSONResponse(content=None)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user