[PD] Bound cached-prefix DCP transfers by pack capacity (#40376)
This commit is contained in:
@@ -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]
|
||||
|
||||
|
||||
@@ -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)]
|
||||
|
||||
Reference in New Issue
Block a user