Reduce prefill EAGLE memory pressure under CP shared KV

Prefill CP only needs the local hidden shard for DeepSeek NextN draft extend. The change adds a draft shared-KV path that captures target hidden locally, feeds only the CP-local slice into the draft model, and keeps draft KV writes/transfers on the same shared logical-to-physical page mapping as target KV.\n\nDebug logs are gated behind SGLANG_CP_DRAFT_SHARED_KV_DEBUG and cover scheduler pool selection, KV manager buffer registration, local physical writes, prefill sender filtering, transfer pages, and decode commit metadata so ETE runs can prove draft KV is sharded rather than full-concatenated on a prefill rank.\n\nConstraint: Prefill runs CP while decode remains DP, so prefill must avoid full hidden/KV materialization but decode still receives full logical KV pages.\nRejected: Keep draft extend on full hidden state | preserves correctness but wastes prefill memory and defeats CP shared-KV intent.\nRejected: Transfer draft KV with a separate mapping | target and draft pools share req_to_token logical indices, so duplicating mapping adds risk without benefit.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not remove the debug logs until ETE evidence confirms draft MLA/index writes and transfer pages are CP-sharded on all ranks.\nTested: Remote compileall for changed CP draft, transfer, scheduler, NSA index, MLA write, and EAGLE files.\nNot-tested: Full GLM-5 EAGLE ETE with SGLANG_CP_DRAFT_SHARED_KV_DEBUG=1 after this logging addition; local pytest intentionally not run.
This commit is contained in:
laoyao0822
2026-05-13 22:29:18 +08:00
parent 3fc7a5c18c
commit 99b669f8b9
16 changed files with 951 additions and 31 deletions
@@ -77,6 +77,42 @@ if TYPE_CHECKING:
CLIP_MAX_NEW_TOKEN = envs.SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION.get()
def _cp_draft_shared_kv_debug(message: str, *args) -> None:
if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get():
logger.info("[CP_DRAFT_SHARED_KV] " + message, *args)
def _seq_summary(values) -> str:
if values is None:
return "None"
try:
size = len(values)
except TypeError:
return str(values)
if size == 0:
return "size=0"
try:
head = list(values[: min(8, size)])
except TypeError:
head = list(values)[: min(8, size)]
try:
min_val = min(values)
max_val = max(values)
return f"size={size} min={min_val} max={max_val} head={head}"
except (TypeError, ValueError):
return f"size={size} head={head}"
def _pool_summary(pool) -> str:
if pool is None:
return "None"
parts = [pool.__class__.__name__]
for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"):
if hasattr(pool, attr):
parts.append(f"{attr}={getattr(pool, attr)}")
return " ".join(parts)
def _kv_locs_to_page_indices_cpu(
kv_locs: torch.Tensor,
page_size: int,
@@ -321,16 +357,39 @@ class DecodePreallocQueue:
kv_data_ptrs, kv_data_lens, kv_item_lens = (
self.token_to_kv_pool.get_contiguous_buf_infos()
)
target_kv_buffer_count = len(kv_data_ptrs)
draft_kv_data_lens = []
draft_kv_item_lens = []
draft_kv_buffer_count = 0
if self.draft_token_to_kv_pool is not None:
# We should also transfer draft model kv cache. The indices are
# always shared with a target model.
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
self.draft_token_to_kv_pool.get_contiguous_buf_infos()
)
draft_kv_buffer_count = len(draft_kv_data_ptrs)
kv_data_ptrs += draft_kv_data_ptrs
kv_data_lens += draft_kv_data_lens
kv_item_lens += draft_kv_item_lens
kv_args.draft_kv_buffer_start = target_kv_buffer_count
kv_args.draft_kv_buffer_count = draft_kv_buffer_count
_cp_draft_shared_kv_debug(
"decode_kv_manager cp_rank=%s target_pool=(%s) draft_pool=(%s) "
"target_bufs=%s draft_bufs=%s total_bufs=%s target_lens=%s "
"draft_lens=%s target_item_lens=%s draft_item_lens=%s",
self.tp_rank,
_pool_summary(self.token_to_kv_pool),
_pool_summary(self.draft_token_to_kv_pool),
target_kv_buffer_count,
draft_kv_buffer_count,
len(kv_data_ptrs),
_seq_summary(kv_data_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_data_lens),
_seq_summary(kv_item_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_item_lens),
)
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
@@ -755,6 +814,19 @@ class DecodePreallocQueue:
self.req_to_metadata_buffer_idx_allocator.alloc()
)
assert decode_req.metadata_buffer_index is not None
_cp_draft_shared_kv_debug(
"decode_prealloc rid=%s room=%s origin_tokens=%s fill_tokens=%s "
"page_size=%s pages=%s state_pages=%s metadata_idx=%s has_draft_pool=%s",
decode_req.req.rid,
decode_req.req.bootstrap_room,
origin_input_len,
len(kv_loc),
page_size,
_seq_summary(page_indices),
_seq_summary(state_indices),
decode_req.metadata_buffer_index,
self.draft_token_to_kv_pool is not None,
)
decode_req.kv_receiver.init(
page_indices, decode_req.metadata_buffer_index, state_indices
)
@@ -1008,6 +1080,18 @@ class DecodeTransferQueue:
decode_req.req.output_topk_index = output_topk_index
decode_req.req.hidden_states_tensor = output_hidden_states
_cp_draft_shared_kv_debug(
"decode_transfer_commit rid=%s room=%s metadata_idx=%s cached_tokens=%s "
"topk_p_shape=%s topk_index_shape=%s hidden_shape=%s",
decode_req.req.rid,
decode_req.req.bootstrap_room,
idx,
decode_req.req.cached_tokens,
tuple(output_topk_p.shape) if output_topk_p is not None else None,
tuple(output_topk_index.shape) if output_topk_index is not None else None,
tuple(output_hidden_states.shape) if output_hidden_states is not None else None,
)
if decode_req.req.return_logprob:
decode_req.req.output_token_logprobs_val.append(
output_token_logprobs_val[0].item()
@@ -51,6 +51,17 @@ def _cp_shared_debug_log(key: str, message: str, *args, limit: int = 64) -> None
logger.info("[CP_SHARED_KV_DEBUG] " + message, *args)
def _cp_draft_shared_kv_debug(message: str, *args, limit: int = 64) -> None:
if not envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get():
return
key = "draft:" + message.split(" ", 1)[0]
count = _CP_SHARED_DEBUG_COUNTS.get(key, 0)
if count >= limit:
return
_CP_SHARED_DEBUG_COUNTS[key] = count + 1
logger.info("[CP_DRAFT_SHARED_KV] " + message, *args)
def _np_summary(arr) -> str:
if arr is None:
return "None"
@@ -246,6 +257,17 @@ class MooncakeKVManager(CommonKVManager):
def register_buffer_to_engine(self):
# Batch register KV data buffers
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",
self.disaggregation_mode,
self.attn_cp_rank,
len(self.kv_args.kv_data_ptrs),
getattr(self.kv_args, "draft_kv_buffer_start", None),
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),
)
self.engine.batch_register(
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_lens
)
@@ -847,6 +869,16 @@ class MooncakeKVManager(CommonKVManager):
chunked_dst_kv_indice = req.dst_kv_indices[
kv_chunk.index_slice
]
_cp_draft_shared_kv_debug(
"transfer_pages cp_rank=%s room=%s prefill_pages=%s "
"logical_positions=%s dst_pages=%s is_last=%s",
self.attn_cp_rank,
kv_chunk.room,
_np_summary(kv_chunk.prefill_kv_indices),
_np_summary(kv_chunk.logical_page_positions),
_np_summary(chunked_dst_kv_indice),
kv_chunk.is_last_chunk,
)
if envs.SGLANG_DEBUG_CP_SHARED_KV.get():
_cp_shared_debug_log(
"transfer_worker_kv",
@@ -1275,6 +1307,22 @@ class MooncakeKVSender(CommonKVSender):
_np_summary(state_logical_page_positions),
is_last_chunk,
)
_cp_draft_shared_kv_debug(
"sender_filter cp_rank=%s room=%s page_start=%s orig_kv_pages=%s "
"filtered_kv_pages=%s kv_positions=%s orig_state_pages=%s "
"filtered_state_pages=%s state_positions=%s is_last=%s draft_bufs=%s",
self.kv_mgr.attn_cp_rank,
self.bootstrap_room,
chunk_page_start,
_np_summary(orig_kv_indices),
_np_summary(kv_indices),
_np_summary(logical_page_positions),
_np_summary(orig_state_indices),
_np_summary(state_indices),
_np_summary(state_logical_page_positions),
is_last_chunk,
getattr(self.kv_mgr.kv_args, "draft_kv_buffer_count", None),
)
# Special handling for cp
elif self.kv_mgr.enable_all_cp_ranks_for_transfer:
kv_indices, index_slice = filter_kv_indices_for_cp_rank(
@@ -27,6 +27,7 @@ from typing import TYPE_CHECKING, List, Optional
import torch
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.environ import envs
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST,
@@ -61,6 +62,42 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _cp_draft_shared_kv_debug(message: str, *args) -> None:
if envs.SGLANG_CP_DRAFT_SHARED_KV_DEBUG.get():
logger.info("[CP_DRAFT_SHARED_KV] " + message, *args)
def _seq_summary(values) -> str:
if values is None:
return "None"
try:
size = len(values)
except TypeError:
return str(values)
if size == 0:
return "size=0"
try:
head = list(values[: min(8, size)])
except TypeError:
head = list(values)[: min(8, size)]
try:
min_val = min(values)
max_val = max(values)
return f"size={size} min={min_val} max={max_val} head={head}"
except (TypeError, ValueError):
return f"size={size} head={head}"
def _pool_summary(pool) -> str:
if pool is None:
return "None"
parts = [pool.__class__.__name__]
for attr in ("size", "page_size", "start_layer", "end_layer", "layer_num"):
if hasattr(pool, attr):
parts.append(f"{attr}={getattr(pool, attr)}")
return " ".join(parts)
def _kv_locs_to_page_indices_cpu(
kv_locs: torch.Tensor,
page_size: int,
@@ -154,16 +191,39 @@ class PrefillBootstrapQueue:
self.token_to_kv_pool.get_contiguous_buf_infos()
)
target_kv_buffer_count = len(kv_data_ptrs)
draft_kv_data_lens = []
draft_kv_item_lens = []
draft_kv_buffer_count = 0
if self.draft_token_to_kv_pool is not None:
# We should also transfer draft model kv cache. The indices are
# always shared with a target model.
draft_kv_data_ptrs, draft_kv_data_lens, draft_kv_item_lens = (
self.draft_token_to_kv_pool.get_contiguous_buf_infos()
)
draft_kv_buffer_count = len(draft_kv_data_ptrs)
kv_data_ptrs += draft_kv_data_ptrs
kv_data_lens += draft_kv_data_lens
kv_item_lens += draft_kv_item_lens
kv_args.draft_kv_buffer_start = target_kv_buffer_count
kv_args.draft_kv_buffer_count = draft_kv_buffer_count
_cp_draft_shared_kv_debug(
"prefill_kv_manager cp_rank=%s target_pool=(%s) draft_pool=(%s) "
"target_bufs=%s draft_bufs=%s total_bufs=%s target_lens=%s "
"draft_lens=%s target_item_lens=%s draft_item_lens=%s",
self.tp_rank,
_pool_summary(self.token_to_kv_pool),
_pool_summary(self.draft_token_to_kv_pool),
target_kv_buffer_count,
draft_kv_buffer_count,
len(kv_data_ptrs),
_seq_summary(kv_data_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_data_lens),
_seq_summary(kv_item_lens[:target_kv_buffer_count]),
_seq_summary(draft_kv_item_lens),
)
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
@@ -792,4 +852,18 @@ class SchedulerDisaggregationPrefillMixin:
f"Skip sending kv chunk for request {req.rid=} {req.bootstrap_room=} because page_indices is empty"
)
return
prefill_queue = getattr(self, "disagg_prefill_bootstrap_queue", None)
_cp_draft_shared_kv_debug(
"prefill_send_kv_chunk rid=%s room=%s start_idx=%s end_idx=%s "
"last_chunk=%s page_size=%s pages=%s state_pages=%s has_draft_pool=%s",
req.rid,
req.bootstrap_room,
start_idx,
end_idx,
last_chunk,
page_size,
_seq_summary(page_indices),
_seq_summary(state_indices),
getattr(prefill_queue, "draft_token_to_kv_pool", None) is not None,
)
req.disagg_kv_sender.send(page_indices, state_indices)