[BugFix][EPD] Fix Mooncake source-MR lifecycle for multi-TP /send (#32071)

This commit is contained in:
Zheng Wengang
2026-07-27 15:06:01 +08:00
committed by GitHub
parent 082b2a10b6
commit 3d3ba4f746
@@ -1813,16 +1813,15 @@ class MMEncoder:
f"(shape={mm_data.shape}, element_size={self._element_size})"
)
# MR was registered once in _run_forward and is shared across all
# sibling-TP /send calls;
mr_already_registered = (
self._forward_results.get(req_id, {}).get("mr_ptr")
== embedding.data_ptr()
)
# Request-level shared MR, registered lazily on the first /send;
# deregistration is deferred to _cleanup_inflight_encode_state.
fwd_state = self._forward_results.setdefault(req_id, {})
mr_already_registered = fwd_state.get("mr_ptr") == embedding.data_ptr()
if not mr_already_registered:
self.engine.register(embedding.data_ptr(), embedding.nbytes)
self._forward_results[req_id]["mr_ptr"] = embedding.data_ptr()
_t_xfer_start = time.monotonic()
await asyncio.to_thread(
xfer_ret = await asyncio.to_thread(
self.engine.transfer_sync,
session_id,
embedding.data_ptr(),
@@ -1834,18 +1833,20 @@ class MMEncoder:
encoder_metrics_collector.observe_transfer(
xfer_ms / 1000.0, backend="mooncake"
)
if not mr_already_registered:
self.engine.deregister(embedding.data_ptr())
# Only emit at INFO when transfer is slow or fell back
# to per-/send register;
if xfer_ret < 0:
raise InternalError(
f"Mooncake transfer_sync failed for {req_id} "
f"(session={session_id}, nbytes={embedding.nbytes}, "
f"ret={xfer_ret})"
)
# Only emit at INFO when transfer is slow or the MR was
# registered lazily by this /send;
if xfer_ms > 200.0 or not mr_already_registered:
logger.info(
f"[{req_id}] mooncake transfer_sync={xfer_ms:.1f}ms "
f"nbytes={embedding.nbytes} shared_mr={mr_already_registered}"
)
mm_data.embedding = None
# Send ack/data
if url is not None:
endpoint = NetworkAddress.parse(url).to_tcp()
@@ -2010,8 +2011,6 @@ class MMEncoder:
task.cancel()
# Also clean up embedding data and forward state
mm_data = self.embedding_to_send.pop(req_id, None)
if mm_data is not None:
mm_data.cached_embedding = None
# Release the rkey after all /send calls have completed.
forward_state = self._forward_results.pop(req_id, None)
if forward_state is not None:
@@ -2023,6 +2022,11 @@ class MMEncoder:
logger.warning(
f"Shared-MR deregister failed for {req_id}: {dereg_err}"
)
forward_state.pop("embedding", None)
# Release the embedding only after the MR is deregistered.
if mm_data is not None:
mm_data.embedding = None
mm_data.cached_embedding = None
self._forward_ready_events.pop(req_id, None)
def _schedule_inflight_encode_cleanup(self, req_id: str):