Preserve draft NSA state during CP disaggregated transfer

CP shared KV already registered the draft model main KV buffer with the
prefill/decode Mooncake managers, but NSA draft state buffers were not part
of the state registration set. HiCache/cache-hit traffic could then transfer
pages from the draft pool without transferring the matching draft
index/scale state, which is a plausible cause of the EAGLE/MTP accept-length
collapse after cache hits.

This appends compatible draft NSA state buffers to the existing state
transfer registration on both prefill and decode, and extends transfer-side
diagnostics so source/destination state-buffer counts are visible. The
mismatch guard degrades to the common prefix of registered state buffers
instead of crashing if a rolling deployment exposes asymmetric registration.

Constraint: Scope is intentionally limited to target_state_type=nsa and draft_state_type=nsa.
Rejected: Treat draft main KV transfer as sufficient | NSA attention also needs draft index/scale state for transferred pages.
Rejected: Add Mamba/SWA draft-state semantics now | those state layouts need separate correctness analysis.
Confidence: medium
Scope-risk: moderate
Directive: Do not remove the draft_state_buffer_start/count fields without checking Mooncake source/destination registration symmetry.
Tested: PYTHONDONTWRITEBYTECODE=1 python3 -m py_compile python/sglang/srt/disaggregation/prefill.py python/sglang/srt/disaggregation/decode.py python/sglang/srt/disaggregation/mooncake/conn.py
Tested: git diff --check
Tested: Remote prefill log showed registered_state_bufs=79 and maybe_send_extra_state src_state_bufs=79 dst_state_bufs=79 with no state-buffer mismatch.
Not-tested: Full accept-length recovery; latest remote run hit an unrelated prefill KV allocator idle-check leak after transfer registration succeeded.
This commit is contained in:
laoyao0822
2026-05-26 23:59:28 +08:00
committed by leavelet
parent fa80c6278e
commit f2834b3403
4 changed files with 621 additions and 4 deletions
@@ -259,7 +259,9 @@ class MooncakeKVManager(CommonKVManager):
if self.kv_args.kv_data_ptrs and self.kv_args.kv_data_lens:
_cp_draft_shared_kv_debug(
"register_buffers mode=%s cp_rank=%s total_kv_bufs=%s "
"draft_start=%s draft_count=%s kv_lens=%s kv_item_lens=%s",
"draft_start=%s draft_count=%s kv_lens=%s kv_item_lens=%s "
"state_type=%s state_bufs=%s state_lens=%s state_item_lens=%s "
"draft_state_type=%s draft_state_bufs=%s",
self.disaggregation_mode,
self.attn_cp_rank,
len(self.kv_args.kv_data_ptrs),
@@ -267,6 +269,12 @@ class MooncakeKVManager(CommonKVManager):
getattr(self.kv_args, "draft_kv_buffer_count", None),
_np_summary(self.kv_args.kv_data_lens),
_np_summary(self.kv_args.kv_item_lens),
getattr(self.kv_args, "state_type", None),
len(getattr(self.kv_args, "state_data_ptrs", []) or []),
_np_summary(getattr(self.kv_args, "state_data_lens", [])),
_np_summary(getattr(self.kv_args, "state_item_lens", [])),
getattr(self.kv_args, "draft_state_type", None),
getattr(self.kv_args, "draft_state_buffer_count", None),
)
self.engine.batch_register(
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens
@@ -645,6 +653,33 @@ class MooncakeKVManager(CommonKVManager):
):
"""Send state or extra pool data with type-specific handling."""
state_type = getattr(self.kv_args, "state_type", "none")
_cp_draft_shared_kv_debug(
"maybe_send_extra_state cp_rank=%s room=%s session=%s state_type=%s "
"src_state_bufs=%s dst_state_bufs=%s src_state_lens=%s src_state_item_lens=%s "
"prefill_state_indices=%s dst_state_indices=%s draft_state_type=%s "
"draft_state_bufs=%s target_registration_state_bufs=%s",
self.attn_cp_rank,
req.room,
req.mooncake_session_id,
state_type,
len(getattr(self.kv_args, "state_data_ptrs", []) or []),
len(dst_state_data_ptrs or []),
_np_summary(getattr(self.kv_args, "state_data_lens", [])),
_np_summary(getattr(self.kv_args, "state_item_lens", [])),
_np_summary(prefill_state_indices),
_np_summary(
dst_state_indices
if dst_state_indices is not None
else getattr(req, "dst_state_indices", [])
),
getattr(self.kv_args, "draft_state_type", None),
getattr(self.kv_args, "draft_state_buffer_count", None),
(
len(target_rank_registration_info.dst_state_data_ptrs)
if target_rank_registration_info is not None
else None
),
)
if state_type == "mamba":
# Check if we need slice transfer for different TP sizes
@@ -689,13 +724,34 @@ class MooncakeKVManager(CommonKVManager):
prefill_state_indices = prefill_state_indices[
: len(effective_dst_state_indices)
]
src_state_data_ptrs = self.kv_args.state_data_ptrs
dst_state_ptrs = dst_state_data_ptrs
state_item_lens = self.kv_args.state_item_lens
if len(src_state_data_ptrs) != len(dst_state_ptrs):
transfer_buf_count = min(len(src_state_data_ptrs), len(dst_state_ptrs))
logger.warning(
"State buffer count mismatch during PD transfer: src=%s dst=%s "
"state_type=%s draft_state_type=%s draft_state_bufs=%s room=%s "
"session=%s; transferring first %s buffers only",
len(src_state_data_ptrs),
len(dst_state_ptrs),
state_type,
getattr(self.kv_args, "draft_state_type", None),
getattr(self.kv_args, "draft_state_buffer_count", None),
req.room,
req.mooncake_session_id,
transfer_buf_count,
)
src_state_data_ptrs = src_state_data_ptrs[:transfer_buf_count]
dst_state_ptrs = dst_state_ptrs[:transfer_buf_count]
state_item_lens = state_item_lens[:transfer_buf_count]
# Reuse _send_kvcache_generic interface to send extra pool data
prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32)
return self._send_kvcache_generic(
mooncake_session_id=req.mooncake_session_id,
src_data_ptrs=self.kv_args.state_data_ptrs,
dst_data_ptrs=dst_state_data_ptrs,
item_lens=self.kv_args.state_item_lens,
src_data_ptrs=src_state_data_ptrs,
dst_data_ptrs=dst_state_ptrs,
item_lens=state_item_lens,
prefill_data_indices=prefill_state_indices,
dst_data_indices=effective_dst_state_indices,
executor=executor,
@@ -1454,6 +1510,24 @@ class MooncakeKVReceiver(CommonKVReceiver):
dst_tp_rank = str(tp_rank).encode("ascii")
dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii")
dst_kv_item_len = str(kv_item_len).encode("ascii")
_cp_draft_shared_kv_debug(
"decode_register_kv_args cp_rank=%s room=%s session=%s "
"kv_bufs=%s aux_bufs=%s state_type=%s state_bufs=%s "
"state_item_lens=%s draft_start=%s draft_count=%s "
"draft_state_type=%s draft_state_bufs=%s",
self.kv_mgr.attn_cp_rank,
self.bootstrap_room,
self.session_id,
len(self.kv_mgr.kv_args.kv_data_ptrs),
len(self.kv_mgr.kv_args.aux_data_ptrs),
getattr(self.kv_mgr.kv_args, "state_type", None),
len(self.kv_mgr.kv_args.state_data_ptrs),
_np_summary(self.kv_mgr.kv_args.state_item_lens),
getattr(self.kv_mgr.kv_args, "draft_kv_buffer_start", None),
getattr(self.kv_mgr.kv_args, "draft_kv_buffer_count", None),
getattr(self.kv_mgr.kv_args, "draft_state_type", None),
getattr(self.kv_mgr.kv_args, "draft_state_buffer_count", None),
)
sock, lock = self._connect_to_bootstrap_server(bootstrap_info)
with lock: