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:
@@ -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
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user