From 667e18d99d961044293fceea15a815962512c69d Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Mon, 10 Aug 2026 22:53:17 -0700 Subject: [PATCH] [PD] Support pipeline-parallel prefill with Mooncake staging buffer (#33807) Co-authored-by: Shangming Cai --- .../sglang/srt/disaggregation/common/conn.py | 1 + .../disaggregation/common/staging_handler.py | 36 +++++--- python/sglang/srt/disaggregation/decode.py | 4 +- .../srt/disaggregation/mooncake/conn.py | 50 ++++++++--- python/sglang/srt/disaggregation/prefill.py | 6 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 2 + .../test_disaggregation_wire.py | 89 ++++++++++++++++++- 7 files changed, 156 insertions(+), 32 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 636d2ae0e..3b19529cd 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -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( diff --git a/python/sglang/srt/disaggregation/common/staging_handler.py b/python/sglang/srt/disaggregation/common/staging_handler.py index 0b7860210..bccfac37d 100644 --- a/python/sglang/srt/disaggregation/common/staging_handler.py +++ b/python/sglang/srt/disaggregation/common/staging_handler.py @@ -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) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 232e266a3..c2c009b81 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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"): diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 255ff045f..66ac63b34 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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 diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 09ba12c24..2a15fa758 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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. diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 715e20ebf..42ead03e0 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -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, diff --git a/test/registered/unit/disaggregation/test_disaggregation_wire.py b/test/registered/unit/disaggregation/test_disaggregation_wire.py index 0b4639661..483dcdaf6 100644 --- a/test/registered/unit/disaggregation/test_disaggregation_wire.py +++ b/test/registered/unit/disaggregation/test_disaggregation_wire.py @@ -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):