[PD] Bound cached-prefix DCP transfers by pack capacity (#40376)

This commit is contained in:
Khoa Pham
2026-09-21 01:17:15 +08:00
committed by GitHub
parent 80da4432d0
commit b3e4d198af
8 changed files with 297 additions and 16 deletions
@@ -171,6 +171,10 @@ class BaseKVSender(ABC):
def pop_decode_prefix_len(self) -> int:
return 0
def get_max_transfer_tokens(self) -> Optional[int]:
"""Optional page-aligned limit for one scheduler KV send."""
return None
def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool:
return num_pages > 0
@@ -35,9 +35,12 @@ from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_disagg,
get_parallel,
get_schedule,
get_serving,
max_prefill_buffer_tokens,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ceil_align
from sglang.srt.utils.network import (
NetworkAddress,
get_local_ip_auto,
@@ -181,6 +184,7 @@ class CommonKVManager(BaseKVManager):
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
)
self._dcp_pack_buffers = None
self._dcp_pack_max_tokens: Optional[int] = None
# for p/d multi node infer
self.bootstrap_host = get_serving().host
self.bootstrap_port = get_disagg().disaggregation_bootstrap_port
@@ -385,12 +389,16 @@ class CommonKVManager(BaseKVManager):
return
from sglang.srt.disaggregation.common.dcp_pack import init_dcp_pack_buffers
max_tokens = max_prefill_buffer_tokens() or get_schedule().max_prefill_tokens
max_tokens = ceil_align(max_tokens, self.kv_args.page_size)
self._dcp_pack_buffers = init_dcp_pack_buffers(
self._register_staging_memory,
self.kv_args,
len(self.transfer_queues),
dcp_size,
max_tokens,
)
self._dcp_pack_max_tokens = max_tokens
def check_status(self, bootstrap_room: int) -> KVPoll:
return self.request_status[bootstrap_room]
@@ -1541,6 +1549,19 @@ class CommonKVSender(BaseKVSender):
def pop_decode_prefix_len(self) -> int:
return self.kv_mgr.req_to_decode_prefix_len.pop(self.bootstrap_room, 0)
def get_max_transfer_tokens(self) -> Optional[int]:
if self.kv_mgr._dcp_pack_max_tokens is None:
return None
for peer, info in self.kv_mgr.transfer_infos.get(
self.bootstrap_room, {}
).items():
if (
not info.is_dummy
and self.kv_mgr.decode_kv_args_table[peer].requires_dcp_relayout
):
return self.kv_mgr._dcp_pack_max_tokens
return None
def should_send_kv_chunk(self, num_pages: int, last_chunk: bool) -> bool:
return num_pages > 0 or last_chunk
@@ -9,7 +9,6 @@ import torch
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
from sglang.srt.runtime_context import get_schedule, max_prefill_buffer_tokens
logger = logging.getLogger(__name__)
@@ -43,6 +42,7 @@ def try_pack_dcp_src(
src_token_indices: npt.NDArray[np.integer],
token_item_lens: Sequence[int],
pack_offset_bytes: int = 0,
pack_capacity_bytes: Optional[int] = None,
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
if pack_offset_bytes < 0:
raise ValueError(
@@ -54,13 +54,17 @@ def try_pack_dcp_src(
return [], empty
required = n * sum(int(item_len) for item_len in token_item_lens)
required_end = pack_offset_bytes + required
if not pack_buffer.fits(required_end):
if (
pack_capacity_bytes is not None and required > pack_capacity_bytes
) or not pack_buffer.fits(required_end):
logger.warning(
"PD DCP pack buffer too small for byte range [%s, %s) (have %s); "
"PD DCP pack buffer too small for byte range [%s, %s) "
"(have %s, region capacity %s); "
"falling back to per-token RDMA",
pack_offset_bytes,
required_end,
pack_buffer.get_size(),
pack_capacity_bytes,
)
return None
@@ -88,14 +92,12 @@ def init_dcp_pack_buffers(
kv_args,
count: int,
dcp_size: int,
max_tokens: int,
) -> List[StagingBuffer]:
from sglang.srt.disaggregation.common.staging_handler import (
_get_custom_mem_pool,
)
max_tokens = max_prefill_buffer_tokens()
if max_tokens <= 0:
max_tokens = get_schedule().max_prefill_tokens
kv_item_lens = kv_args.kv_item_lens
if kv_args.num_draft_entries > 0:
kv_item_lens = kv_item_lens[: len(kv_item_lens) - kv_args.num_draft_entries]
@@ -1867,6 +1867,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
src_token_indices=src_token_indices,
token_item_lens=token_item_lens[:num_target],
pack_offset_bytes=rank * rank_stride,
pack_capacity_bytes=rank_stride,
)
return packed_source_by_dcp_rank[rank]
+7 -1
View File
@@ -1449,15 +1449,21 @@ class SchedulerDisaggregationPrefillMixin:
payloads[st]() if st in payloads else None for st in state_types
]
transfer_chunk_tokens = req.disagg_kv_sender.get_max_transfer_tokens()
if self.enable_staging:
# One sender.send per grid slot; the sender's cumulative page
# counter marks only the final sub-send of the final chunk as
# is_last, routing aux/state correctly.
transfer_chunk_tokens = staging_grid_tokens(
get_schedule().chunked_prefill_size, page_size
)
if transfer_chunk_tokens is not None:
# Prefill cache hits can leave more KV to transfer than the DCP pack buffer holds.
segments = compute_grid_segments(
start_idx,
end_idx,
req.disagg_decode_prefix_len,
staging_grid_tokens(get_schedule().chunked_prefill_size, page_size),
transfer_chunk_tokens,
)
else:
segments = [(start_idx, end_idx)]