diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index b095d7c3e..1ab1562fb 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -791,12 +791,33 @@ class NixlKVManager(CommonKVManager): else: # One prefill rank feeds multiple decode ranks: interleave num_groups # 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 - num_heads_to_send = dst_heads_per_rank - src_head_start = ( - dst_tp_rank_in_group * dst_heads_per_rank - ) % src_heads_per_rank + # + # Under GQA the decode side can have MORE attn-TP ranks than there are + # KV heads (decode_tp_size > total_kv_heads). In that case consecutive + # decode ranks replicate a shared KV head, so the src dlist must + # interleave one group per UNIQUE source head-slice, not one per decode + # 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 dst_head_offset = 0 diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index f97ef2faa..67e727f41 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -1019,5 +1019,115 @@ class TestNixlStaging(CustomTestCase): 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__": unittest.main()