[PD] Support pipeline-parallel prefill with Mooncake staging buffer (#33807)

Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
YAMY
2026-08-11 13:53:17 +08:00
committed by GitHub
co-authored by Shangming Cai
parent 13aeb91b6e
commit 667e18d99d
7 changed files with 156 additions and 32 deletions
@@ -1369,6 +1369,7 @@ class CommonKVReceiver(BaseKVReceiver):
response = _get_bootstrap_session(self.bootstrap_addr).get(url, timeout=5)
if response.status_code == 200:
bootstrap_info = response.json()
bootstrap_info["pp_rank"] = int(target_pp_rank)
return bootstrap_info
else:
logger.error(
@@ -107,11 +107,15 @@ class DecodeStagingHandler:
self._wm_subscribers[key] = (receiver, session_id)
def num_writers_for(self, receiver) -> int:
"""Compute num_writers for a specific request based on its prefill TP."""
prefill_tp = receiver.prefill_info.attn_tp_size
"""Compute all TP and PP writers expected for a staging chunk."""
prefill_info = receiver.prefill_info
prefill_tp = prefill_info.attn_tp_size
if prefill_tp > self.decode_tp:
return prefill_tp // max(1, self.decode_tp)
return 1
tp_writers = prefill_tp // max(1, self.decode_tp)
else:
tp_writers = 1
pp_writers = prefill_info.pp_size // self.kv_manager.pp_size
return tp_writers * pp_writers
@classmethod
def create(cls, kv_manager, scheduler, tp_rank: int) -> DecodeStagingHandler:
@@ -617,6 +621,7 @@ class PrefillStagingStrategy:
target_info.dst_tp_rank,
target_info.dst_attn_tp_size,
target_info.dst_kv_item_len,
target_info.dst_kv_layer_ids,
staging_buffer=self.staging_buffer,
)
except Exception as e:
@@ -716,6 +721,7 @@ def handle_staging_req(
chunk_idx = int(msg[2].decode("ascii"))
chunk_num_pages = int(msg[3].decode("ascii"))
session_id = msg[4].decode("ascii")
requester_pp_rank = int(msg[5].decode("ascii")) if len(msg) > 5 else None
if staging_allocator is None:
logger.warning(
@@ -798,6 +804,8 @@ def handle_staging_req(
bootstrap_infos = room_bootstrap.get(room)
if bootstrap_infos:
for bi in bootstrap_infos:
if requester_pp_rank is not None and bi["pp_rank"] != requester_pp_rank:
continue
try:
sock, lock = receiver._connect_to_bootstrap_server(bi)
with lock:
@@ -823,6 +831,7 @@ def prefetch_staging_reqs(
chunked_prefill_size: int,
staging_requested: set,
prefetch_sockets: dict,
requester_pp_rank: Optional[int] = None,
) -> None:
"""Send STAGING_REQ for all chunks before the prefill forward starts.
@@ -868,14 +877,15 @@ def prefetch_staging_reqs(
sock.setsockopt(zmq.IPV6, 1)
sock.connect(ep)
prefetch_sockets[ep] = sock
prefetch_sockets[ep].send_multipart(
[
b"STAGING_REQ",
str(room).encode("ascii"),
str(chunk_idx).encode("ascii"),
str(chunk_pages).encode("ascii"),
session_id.encode("ascii"),
]
)
request = [
b"STAGING_REQ",
str(room).encode("ascii"),
str(chunk_idx).encode("ascii"),
str(chunk_pages).encode("ascii"),
session_id.encode("ascii"),
]
if requester_pp_rank is not None:
request.append(str(requester_pp_rank).encode("ascii"))
prefetch_sockets[ep].send_multipart(request)
except Exception:
staging_requested.discard(stg_key)
+3 -1
View File
@@ -510,7 +510,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
per_rank_kv_heads = getattr(kv_pool_for_heads, "head_num", 0)
if per_rank_kv_heads > 0:
kv_args.kv_head_num = per_rank_kv_heads
kv_args.total_kv_head_num = per_rank_kv_heads * attn_tp_size
kv_args.total_kv_head_num = (
self.scheduler.model_config.get_total_num_kv_heads()
)
if hasattr(kv_manager, "set_kv_buffer_tensors"):
kv_pool = kv_pool_for_heads
if hasattr(kv_pool, "k_buffer") and hasattr(kv_pool, "v_buffer"):
@@ -520,6 +520,7 @@ class MooncakeKVManager(CommonKVManager):
get_schedule().chunked_prefill_size,
self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets,
requester_pp_rank=self.pp_rank,
)
def send_kvcache_staged(
@@ -531,6 +532,7 @@ class MooncakeKVManager(CommonKVManager):
dst_tp_rank: int,
dst_attn_tp_size: int,
dst_kv_item_len: int,
dst_layer_ids: List[int],
staging_buffer=None,
) -> int:
"""Transfer KV cache via staging buffers (gather -> bulk RDMA -> scatter on decode)."""
@@ -563,7 +565,19 @@ class MooncakeKVManager(CommonKVManager):
num_tokens = len(prefill_kv_indices) * page_size
per_layer_bytes = num_tokens * num_heads_to_send * head_dim * dtype_size
per_rank_bytes = per_layer_bytes * num_layers * 2
local_bytes = per_layer_bytes * num_layers * 2
if self.pp_size > 1:
pairs = build_transfer_entry_pairs(
self.kv_args.kv_layer_ids,
dst_layer_ids,
num_layers * 2,
len(dst_layer_ids),
)
dst_num_layers = len(dst_layer_ids) // 2
else:
pairs = None
dst_num_layers = num_layers
num_writers, writer_rank_bytes, total_staging_needed = compute_staging_layout(
self.attn_tp_size,
@@ -572,21 +586,20 @@ class MooncakeKVManager(CommonKVManager):
total_kv_heads,
num_tokens,
head_dim * dtype_size,
num_layers,
dst_num_layers,
)
writer_idx = local_tp_rank % num_writers if num_writers > 1 else 0
rank_offset = sum(writer_rank_bytes[:writer_idx])
if not staging_buffer.fits(per_rank_bytes):
if not staging_buffer.fits(local_bytes):
logger.warning(
f"Prefill staging too small for {per_rank_bytes} bytes, falling back"
f"Prefill staging too small for {local_bytes} bytes, falling back"
)
return -1
if dst_staging_size < total_staging_needed:
logger.warning(
f"Decode staging too small: need {total_staging_needed} bytes "
f"({num_writers if self.attn_tp_size > dst_attn_tp_size else 1} writers "
f"x {per_rank_bytes} bytes/rank), have {dst_staging_size}, falling back"
f"for {dst_num_layers} layers, have {dst_staging_size}, falling back"
)
return -1
@@ -605,16 +618,29 @@ class MooncakeKVManager(CommonKVManager):
self.kv_args.gpu_id,
)
dst_write_ptr = dst_staging_ptr + rank_offset
ret = self._transfer_data(
mooncake_session_id,
[(staging_buffer.get_ptr(), dst_write_ptr, per_rank_bytes)],
)
if pairs is None:
transfer_blocks = [
(
staging_buffer.get_ptr(),
dst_staging_ptr + rank_offset,
local_bytes,
)
]
else:
transfer_blocks = [
(
staging_buffer.get_ptr() + src_idx * per_layer_bytes,
dst_staging_ptr + rank_offset + dst_idx * per_layer_bytes,
per_layer_bytes,
)
for src_idx, dst_idx in pairs
]
ret = self._transfer_data(mooncake_session_id, transfer_blocks)
if ret != 0:
raise RuntimeError(
f"[Staging] Bulk RDMA transfer failed with ret={ret}. "
f"src_ptr=0x{staging_buffer.get_ptr():x}, "
f"dst_ptr=0x{dst_write_ptr:x}, size={per_rank_bytes}. "
f"dst_ptr=0x{dst_staging_ptr + rank_offset:x}, size={local_bytes}. "
f"The decode staging buffer may not be properly registered."
)
return ret
+3 -3
View File
@@ -175,10 +175,10 @@ class PrefillBootstrapQueue:
f"chunked_prefill_size that is a multiple of page_size "
f"({page_size}); got {chunked_prefill_size}."
)
if self.pp_size > 1:
# Staging writer accounting has no pp dimension.
if self.pp_size > 1 and self.transfer_backend != TransferBackend.MOONCAKE:
raise RuntimeError(
"SGLANG_DISAGG_STAGING_BUFFER does not support pp_size > 1."
"SGLANG_DISAGG_STAGING_BUFFER with pp_size > 1 is only "
"supported by Mooncake."
)
if get_parallel().enable_prefill_context_parallel:
# CP rewrites index_slice per rank, breaking the chunk grid.
@@ -278,6 +278,8 @@ class SchedulerPPMixin:
)
self._pp_commit_comm_work(self.send_proxy_work)
if cur_batch:
if self.enable_staging:
self.maybe_prefetch_staging_for_batch(cur_batch)
result, self.launch_event = self._pp_launch_batch(
mb_id,
cur_batch,