[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)
|
||||
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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user