[Disagg][NIXL] Fix heterogeneous attn-TP KV transfer for replicated GQA heads (NIXL_ERR_NOT_FOUND) (#31968)
This commit is contained in:
@@ -791,12 +791,33 @@ class NixlKVManager(CommonKVManager):
|
|||||||
else:
|
else:
|
||||||
# One prefill rank feeds multiple decode ranks: interleave num_groups
|
# One prefill rank feeds multiple decode ranks: interleave num_groups
|
||||||
# head-groups in the src dlist so each decode rank picks its slice.
|
# head-groups in the src dlist so each decode rank picks its slice.
|
||||||
dst_tp_rank_in_group = decode_kv_args.decode_tp_rank % decode_tp_size
|
#
|
||||||
num_groups = decode_tp_size // prefill_tp_size
|
# Under GQA the decode side can have MORE attn-TP ranks than there are
|
||||||
num_heads_to_send = dst_heads_per_rank
|
# KV heads (decode_tp_size > total_kv_heads). In that case consecutive
|
||||||
src_head_start = (
|
# decode ranks replicate a shared KV head, so the src dlist must
|
||||||
dst_tp_rank_in_group * dst_heads_per_rank
|
# interleave one group per UNIQUE source head-slice, not one per decode
|
||||||
) % src_heads_per_rank
|
# rank -- otherwise it addresses past the registered KV region and
|
||||||
|
# prep_xfer_dlist raises NIXL_ERR_NOT_FOUND.
|
||||||
|
#
|
||||||
|
# Reuse the shared replicated-KV head map (integer division under
|
||||||
|
# replication, not modulo) that the mooncake backend already relies
|
||||||
|
# on, so the two backends stay in sync.
|
||||||
|
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||||
|
compute_head_slice_params,
|
||||||
|
)
|
||||||
|
|
||||||
|
src_head_start, num_heads_to_send, _, _ = compute_head_slice_params(
|
||||||
|
prefill_tp_size,
|
||||||
|
decode_tp_size,
|
||||||
|
self.kv_args.engine_rank,
|
||||||
|
decode_kv_args.decode_tp_rank,
|
||||||
|
total_kv_heads,
|
||||||
|
)
|
||||||
|
# num_groups (distinct head-groups packed in one prefill rank's src
|
||||||
|
# region) and head_group_idx (this peer's group) are NIXL-specific and
|
||||||
|
# not returned by the shared helper, so derive them here.
|
||||||
|
dst_replication = max(1, decode_tp_size // total_kv_heads)
|
||||||
|
num_groups = decode_tp_size // prefill_tp_size // dst_replication
|
||||||
head_group_idx = src_head_start // dst_heads_per_rank
|
head_group_idx = src_head_start // dst_heads_per_rank
|
||||||
dst_head_offset = 0
|
dst_head_offset = 0
|
||||||
|
|
||||||
|
|||||||
@@ -1019,5 +1019,115 @@ class TestNixlStaging(CustomTestCase):
|
|||||||
self.assertIsNone(handle)
|
self.assertIsNone(handle)
|
||||||
|
|
||||||
|
|
||||||
|
class DlistCaptureAgent:
|
||||||
|
"""Records prep_xfer_dlist descriptor arrays so tests can inspect them."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.calls = [] # (peer_name, np.ndarray, mem_kind)
|
||||||
|
|
||||||
|
def prep_xfer_dlist(self, peer_name, array, mem_kind):
|
||||||
|
self.calls.append((peer_name, np.asarray(array), mem_kind))
|
||||||
|
return f"handle_{len(self.calls)}"
|
||||||
|
|
||||||
|
|
||||||
|
class TestNixlHeteroTpReplicatedKV(CustomTestCase):
|
||||||
|
"""Regression guard for #31295.
|
||||||
|
|
||||||
|
Prefill attention-TP1 -> decode TP4 on a model with only 2 KV heads forces
|
||||||
|
GQA replication: decode ranks 0,1 share KV head 0 and ranks 2,3 share KV
|
||||||
|
head 1. The shared source dlist must interleave one group per *unique*
|
||||||
|
source head-slice (2), and each peer's head_group_idx must map replicated
|
||||||
|
decode ranks via integer division (0,0,1,1). The pre-fix code used
|
||||||
|
``num_groups = decode_tp // prefill_tp`` (=4) -- addressing 2x past the
|
||||||
|
registered source region, which NIXL rejects with NIXL_ERR_NOT_FOUND -- and
|
||||||
|
a modulo head map (0,1,0,1).
|
||||||
|
"""
|
||||||
|
|
||||||
|
TOTAL_KV_HEADS = 2
|
||||||
|
DECODE_TP = 4
|
||||||
|
PAGE_SIZE = 1
|
||||||
|
BYTES_PER_HEAD = 128 # per token, per head slice
|
||||||
|
SRC_KV_ITEM_LEN = TOTAL_KV_HEADS * BYTES_PER_HEAD # both heads on one prefill rank
|
||||||
|
DST_KV_ITEM_LEN = BYTES_PER_HEAD # one replicated head per decode rank
|
||||||
|
NUM_SLOTS = 4
|
||||||
|
SRC_PTRS = [0x10000, 0x20000] # K, V for the single local layer
|
||||||
|
REGION_LEN = NUM_SLOTS * SRC_KV_ITEM_LEN
|
||||||
|
|
||||||
|
def _make_manager(self):
|
||||||
|
mgr = object.__new__(NixlKVManager)
|
||||||
|
mgr.agent = DlistCaptureAgent()
|
||||||
|
mgr.attn_tp_size = 1 # prefill attention TP = 1 (DP attention)
|
||||||
|
mgr.prep_handle_slice_src = None
|
||||||
|
mgr.prep_handles_slice_dst = {}
|
||||||
|
mgr.kv_args = SimpleNamespace(
|
||||||
|
gpu_id=0,
|
||||||
|
engine_rank=0,
|
||||||
|
page_size=self.PAGE_SIZE,
|
||||||
|
prefill_start_layer=0,
|
||||||
|
total_kv_head_num=self.TOTAL_KV_HEADS,
|
||||||
|
kv_head_num=self.TOTAL_KV_HEADS,
|
||||||
|
kv_item_lens=[self.SRC_KV_ITEM_LEN, self.SRC_KV_ITEM_LEN],
|
||||||
|
kv_data_ptrs=list(self.SRC_PTRS),
|
||||||
|
kv_data_lens=[self.REGION_LEN, self.REGION_LEN],
|
||||||
|
)
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def _decode_args(self, decode_tp_rank):
|
||||||
|
return SimpleNamespace(
|
||||||
|
agent_name=f"decode_{decode_tp_rank}",
|
||||||
|
decode_tp_size=self.DECODE_TP,
|
||||||
|
decode_tp_rank=decode_tp_rank,
|
||||||
|
dst_kv_item_len=self.DST_KV_ITEM_LEN,
|
||||||
|
dst_kv_ptrs=[0x30000, 0x40000],
|
||||||
|
dst_num_slots=self.NUM_SLOTS,
|
||||||
|
gpu_id=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_src_dlist_stays_within_registered_region_and_num_groups(self):
|
||||||
|
# Src dlist is built once (shared across peers) on the first call.
|
||||||
|
mgr = self._make_manager()
|
||||||
|
mgr._init_hetero_tp_prep_handle(
|
||||||
|
peer_name="decode_0", decode_kv_args=self._decode_args(0)
|
||||||
|
)
|
||||||
|
|
||||||
|
# num_groups must be 2 (one per unique KV head), not decode_tp//prefill_tp=4.
|
||||||
|
src_handle, num_groups, _num_ptr_pairs, _num_slots = mgr.prep_handle_slice_src
|
||||||
|
self.assertEqual(num_groups, 2)
|
||||||
|
|
||||||
|
# Every source descriptor [addr, addr+len) must lie inside a registered
|
||||||
|
# base region [ptr, ptr+REGION_LEN). Pre-fix, num_groups=4 pushed the
|
||||||
|
# top group's addresses past the region -> NIXL_ERR_NOT_FOUND.
|
||||||
|
src_call = next(c for c in mgr.agent.calls if c[0] == "")
|
||||||
|
src_array = src_call[1]
|
||||||
|
regions = [(p, p + self.REGION_LEN) for p in self.SRC_PTRS]
|
||||||
|
for addr, length, _dev in src_array:
|
||||||
|
addr = int(addr)
|
||||||
|
length = int(length)
|
||||||
|
self.assertTrue(
|
||||||
|
any(lo <= addr and addr + length <= hi for lo, hi in regions),
|
||||||
|
f"descriptor [{addr:#x}, {addr + length:#x}) escapes all "
|
||||||
|
f"registered source regions {[(hex(lo), hex(hi)) for lo, hi in regions]}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_head_group_idx_maps_replicated_ranks_by_integer_division(self):
|
||||||
|
# Each decode rank's per-peer dst handle records its head_group_idx.
|
||||||
|
# Expected replicated-KV mapping: ranks 0,1 -> group 0; ranks 2,3 -> group 1.
|
||||||
|
expected = {0: 0, 1: 0, 2: 1, 3: 1}
|
||||||
|
for rank in range(self.DECODE_TP):
|
||||||
|
mgr = self._make_manager()
|
||||||
|
mgr._init_hetero_tp_prep_handle(
|
||||||
|
peer_name=f"decode_{rank}", decode_kv_args=self._decode_args(rank)
|
||||||
|
)
|
||||||
|
_dst_handle, _num_slots_dst, head_group_idx = mgr.prep_handles_slice_dst[
|
||||||
|
f"decode_{rank}"
|
||||||
|
]
|
||||||
|
self.assertEqual(
|
||||||
|
head_group_idx,
|
||||||
|
expected[rank],
|
||||||
|
f"decode rank {rank} mapped to group {head_group_idx}, "
|
||||||
|
f"expected {expected[rank]} (modulo bug gives 0,1,0,1)",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user