[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,
@@ -1,12 +1,16 @@
import struct
import threading
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from unittest.mock import Mock, patch
import numpy as np
import torch
from sglang.srt.disaggregation.base.conn import KVArgs, StateType
from sglang.srt.disaggregation.common.staging_handler import (
handle_staging_req,
)
from sglang.srt.disaggregation.common.utils import (
group_concurrent_contiguous,
pack_int_lists,
@@ -15,7 +19,8 @@ from sglang.srt.disaggregation.common.utils import (
unpack_list_of_buffers,
)
from sglang.srt.disaggregation.mooncake.conn import (
KVArgsRegisterInfo as MooncakeKVArgsRegisterInfo,
KVArgsRegisterInfo,
MooncakeKVManager,
)
from sglang.srt.disaggregation.utils import (
MetadataBuffers,
@@ -57,7 +62,7 @@ class TestDisaggregationWire(unittest.TestCase):
b"2",
]
info = MooncakeKVArgsRegisterInfo.from_zmq(msg)
info = KVArgsRegisterInfo.from_zmq(msg)
self.assertEqual(info.staging_base_ptr, 0x3000)
self.assertEqual(info.staging_total_size, 4096)
@@ -134,6 +139,84 @@ class TestGroupConcurrentContiguous(unittest.TestCase):
group_concurrent_contiguous(self._arr([1, 2, 3]), self._arr([1, 2]))
class TestMooncakePPStaging(unittest.TestCase):
def test_staging_response_targets_requesting_pp_rank(self):
sock = Mock()
receiver = SimpleNamespace(
chunk_staging_infos=[],
_connect_to_bootstrap_server=Mock(return_value=(sock, threading.Lock())),
)
allocator = SimpleNamespace(
assign=Mock(return_value=(3, 128, 0)), total_size=1 << 20
)
kv_args = SimpleNamespace(
page_size=64,
kv_item_lens=[4096, 4096],
total_kv_head_num=4,
engine_rank=0,
)
target = {"pp_rank": 3}
handle_staging_req(
[b"STAGING_REQ", b"7", b"0", b"1", b"peer", b"3"],
allocator,
kv_args,
attn_tp_size=16,
prefill_attn_tp_size=1,
kv_buffer_tensors=None,
room_receivers={7: receiver},
room_bootstrap={7: [{"pp_rank": 2}, target]},
)
receiver._connect_to_bootstrap_server.assert_called_once_with(target)
sock.send_multipart.assert_called_once()
@patch(
"sglang.srt.disaggregation.common.staging_buffer.gather_all_layers_to_staging"
)
def test_pp_stage_writes_its_global_layer_slots(self, gather):
manager = object.__new__(MooncakeKVManager)
tensor = SimpleNamespace(shape=(1, 1, 8), element_size=lambda: 2)
manager.kv_buffer_tensors = {
"k_buffers": [tensor],
"v_buffers": [tensor],
"page_size": 2,
}
manager.attn_tp_size = 1
manager.pp_size = 16
manager.kv_args = SimpleNamespace(
engine_rank=0,
gpu_id=0,
total_kv_head_num=4,
kv_head_num=4,
kv_layer_ids=[7, 7],
)
manager._transfer_data = Mock(return_value=0)
staging = SimpleNamespace(fits=lambda size: True, get_ptr=lambda: 0x9000)
ret = manager.send_kvcache_staged(
"peer",
np.array([1, 2], dtype=np.int32),
dst_staging_ptr=0x100000,
dst_staging_size=1 << 20,
dst_tp_rank=0,
dst_attn_tp_size=16,
dst_kv_item_len=128,
dst_layer_ids=[3, 7, 11, 3, 7, 11],
staging_buffer=staging,
)
self.assertEqual(ret, 0)
gather.assert_called_once()
manager._transfer_data.assert_called_once_with(
"peer",
[
(0x9000, 0x100000 + 64, 64),
(0x9000 + 64, 0x100000 + 4 * 64, 64),
],
)
class TestEagleDsaSeedTransfer(unittest.TestCase):
@staticmethod
def _make_req(seed, metadata_buffer_index=0):