[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) response = _get_bootstrap_session(self.bootstrap_addr).get(url, timeout=5)
if response.status_code == 200: if response.status_code == 200:
bootstrap_info = response.json() bootstrap_info = response.json()
bootstrap_info["pp_rank"] = int(target_pp_rank)
return bootstrap_info return bootstrap_info
else: else:
logger.error( logger.error(
@@ -107,11 +107,15 @@ class DecodeStagingHandler:
self._wm_subscribers[key] = (receiver, session_id) self._wm_subscribers[key] = (receiver, session_id)
def num_writers_for(self, receiver) -> int: def num_writers_for(self, receiver) -> int:
"""Compute num_writers for a specific request based on its prefill TP.""" """Compute all TP and PP writers expected for a staging chunk."""
prefill_tp = receiver.prefill_info.attn_tp_size prefill_info = receiver.prefill_info
prefill_tp = prefill_info.attn_tp_size
if prefill_tp > self.decode_tp: if prefill_tp > self.decode_tp:
return prefill_tp // max(1, self.decode_tp) tp_writers = prefill_tp // max(1, self.decode_tp)
return 1 else:
tp_writers = 1
pp_writers = prefill_info.pp_size // self.kv_manager.pp_size
return tp_writers * pp_writers
@classmethod @classmethod
def create(cls, kv_manager, scheduler, tp_rank: int) -> DecodeStagingHandler: def create(cls, kv_manager, scheduler, tp_rank: int) -> DecodeStagingHandler:
@@ -617,6 +621,7 @@ class PrefillStagingStrategy:
target_info.dst_tp_rank, target_info.dst_tp_rank,
target_info.dst_attn_tp_size, target_info.dst_attn_tp_size,
target_info.dst_kv_item_len, target_info.dst_kv_item_len,
target_info.dst_kv_layer_ids,
staging_buffer=self.staging_buffer, staging_buffer=self.staging_buffer,
) )
except Exception as e: except Exception as e:
@@ -716,6 +721,7 @@ def handle_staging_req(
chunk_idx = int(msg[2].decode("ascii")) chunk_idx = int(msg[2].decode("ascii"))
chunk_num_pages = int(msg[3].decode("ascii")) chunk_num_pages = int(msg[3].decode("ascii"))
session_id = msg[4].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: if staging_allocator is None:
logger.warning( logger.warning(
@@ -798,6 +804,8 @@ def handle_staging_req(
bootstrap_infos = room_bootstrap.get(room) bootstrap_infos = room_bootstrap.get(room)
if bootstrap_infos: if bootstrap_infos:
for bi in bootstrap_infos: for bi in bootstrap_infos:
if requester_pp_rank is not None and bi["pp_rank"] != requester_pp_rank:
continue
try: try:
sock, lock = receiver._connect_to_bootstrap_server(bi) sock, lock = receiver._connect_to_bootstrap_server(bi)
with lock: with lock:
@@ -823,6 +831,7 @@ def prefetch_staging_reqs(
chunked_prefill_size: int, chunked_prefill_size: int,
staging_requested: set, staging_requested: set,
prefetch_sockets: dict, prefetch_sockets: dict,
requester_pp_rank: Optional[int] = None,
) -> None: ) -> None:
"""Send STAGING_REQ for all chunks before the prefill forward starts. """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.setsockopt(zmq.IPV6, 1)
sock.connect(ep) sock.connect(ep)
prefetch_sockets[ep] = sock prefetch_sockets[ep] = sock
prefetch_sockets[ep].send_multipart( request = [
[
b"STAGING_REQ", b"STAGING_REQ",
str(room).encode("ascii"), str(room).encode("ascii"),
str(chunk_idx).encode("ascii"), str(chunk_idx).encode("ascii"),
str(chunk_pages).encode("ascii"), str(chunk_pages).encode("ascii"),
session_id.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: except Exception:
staging_requested.discard(stg_key) 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) per_rank_kv_heads = getattr(kv_pool_for_heads, "head_num", 0)
if per_rank_kv_heads > 0: if per_rank_kv_heads > 0:
kv_args.kv_head_num = per_rank_kv_heads 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"): if hasattr(kv_manager, "set_kv_buffer_tensors"):
kv_pool = kv_pool_for_heads kv_pool = kv_pool_for_heads
if hasattr(kv_pool, "k_buffer") and hasattr(kv_pool, "v_buffer"): if hasattr(kv_pool, "k_buffer") and hasattr(kv_pool, "v_buffer"):
@@ -520,6 +520,7 @@ class MooncakeKVManager(CommonKVManager):
get_schedule().chunked_prefill_size, get_schedule().chunked_prefill_size,
self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_requested,
self._staging_ctx.prefetch_sockets, self._staging_ctx.prefetch_sockets,
requester_pp_rank=self.pp_rank,
) )
def send_kvcache_staged( def send_kvcache_staged(
@@ -531,6 +532,7 @@ class MooncakeKVManager(CommonKVManager):
dst_tp_rank: int, dst_tp_rank: int,
dst_attn_tp_size: int, dst_attn_tp_size: int,
dst_kv_item_len: int, dst_kv_item_len: int,
dst_layer_ids: List[int],
staging_buffer=None, staging_buffer=None,
) -> int: ) -> int:
"""Transfer KV cache via staging buffers (gather -> bulk RDMA -> scatter on decode).""" """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 num_tokens = len(prefill_kv_indices) * page_size
per_layer_bytes = num_tokens * num_heads_to_send * head_dim * dtype_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( num_writers, writer_rank_bytes, total_staging_needed = compute_staging_layout(
self.attn_tp_size, self.attn_tp_size,
@@ -572,21 +586,20 @@ class MooncakeKVManager(CommonKVManager):
total_kv_heads, total_kv_heads,
num_tokens, num_tokens,
head_dim * dtype_size, head_dim * dtype_size,
num_layers, dst_num_layers,
) )
writer_idx = local_tp_rank % num_writers if num_writers > 1 else 0 writer_idx = local_tp_rank % num_writers if num_writers > 1 else 0
rank_offset = sum(writer_rank_bytes[:writer_idx]) 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( 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 return -1
if dst_staging_size < total_staging_needed: if dst_staging_size < total_staging_needed:
logger.warning( logger.warning(
f"Decode staging too small: need {total_staging_needed} bytes " 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"for {dst_num_layers} layers, have {dst_staging_size}, falling back"
f"x {per_rank_bytes} bytes/rank), have {dst_staging_size}, falling back"
) )
return -1 return -1
@@ -605,16 +618,29 @@ class MooncakeKVManager(CommonKVManager):
self.kv_args.gpu_id, self.kv_args.gpu_id,
) )
dst_write_ptr = dst_staging_ptr + rank_offset if pairs is None:
ret = self._transfer_data( transfer_blocks = [
mooncake_session_id, (
[(staging_buffer.get_ptr(), dst_write_ptr, per_rank_bytes)], 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: if ret != 0:
raise RuntimeError( raise RuntimeError(
f"[Staging] Bulk RDMA transfer failed with ret={ret}. " f"[Staging] Bulk RDMA transfer failed with ret={ret}. "
f"src_ptr=0x{staging_buffer.get_ptr():x}, " 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." f"The decode staging buffer may not be properly registered."
) )
return ret return ret
+3 -3
View File
@@ -175,10 +175,10 @@ class PrefillBootstrapQueue:
f"chunked_prefill_size that is a multiple of page_size " f"chunked_prefill_size that is a multiple of page_size "
f"({page_size}); got {chunked_prefill_size}." f"({page_size}); got {chunked_prefill_size}."
) )
if self.pp_size > 1: if self.pp_size > 1 and self.transfer_backend != TransferBackend.MOONCAKE:
# Staging writer accounting has no pp dimension.
raise RuntimeError( 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: if get_parallel().enable_prefill_context_parallel:
# CP rewrites index_slice per rank, breaking the chunk grid. # 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) self._pp_commit_comm_work(self.send_proxy_work)
if cur_batch: if cur_batch:
if self.enable_staging:
self.maybe_prefetch_staging_for_batch(cur_batch)
result, self.launch_event = self._pp_launch_batch( result, self.launch_event = self._pp_launch_batch(
mb_id, mb_id,
cur_batch, cur_batch,
@@ -1,12 +1,16 @@
import struct import struct
import threading
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import Mock, patch
import numpy as np import numpy as np
import torch import torch
from sglang.srt.disaggregation.base.conn import KVArgs, StateType 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 ( from sglang.srt.disaggregation.common.utils import (
group_concurrent_contiguous, group_concurrent_contiguous,
pack_int_lists, pack_int_lists,
@@ -15,7 +19,8 @@ from sglang.srt.disaggregation.common.utils import (
unpack_list_of_buffers, unpack_list_of_buffers,
) )
from sglang.srt.disaggregation.mooncake.conn import ( from sglang.srt.disaggregation.mooncake.conn import (
KVArgsRegisterInfo as MooncakeKVArgsRegisterInfo, KVArgsRegisterInfo,
MooncakeKVManager,
) )
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
MetadataBuffers, MetadataBuffers,
@@ -57,7 +62,7 @@ class TestDisaggregationWire(unittest.TestCase):
b"2", b"2",
] ]
info = MooncakeKVArgsRegisterInfo.from_zmq(msg) info = KVArgsRegisterInfo.from_zmq(msg)
self.assertEqual(info.staging_base_ptr, 0x3000) self.assertEqual(info.staging_base_ptr, 0x3000)
self.assertEqual(info.staging_total_size, 4096) 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])) 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): class TestEagleDsaSeedTransfer(unittest.TestCase):
@staticmethod @staticmethod
def _make_req(seed, metadata_buffer_index=0): def _make_req(seed, metadata_buffer_index=0):