372 lines
14 KiB
Python
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()
|