feat(disagg): wire per-layer overlapped transfer end-to-end (lever A, A3-step3)

Complete the lever-A hot-path integration behind SGLANG_CP_SHARED_KV_PER_LAYER_TRANSFER:
- PerLayerTransferContext: add num_layers + a completion event so finish() waits
  until ALL layers are processed before wait_batch_transfers (never races ahead of
  the worker threads and silently drops in-flight layers); times out to FAILURE.
- PerLayerTransferManager: has_room (for the swap) + drop (abort/failure drain so
  outstanding RDMA finishes before pages are reclaimed).
- MooncakeKVManager.register_per_layer_transfer: build + register a context before
  the forward, reusing send()'s exact CP filter (no re-derivation).
- transfer_worker: when a room is per-layer-active, wait those transfers (finish)
  instead of the monolithic send_kvcache -- no double-send; aux/state/completion
  unchanged. The skip path drops the context on abort/failure.
- prefill scheduler: _register_per_layer_transfers(batch) before run_batch, scoped
  (first impl) to single-forward, no-cached-prefix requests (the notifier transfers
  forward-written pages; chunked/cached-prefix are HiCache-loaded -> lever B).

Unit-tested (25 cases). e2e output-equality + TTFT verification next.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-07 09:51:05 +00:00
co-authored by Claude Opus 4.8
parent aa6acc9485
commit ae18e3adc8
4 changed files with 179 additions and 28 deletions
@@ -391,7 +391,52 @@ class MooncakeKVManager(CommonKVManager):
dst_kv_blocks,
)
return PerLayerTransferContext(self.engine, mooncake_session_id, get_blocks)
return PerLayerTransferContext(
self.engine, mooncake_session_id, get_blocks, num_layers=n_layers
)
def register_per_layer_transfer(self, room, page_indices) -> bool:
"""Lever A: before the forward, build + register a per-layer transfer context
for `room` so the per-layer notifier overlaps its main-KV transfer with the
forward. Reuses send()'s CP filter exactly (no re-derivation). Scoped to
CP-shared-KV (the target scenario). Returns True iff a context was registered;
a False just falls back to the monolithic post-forward transfer (still correct).
page_indices = the request's full new-token logical page ids (req_to_token)."""
mgr = getattr(self, "per_layer_transfer_manager", None)
if mgr is None or not self.server_args.enable_nsa_prefill_cp_shared_kv:
return False
infos = self.transfer_infos.get(room)
if not infos:
return False
from sglang.srt.disaggregation.utils import filter_kv_pages_for_cp_shared_kv
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
layout = CpSharedKVLayout(
page_size=self.kv_args.page_size,
cp_size=self.attn_cp_size,
cp_rank=self.attn_cp_rank,
)
pages = np.asarray(page_indices, dtype=np.int32)
for info in infos.values():
if getattr(info, "is_dummy", False):
continue
owned_pages, positions = filter_kv_pages_for_cp_shared_kv(
layout=layout, logical_pages=pages, chunk_page_start=0
)
dst_indices = np.asarray(info.dst_kv_indices, dtype=np.int32)[positions]
ctx = self.build_per_layer_context(
info.mooncake_session_id, owned_pages, dst_indices
)
if ctx is not None:
mgr.register(room, ctx)
logger.info(
"[CP_PER_LAYER_TRANSFER] registered room=%s owned_pages=%d",
room,
len(owned_pages),
)
return True
return False
return False
def _send_kvcache_generic(
self,
@@ -1090,6 +1135,11 @@ class MooncakeKVManager(CommonKVManager):
"Skipping chunk for room %s because it has already failed or been aborted",
kv_chunk.room,
)
# Lever A: drain + drop any per-layer context for this room so an
# outstanding RDMA completes before its KV pages can be reclaimed.
per_layer_mgr = getattr(self, "per_layer_transfer_manager", None)
if per_layer_mgr is not None:
per_layer_mgr.drop(kv_chunk.room)
continue
reqs_to_be_processed = (
self.transfer_infos[kv_chunk.room].values()
@@ -1171,7 +1221,22 @@ class MooncakeKVManager(CommonKVManager):
target_rank_registration_info: KVArgsRegisterInfo = (
self.decode_kv_args_table[req.mooncake_session_id]
)
if self.is_mla_backend or (
per_layer_mgr = getattr(
self, "per_layer_transfer_manager", None
)
if per_layer_mgr is not None and per_layer_mgr.has_room(
kv_chunk.room
):
# Lever A: the main KV was transferred per-layer, overlapped
# with the forward; wait those transfers here instead of the
# monolithic send (no double-send). aux/state below unchanged.
ret = per_layer_mgr.finish(kv_chunk.room)
logger.info(
"[CP_PER_LAYER_TRANSFER] finished room=%s ret=%s",
kv_chunk.room,
ret,
)
elif self.is_mla_backend or (
self.attn_tp_size
== target_rank_registration_info.dst_attn_tp_size
):