Files
sglang/test/registered/unit/disaggregation/test_dcp_pack.py
T

372 lines
14 KiB
Python

import unittest
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock, patch
import numpy as np
import torch
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.common.dcp_pack import (
dcp_pack_buffer_bytes,
)
from sglang.srt.disaggregation.common.utils import (
build_dcp_token_transfer_plan,
group_concurrent_contiguous,
)
from sglang.srt.disaggregation.nixl.conn import NixlKVManager, NixlKVSender
from sglang.srt.disaggregation.prefill import SchedulerDisaggregationPrefillMixin
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
def _plan(*, src, dst, page_size, dcp_size, dcp_rank, **kwargs):
return build_dcp_token_transfer_plan(
np.asarray(src, dtype=np.int32),
np.asarray(dst, dtype=np.int32),
physical_page_size=page_size,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
**kwargs,
)
class TestDcpTokenTransferPlan(CustomTestCase):
def test_one_virtual_page_explicit_rows(self):
# P=2, N=4. Prefill pages 5,2,11,4; decode virtual page 7.
# pos 0..7 src rows: 10,11, 4,5, 22,23, 8,9
# draft dest page is P*N=8 → 56..63
# each rank stores local rows 14,15 (page P=2)
expected_draft_src = [10, 11, 4, 5, 22, 23, 8, 9]
expected_draft_dst = list(range(56, 64))
expected_target_src = {
0: [10, 22],
1: [11, 23],
2: [4, 8],
3: [5, 9],
}
seen_src = []
for rank, src in expected_target_src.items():
plan = _plan(
src=[5, 2, 11, 4],
dst=[7],
page_size=2,
dcp_size=4,
dcp_rank=rank,
num_kv_tokens=8,
)
np.testing.assert_array_equal(
plan.draft_src_token_indices, expected_draft_src
)
np.testing.assert_array_equal(
plan.draft_dst_token_indices, expected_draft_dst
)
np.testing.assert_array_equal(plan.target_src_token_indices, src)
np.testing.assert_array_equal(plan.target_dst_token_indices, [14, 15])
seen_src.extend(plan.target_src_token_indices.tolist())
self.assertEqual(sorted(seen_src), sorted(expected_draft_src))
def test_second_chunk_crosses_dest_pages(self):
# P=2, N=2 (virtual page = 4). Decode already holds a 4-token prefix;
# dst=[4, 6] is the full send-range page list. This chunk is the second
# prefill page of the send range (src_page_offset=1), so its 4 tokens
# sit at send-range pos 2..5 (absolute 6..9) and straddle virtual page
# 4 (rows 16..19) and virtual page 6 (rows 24..27).
plan = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=0,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
)
np.testing.assert_array_equal(plan.draft_src_token_indices, [18, 19, 6, 7])
np.testing.assert_array_equal(plan.draft_dst_token_indices, [18, 19, 24, 25])
# rank 0 owns absolute pos 6, 8 -> per-rank slots 1, 2 -> pages 4, 6.
np.testing.assert_array_equal(plan.target_src_token_indices, [18, 6])
np.testing.assert_array_equal(plan.target_dst_token_indices, [9, 12])
plan_r1 = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=1,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
)
np.testing.assert_array_equal(plan_r1.draft_src_token_indices, [18, 19, 6, 7])
np.testing.assert_array_equal(plan_r1.draft_dst_token_indices, [18, 19, 24, 25])
np.testing.assert_array_equal(plan_r1.target_src_token_indices, [19, 7])
np.testing.assert_array_equal(plan_r1.target_dst_token_indices, [9, 12])
def test_rejects_unaligned_prefix(self):
with self.assertRaisesRegex(ValueError, "align"):
_plan(
src=[0],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
decode_prefix_len=1,
num_kv_tokens=2,
)
def test_empty_tokens(self):
plan = _plan(
src=[0], dst=[0], page_size=2, dcp_size=4, dcp_rank=0, num_kv_tokens=0
)
self.assertTrue(plan.empty())
class TestPackedDcpGrouping(CustomTestCase):
def test_target_needs_pack_draft_does_not(self):
plan = _plan(
src=[0, 1, 2, 3],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
num_kv_tokens=8,
)
np.testing.assert_array_equal(plan.target_src_token_indices, [0, 4])
np.testing.assert_array_equal(plan.target_dst_token_indices, [0, 1])
target_src, _ = group_concurrent_contiguous(
plan.target_src_token_indices, plan.target_dst_token_indices
)
self.assertEqual(target_src, [[0], [4]])
packed_src, packed_dst = group_concurrent_contiguous(
np.arange(2, dtype=np.int64), plan.target_dst_token_indices
)
self.assertEqual(packed_src, [[0, 1]])
self.assertEqual(packed_dst, [[0, 1]])
draft_src, draft_dst = group_concurrent_contiguous(
plan.draft_src_token_indices, plan.draft_dst_token_indices
)
self.assertEqual(draft_src, [[0, 1, 2, 3, 4, 5, 6, 7]])
self.assertEqual(draft_dst, [[0, 1, 2, 3, 4, 5, 6, 7]])
def _dcp_kv_manager_stub(*, page_size, kv_item_lens, num_draft_entries):
return SimpleNamespace(
kv_args=SimpleNamespace(
page_size=page_size,
kv_item_lens=kv_item_lens,
num_draft_entries=num_draft_entries,
)
)
class TestPrepareDcpTokenItemLens(CustomTestCase):
def test_draft_tail_scales_by_dst_dcp_size(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 32, 64 * 16],
num_draft_entries=1,
)
token_lens = CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 32, 4 * 64 * 16], dst_dcp_size=4
)
self.assertEqual(token_lens, [32, 32, 16])
def test_rejects_unscaled_draft_item_len(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 16],
num_draft_entries=1,
)
with self.assertRaisesRegex(RuntimeError, "geometry differs at entry 1"):
CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 16], dst_dcp_size=4
)
class TestDcpCachedPrefixSend(CustomTestCase):
def test_cached_prefix_fits_pack_capacity_and_preserves_pages_and_state(self):
"""Bound DCP sends by the allocation without splitting TP sends on the same worker."""
total, page_size, token_bytes = 2055, 64, 16
mgr = object.__new__(NixlKVManager)
mgr.kv_args = SimpleNamespace(
kv_item_lens=[page_size * token_bytes],
num_draft_entries=0,
page_size=page_size,
gpu_id=0,
state_types=[StateType.MAMBA],
)
mgr._dcp_pack_buffers = None
mgr._dcp_pack_max_tokens = None
mgr.transfer_queues = [None]
mgr._register_staging_memory = Mock()
mgr.request_status = {}
mgr.is_dummy_cp_rank = False
mgr.enable_all_cp_ranks_for_transfer = False
mgr.decode_kv_args_table = {
peer: SimpleNamespace(requires_dcp_relayout=relayout)
for peer, relayout in (("dcp", True), ("tp", False))
}
def allocate(size, *args, **kwargs):
return SimpleNamespace(get_ptr=lambda: 0x1000, get_size=lambda: size)
with (
get_context().override_server_args(chunked_prefill_size=250),
patch(
"sglang.srt.disaggregation.common.staging_handler._get_custom_mem_pool",
return_value=(None, None),
),
patch(
"sglang.srt.disaggregation.common.dcp_pack.StagingBuffer",
side_effect=allocate,
),
):
mgr._init_dcp_pack_buffers_once(dcp_size=4)
limit = mgr._dcp_pack_buffers[0].get_size() // token_bytes
for peer, prefix in (("dcp", 0), ("dcp", 256), ("tp", 0), ("tp", 256)):
with self.subTest(peer=peer, decode_prefix=prefix):
mgr.transfer_infos = {
1: {
"dummy": SimpleNamespace(is_dummy=True),
peer: SimpleNamespace(is_dummy=False),
}
}
mgr.add_transfer_request = Mock()
with get_context().override_server_args(dp_size=1):
sender = NixlKVSender(mgr, "unused", 1, [0], 0)
sender.init((total - prefix + page_size - 1) // page_size, 3)
req = SimpleNamespace(
rid="cached-prefix",
kv=SimpleNamespace(req_pool_idx=0),
origin_input_ids=[0] * total,
extend_range=SimpleNamespace(end=total),
start_send_idx=prefix,
disagg_decode_prefix_len=prefix,
disagg_kv_sender=sender,
)
scheduler = SimpleNamespace(
enable_staging=False,
token_to_kv_pool_allocator=SimpleNamespace(
page_size=page_size,
translate_kv_indices_for_transfer=lambda x: x,
),
req_to_token_pool=SimpleNamespace(
req_to_token=torch.arange(total).reshape(1, -1),
req_index_to_mamba_index_mapping=torch.tensor([17]),
translate_mamba_indices=lambda x: x,
),
disagg_metadata_buffers=Mock(),
disagg_prefill_bootstrap_queue=SimpleNamespace(kv_manager=mgr),
disagg_prefill_pending_chunk_rids=set(),
)
SchedulerDisaggregationPrefillMixin._send_kv_chunk(
scheduler, req, last_chunk=True
)
calls = mgr.add_transfer_request.call_args_list
token_counts = [c.args[7] for c in calls]
if peer == "dcp":
self.assertLessEqual(max(token_counts), limit)
else:
self.assertEqual(len(calls), 1)
self.assertEqual(sum(token_counts), total - prefix)
np.testing.assert_array_equal(
np.concatenate([c.args[1] for c in calls]),
np.arange(prefix // 64, 33),
)
page_offset = 0
for call in calls:
pages = len(call.args[1])
self.assertEqual(
call.args[2], slice(page_offset, page_offset + pages)
)
page_offset += pages
self.assertEqual(
[c.args[3] for c in calls], [False] * (len(calls) - 1) + [True]
)
self.assertTrue(all(c.args[6] is None for c in calls[:-1]))
self.assertEqual(int(calls[-1].args[6][0][0]), 17)
class TestDcpPackBufferBytes(CustomTestCase):
def test_sizes_fixed_regions_for_each_dcp_rank(self):
self.assertEqual(
dcp_pack_buffer_bytes(
[64 * 16, 64 * 16],
page_size=64,
max_tokens=10,
dcp_size=4,
),
4 * 3 * (16 + 16),
)
def test_rejects_invalid_item_lens(self):
with self.assertRaisesRegex(ValueError, "at least one page"):
dcp_pack_buffer_bytes([0], page_size=64, max_tokens=8)
with self.assertRaisesRegex(ValueError, "page-aligned"):
dcp_pack_buffer_bytes([100], page_size=64, max_tokens=8)
class TestTryDcpPack(CustomTestCase):
def test_try_pack_uses_requested_region_and_dense_indices(self):
"""A gather must fit its rank region even when the total buffer has space."""
dim = 4
kv = torch.arange(16 * dim, dtype=torch.float32).view(16, 1, dim)
item_len = int(kv[0].nbytes)
pack = torch.zeros(8 * item_len, dtype=torch.uint8)
gather_stream = Mock()
buf = type(
"Buf",
(),
{
"buffer": pack,
"fits": lambda self, n: n <= pack.numel(),
"get_ptr": lambda self: 0x1000,
"get_size": lambda self: pack.numel(),
"get_gather_stream": lambda self: gather_stream,
},
)()
src = np.array([1, 5, 9, 13], dtype=np.int64)
pack_offset = 4 * item_len
mgr = object.__new__(NixlKVManager)
mgr.kv_args = SimpleNamespace(kv_data_ptrs=[kv.data_ptr()], num_draft_entries=0)
dst = SimpleNamespace(
dst_dcp_rank=1, dst_dcp_size=2, dcp_token_item_lens=[item_len]
)
with (
patch(
"sglang.srt.disaggregation.common.dcp_pack.torch.cuda.default_stream"
),
patch(
"sglang.srt.disaggregation.common.dcp_pack.torch.cuda.stream",
return_value=nullcontext(),
),
patch(
"sglang.srt.disaggregation.common.dcp_pack.copy_mla_rows_into_pack"
) as copy_mock,
):
packed = mgr._pack_dcp_rank_once(buf, dst, src, {})
dst.dst_dcp_size = 4
overflow = mgr._pack_dcp_rank_once(buf, dst, src, {})
self.assertIsNone(overflow)
gather_stream.synchronize.assert_called_once_with()
self.assertIsNotNone(packed)
ptrs, indices = packed
self.assertEqual(ptrs, [0x1000 + pack_offset])
np.testing.assert_array_equal(indices, np.arange(4))
pack_view = copy_mock.call_args.args[2]
self.assertEqual(pack_view.storage_offset(), pack_offset)
self.assertEqual(pack_view.numel(), src.size * item_len)
if __name__ == "__main__":
unittest.main()