[PD] Pack DCP1→DCP-N PD KV transfers into dest-contiguous RDMA blocks (#35762)

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Khoa Pham
2026-08-28 21:41:30 -07:00
committed by GitHub
co-authored by Cursor Claude Opus 5
parent 5f216fc33f
commit 3760296be8
10 changed files with 581 additions and 60 deletions
@@ -61,6 +61,7 @@ __all__ = ["reshape_and_cache_flash"]
_TRITON_KERNELS = [
("cache_ops", "concat_and_cast_mha_k_triton"),
("cache_ops", "launch_reshape_and_cache_flash"),
("pd_dcp_gather", "copy_mla_rows_into_pack"),
("kv_indices", "create_flashinfer_kv_indices_triton"),
("kv_indices", "create_flashmla_kv_indices_triton"),
("kv_indices", "create_chunked_prefix_cache_kv_indices"),
@@ -0,0 +1,66 @@
from typing import Sequence
import torch
import triton
import triton.language as tl
@triton.jit
def _copy_mla_rows_into_pack_kernel(
src_metadata,
row_indices,
pack,
num_rows,
BLOCK_SIZE: tl.constexpr,
):
layer_id = tl.program_id(0)
block_id = tl.program_id(1)
metadata_offset = layer_id * 3
src = tl.load(src_metadata + metadata_offset).to(pack.dtype)
row_nbytes = tl.load(src_metadata + metadata_offset + 1)
pack_offset = tl.load(src_metadata + metadata_offset + 2)
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
layer_nbytes = num_rows * row_nbytes
mask = offsets < layer_nbytes
row = offsets // row_nbytes
byte = offsets % row_nbytes
src_row = tl.load(row_indices + row, mask=mask, other=0)
values = tl.load(src + src_row * row_nbytes + byte, mask=mask)
tl.store(pack + pack_offset + offsets, values, mask=mask)
def copy_mla_rows_into_pack(
kv_data_ptrs: Sequence[int],
row_indices: torch.Tensor,
pack: torch.Tensor,
token_item_lens: Sequence[int],
) -> None:
if len(kv_data_ptrs) != len(token_item_lens):
raise ValueError(
"kv_data_ptrs and token_item_lens length mismatch: "
f"{len(kv_data_ptrs)} vs {len(token_item_lens)}"
)
if not kv_data_ptrs:
return
n = int(row_indices.numel())
metadata = []
offset = 0
for ptr, item_len in zip(kv_data_ptrs, token_item_lens):
item_len = int(item_len)
if item_len <= 0:
raise ValueError(f"MLA token item length must be positive, got {item_len}")
metadata.extend((int(ptr), item_len, offset))
offset += n * item_len
src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device)
max_item_len = max(int(item_len) for item_len in token_item_lens)
grid = (len(kv_data_ptrs), triton.cdiv(n * max_item_len, 1024))
_copy_mla_rows_into_pack_kernel[grid](
src_metadata,
row_indices,
pack,
n,
BLOCK_SIZE=1024,
)
@@ -167,6 +167,7 @@ class CommonKVManager(BaseKVManager):
self.enable_deferred_decode_kv_release = (
envs.SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE.get()
)
self._dcp_pack_buffers = None
# for p/d multi node infer
self.bootstrap_host = get_serving().host
self.bootstrap_port = get_disagg().disaggregation_bootstrap_port
@@ -339,6 +340,25 @@ class CommonKVManager(BaseKVManager):
)
return src_token_lens
def _register_staging_memory(self, ptr: int, size: int) -> None:
raise NotImplementedError(
f"{type(self).__name__} does not support staging memory registration"
)
def _init_dcp_pack_buffers_once(self, dcp_size: int) -> None:
if self._dcp_pack_buffers is not None:
return
if not self.kv_args.kv_item_lens:
return
from sglang.srt.disaggregation.common.dcp_pack import init_dcp_pack_buffers
self._dcp_pack_buffers = init_dcp_pack_buffers(
self._register_staging_memory,
self.kv_args,
len(self.transfer_queues),
dcp_size,
)
def check_status(self, bootstrap_room: int) -> KVPoll:
return self.request_status[bootstrap_room]
@@ -0,0 +1,120 @@
from __future__ import annotations
import logging
from typing import List, Optional, Sequence, Tuple
import numpy as np
import numpy.typing as npt
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__)
def dcp_pack_buffer_bytes(
kv_item_lens: Sequence[int], page_size: int, max_tokens: int, dcp_size: int = 1
) -> int:
if page_size <= 0:
raise ValueError(f"page_size must be positive, got {page_size}")
if dcp_size <= 0:
raise ValueError(f"dcp_size must be positive, got {dcp_size}")
if any(item_len < page_size for item_len in kv_item_lens):
raise ValueError(
"PD DCP pack requires each kv_item_len to span at least one page, "
f"got {list(kv_item_lens)} with page_size={page_size}"
)
if any(item_len % page_size != 0 for item_len in kv_item_lens):
raise ValueError(
"PD DCP pack requires page-aligned kv_item_lens, "
f"got {list(kv_item_lens)} with page_size={page_size}"
)
token_item_lens = [item_len // page_size for item_len in kv_item_lens]
rank_tokens = (max_tokens + dcp_size - 1) // dcp_size
return dcp_size * rank_tokens * sum(token_item_lens)
def try_pack_dcp_src(
*,
pack_buffer: StagingBuffer,
kv_data_ptrs: Sequence[int],
src_token_indices: npt.NDArray[np.integer],
token_item_lens: Sequence[int],
pack_offset_bytes: int = 0,
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
if pack_offset_bytes < 0:
raise ValueError(
f"pack_offset_bytes must be non-negative, got {pack_offset_bytes}"
)
n = int(src_token_indices.size)
if n == 0:
empty = np.empty((0,), dtype=np.int64)
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):
logger.warning(
"PD DCP pack buffer too small for byte range [%s, %s) (have %s); "
"falling back to per-token RDMA",
pack_offset_bytes,
required_end,
pack_buffer.get_size(),
)
return None
pack = pack_buffer.buffer.narrow(0, pack_offset_bytes, required)
row_indices = torch.as_tensor(
src_token_indices, device=pack.device, dtype=torch.int64
)
gather_stream = pack_buffer.get_gather_stream()
gather_stream.wait_stream(torch.cuda.default_stream(pack.device))
with torch.cuda.stream(gather_stream):
copy_mla_rows_into_pack(kv_data_ptrs, row_indices, pack, token_item_lens)
gather_stream.synchronize()
packed_ptrs: List[int] = []
offset = 0
base = pack_buffer.get_ptr() + pack_offset_bytes
for item_len in token_item_lens:
packed_ptrs.append(base + offset)
offset += n * int(item_len)
return packed_ptrs, np.arange(n, dtype=np.int64)
def init_dcp_pack_buffers(
register_fn,
kv_args,
count: int,
dcp_size: 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
# Note(kpham-sgl): size = dcp_size x ceil(max_tokens / dcp_size)
# x sum(per-layer token bytes). At 32,768 tokens and 61 MLA layers
# x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
size_bytes = dcp_pack_buffer_bytes(
kv_args.kv_item_lens, kv_args.page_size, max_tokens, dcp_size
)
gpu_id = kv_args.gpu_id
device = f"cuda:{gpu_id}"
custom_mem_pool, _ = _get_custom_mem_pool(device)
buffers = []
for _ in range(count):
buf = StagingBuffer(size_bytes, device, gpu_id, custom_mem_pool=custom_mem_pool)
register_fn(buf.get_ptr(), buf.get_size())
buffers.append(buf)
logger.info(
"PD DCP pack buffers allocated: %d x %.1f MB (max_tokens=%d)",
count,
size_bytes / (1024 * 1024),
max_tokens,
)
return buffers
@@ -133,6 +133,7 @@ class StagingBuffer:
self.size_bytes = size_bytes
self.device = device
self.gpu_id = gpu_id
self._gather_stream: Optional[torch.cuda.Stream] = None
torch.cuda.set_device(gpu_id)
if custom_mem_pool is not None:
@@ -158,6 +159,11 @@ class StagingBuffer:
def fits(self, required_bytes: int) -> bool:
return required_bytes <= self.size_bytes
def get_gather_stream(self) -> torch.cuda.Stream:
if self._gather_stream is None:
self._gather_stream = torch.cuda.Stream(device=self.device)
return self._gather_stream
class StagingAllocator:
"""Decode-side dynamic staging ring buffer allocator with overcommit.
@@ -350,16 +356,12 @@ def _gather_all_layers_torch(
gather_idx = token_indices.view(-1, 1, 1).expand(num_tokens, num_heads, head_dim)
if not hasattr(staging_buffer, "_gather_stream"):
staging_buffer._gather_stream = torch.cuda.Stream(device=device)
staging_buffer._gather_stream.wait_stream(
torch.cuda.default_stream(torch.device(device))
)
gather_stream = staging_buffer.get_gather_stream()
gather_stream.wait_stream(torch.cuda.default_stream(torch.device(device)))
staging_view = staging_buffer.buffer
offset = 0
with torch.cuda.stream(staging_buffer._gather_stream):
with torch.cuda.stream(gather_stream):
for layer_id in range(num_layers):
dst = (
staging_view[offset : offset + per_layer_bytes]
@@ -389,7 +391,7 @@ def _gather_all_layers_torch(
)
offset += per_layer_bytes
staging_buffer._gather_stream.synchronize()
gather_stream.synchronize()
return offset
@@ -430,17 +432,13 @@ def _gather_all_layers_triton(
int_dtype = int_dtype_map.get(dtype_size, torch.int16)
staging_typed = staging_buffer.buffer[:total_bytes].view(int_dtype)
if not hasattr(staging_buffer, "_gather_stream"):
staging_buffer._gather_stream = torch.cuda.Stream(device=device)
staging_buffer._gather_stream.wait_stream(
torch.cuda.default_stream(torch.device(device))
)
gather_stream = staging_buffer.get_gather_stream()
gather_stream.wait_stream(torch.cuda.default_stream(torch.device(device)))
BLOCK_SIZE = 1024
grid = (2 * num_layers, triton.cdiv(per_layer_elems, BLOCK_SIZE))
with torch.cuda.stream(staging_buffer._gather_stream):
with torch.cuda.stream(gather_stream):
_fused_gather_to_staging_kernel[grid](
layer_ptrs,
page_idx_tensor,
@@ -454,7 +452,7 @@ def _gather_all_layers_triton(
BLOCK_SIZE,
)
staging_buffer._gather_stream.synchronize()
gather_stream.synchronize()
return total_bytes
@@ -353,13 +353,16 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
"slot_layer_ids": list(slot_layer_ids or []),
}
def _register_staging_memory(self, ptr: int, size: int) -> None:
self.engine.batch_register([ptr], [size])
def _init_staging_buffers(self, count: int):
from sglang.srt.disaggregation.common.staging_handler import (
init_staging_buffers,
)
self._staging_ctx.buffers = init_staging_buffers(
lambda ptr, size: self.engine.batch_register([ptr], [size]),
self._register_staging_memory,
self.kv_args,
count,
get_schedule().chunked_prefill_size,
@@ -372,7 +375,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
)
self._staging_ctx.allocator = init_staging_allocator(
lambda ptr, size: self.engine.batch_register([ptr], [size]),
self._register_staging_memory,
self.kv_args,
)
self.kv_buffer_tensors = None
@@ -910,6 +913,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
num_kv_tokens: int,
executor: concurrent.futures.ThreadPoolExecutor,
dst_layer_ids: List[int],
pack_buffer=None,
) -> int:
if num_kv_tokens is None:
raise ValueError("PD DCP transfer requires num_kv_tokens")
@@ -942,10 +946,24 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
self.kv_args.kv_data_ptrs,
dst_kv_ptrs,
)
src_token_indices = plan.src_token_indices
dst_token_indices = plan.dst_token_indices
if pack_buffer is not None:
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
packed = try_pack_dcp_src(
pack_buffer=pack_buffer,
kv_data_ptrs=src_kv_ptrs,
src_token_indices=src_token_indices,
token_item_lens=dcp_token_item_lens[: len(src_kv_ptrs)],
)
if packed is not None:
src_kv_ptrs, src_token_indices = packed
layers_current_pp_stage = len(src_kv_ptrs)
src_groups, dst_groups = group_concurrent_contiguous(
plan.src_token_indices,
plan.dst_token_indices,
src_token_indices,
dst_token_indices,
)
layers_params = [
@@ -1771,6 +1789,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
target_rank_registration_info.dcp_token_item_lens
)
assert dcp_token_item_lens is not None
pack_buffer = (
self._dcp_pack_buffers[worker_index]
if self._dcp_pack_buffers
else None
)
ret = self.send_kvcache_dcp(
req.mooncake_session_id,
kv_chunk.prefill_kv_indices,
@@ -1786,6 +1809,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
dst_layer_ids=(
target_rank_registration_info.dst_kv_layer_ids
),
pack_buffer=pack_buffer,
)
elif (
self.is_mla_backend
@@ -2073,6 +2097,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
* len(self.kv_args.kv_item_lens)
)
)
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
with self.session_lock:
if mooncake_session_id in self.failed_sessions:
+81 -37
View File
@@ -506,7 +506,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
)
threading.Thread(
target=self.transfer_worker,
args=(queue, staging_buffer),
args=(queue, staging_buffer, i),
daemon=True,
).start()
self._start_bootstrap_thread()
@@ -545,9 +545,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
init_staging_buffers,
)
gpu_id = self.kv_args.gpu_id
self._staging_ctx.buffers = init_staging_buffers(
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
self._register_staging_memory,
self.kv_args,
count,
get_schedule().chunked_prefill_size,
@@ -558,15 +557,14 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
init_staging_allocator,
)
gpu_id = self.kv_args.gpu_id
self._staging_ctx.allocator = init_staging_allocator(
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
self._register_staging_memory,
self.kv_args,
)
def _register_staging_memory(self, ptr: int, size: int, gpu_id: int):
def _register_staging_memory(self, ptr: int, size: int):
"""Register a staging buffer with the NIXL agent."""
addrs = [(ptr, size, gpu_id, "")]
addrs = [(ptr, size, self.kv_args.gpu_id, "")]
descs = self.agent.register_memory(addrs, "VRAM")
if not descs:
raise RuntimeError(
@@ -1080,7 +1078,7 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
dst_mem_kind=dst_mem_kind,
)
def transfer_worker(self, queue: FastQueue, staging_buffer=None):
def transfer_worker(self, queue: FastQueue, staging_buffer=None, worker_index=0):
# Per-worker staging strategy: lazy-created on first chunk so we
# see kv_buffer_tensors (set by ModelRunner after engine init).
# Never cache on self -- multiple workers would race the ring.
@@ -1120,6 +1118,10 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
self.update_status(room, KVPoll.Transferring)
reqs_to_be_processed = list(self.transfer_infos[room].values())
# Note(kpham-sgl): Pack each DCP rank once into its fixed region.
# NIXL reads regions asynchronously; the chunk barrier prevents
# reuse until every transfer completes.
packed_source_by_dcp_rank = {}
# Set when staging allocation/watermark is not yet ready and
# the chunk has been re-enqueued. We then break out of the
@@ -1209,15 +1211,37 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
if kv_xfer_handle is None:
if is_dcp_transfer:
kv_xfer_handle = self.send_kvcache_dcp(
req.agent_name,
pack_buffer = (
self._dcp_pack_buffers[worker_index]
if self._dcp_pack_buffers
else None
)
if kv_chunk.num_kv_tokens is None:
raise ValueError(
"PD DCP transfer requires num_kv_tokens"
)
plan = build_dcp_token_transfer_plan(
src_prefill_kv_indices,
dst_info,
chunked_dst_kv_indice,
src_page_offset=kv_chunk.index_slice.start or 0,
physical_page_size=self.kv_args.page_size,
dcp_size=dst_info.dst_dcp_size,
dcp_rank=dst_info.dst_dcp_rank,
src_page_offset=(kv_chunk.index_slice.start or 0),
decode_prefix_len=req.decode_prefix_len or 0,
num_kv_tokens=kv_chunk.num_kv_tokens,
notif=notif,
)
packed_src = self._pack_dcp_rank_once(
pack_buffer,
dst_info,
plan.src_token_indices,
packed_source_by_dcp_rank,
)
kv_xfer_handle = self.send_kvcache_dcp(
req.agent_name,
dst_info,
plan,
notif,
packed_src,
)
elif (
self.is_mla_backend
@@ -1434,6 +1458,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout(
decode_kv_args.dst_dcp_size, decode_kv_args.dst_dcp_rank
)
if decode_kv_args.requires_dcp_relayout:
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
self.decode_kv_args_table[agent_name] = decode_kv_args
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
if self.disaggregation_mode == DisaggregationMode.PREFILL:
@@ -1632,36 +1658,51 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
dst_mem_kind=dst_mem_kind,
)
def _pack_dcp_rank_once(
self,
pack_buffer,
dst_info: KVArgsRegisterInfo,
src_token_indices,
packed_source_by_dcp_rank,
):
"""Pack one source region per DCP rank for the current chunk.
Note(kpham-sgl): TP ranks may share a DCP rank (`tp_rank % dcp_size`),
so they reuse one packed source while sending to distinct destination GPUs.
"""
rank = dst_info.dst_dcp_rank
if rank in packed_source_by_dcp_rank:
return packed_source_by_dcp_rank[rank]
if pack_buffer is None or src_token_indices.size == 0:
packed_source_by_dcp_rank[rank] = None
return None
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
token_item_lens = dst_info.dcp_token_item_lens
assert token_item_lens is not None
rank_stride = pack_buffer.get_size() // dst_info.dst_dcp_size
packed_source_by_dcp_rank[rank] = try_pack_dcp_src(
pack_buffer=pack_buffer,
kv_data_ptrs=self.kv_args.kv_data_ptrs,
src_token_indices=src_token_indices,
token_item_lens=token_item_lens[: len(self.kv_args.kv_data_ptrs)],
pack_offset_bytes=rank * rank_stride,
)
return packed_source_by_dcp_rank[rank]
def send_kvcache_dcp(
self,
peer_name: str,
prefill_kv_indices: npt.NDArray[np.int32],
dst_info: KVArgsRegisterInfo,
dst_kv_indices: npt.NDArray[np.int32],
*,
src_page_offset: int,
decode_prefix_len: int,
num_kv_tokens: int,
plan,
notif: str,
packed_src,
):
if self.src_mem_kind is None:
raise RuntimeError("Missing NIXL source KV memory kind")
if dst_info.dst_homogeneous_mem_kind is None:
raise RuntimeError("Missing NIXL destination KV memory kind")
if num_kv_tokens is None:
raise ValueError("PD DCP transfer requires num_kv_tokens")
physical_page_size = self.kv_args.page_size
plan = build_dcp_token_transfer_plan(
prefill_kv_indices,
dst_kv_indices,
physical_page_size=physical_page_size,
dcp_size=dst_info.dst_dcp_size,
dcp_rank=dst_info.dst_dcp_rank,
src_page_offset=src_page_offset,
decode_prefix_len=decode_prefix_len,
num_kv_tokens=num_kv_tokens,
)
if plan.src_token_indices.size == 0:
self.agent.send_notif(peer_name, notif.encode("ascii"))
return None
@@ -1671,15 +1712,18 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
dst_kv_ptrs = [
dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices
]
src_kv_ptrs = self.kv_args.kv_data_ptrs
src_token_indices = plan.src_token_indices
if packed_src is not None:
src_kv_ptrs, src_token_indices = packed_src
token_item_lens = token_item_lens[: len(src_kv_ptrs)]
# Prepared handles encode page-level offsets, while DCP relayout needs
# flat descriptors for the selected token rows.
return self._send_kvcache_generic(
peer_name=peer_name,
src_data_ptrs=self.kv_args.kv_data_ptrs,
src_data_ptrs=src_kv_ptrs,
dst_data_ptrs=dst_kv_ptrs,
item_lens=token_item_lens,
prefill_data_indices=plan.src_token_indices,
prefill_data_indices=src_token_indices,
dst_data_indices=plan.dst_token_indices,
dst_gpu_id=dst_info.gpu_id,
notif=notif,