[PD] Support pipeline-parallel prefill with Mooncake staging buffer (#33807)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user