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