Reduce CP shared KV materialize and direct-write overhead
Shared KV now relies on page-aligned CP metadata and compute-owner page allocation so persistent MLA KV and NSA index shards can be written by the rank that computed them. The compatibility read path keeps the dense full-view contract for existing topk and attention kernels, but removes duplicated prev/next index materialize, adds optional tai materialize integration, and tightens tests/docs around the fallback boundaries. Constraint: Decode remains non-CP while prefill CP owns the shared-KV changes Constraint: Existing attention/topk kernels still expect dense full-view KV/index inputs Rejected: Change attention kernels to read owner-sharded KV directly | larger semantic change reserved for later phases Rejected: Merge index K/scale storage with MLA KV storage | would couple topk and attention cache lifecycles before materialize overhead is isolated Confidence: medium Scope-risk: broad Directive: Do not remove fallback logging or debug-gated assertions without reproducing long-context chunked/radix-hit paths Tested: git diff --check --cached Not-tested: Local pytest/runtime server verification not run in this commit step per current workflow constraints
This commit is contained in:
@@ -52,6 +52,27 @@ class TestCpSharedKVLayout(unittest.TestCase):
|
||||
|
||||
|
||||
class TestCPSharedPagedAllocator(unittest.TestCase):
|
||||
def test_compute_owner_alloc_fallback_logs_every_event(self):
|
||||
from sglang.srt.mem_cache import common
|
||||
|
||||
with self.assertLogs("sglang.srt.mem_cache.common", level="INFO") as cm:
|
||||
common._log_cp_shared_kv_alloc_fallback(
|
||||
"too_short_for_page_aligned",
|
||||
"falling back for test event %s",
|
||||
1,
|
||||
)
|
||||
common._log_cp_shared_kv_alloc_fallback(
|
||||
"too_short_for_page_aligned",
|
||||
"falling back for test event %s",
|
||||
2,
|
||||
)
|
||||
|
||||
self.assertEqual(len(cm.output), 2)
|
||||
self.assertIn("too_short_for_page_aligned", cm.output[0])
|
||||
self.assertIn("test event 1", cm.output[0])
|
||||
self.assertIn("too_short_for_page_aligned", cm.output[1])
|
||||
self.assertIn("test event 2", cm.output[1])
|
||||
|
||||
def test_shared_allocator_exposes_logical_capacity(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
|
||||
@@ -73,6 +94,229 @@ class TestCPSharedPagedAllocator(unittest.TestCase):
|
||||
allocator.free(locs)
|
||||
self.assertEqual(allocator.available_size(), 64 * 8)
|
||||
|
||||
def test_compute_owner_page_assignment_matches_in_seq_zigzag(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 16,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
owners,
|
||||
[0, 0, 1, 1, 2, 2, 3, 3, 3, 3, 2, 2, 1, 1, 0, 0],
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_falls_back_for_short_extend(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 7,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertIsNone(owners)
|
||||
self.assertEqual(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 7,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
),
|
||||
"too_short_for_page_aligned",
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_allows_radix_hit_suffix_with_one_page_per_rank(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 4,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(owners, [0, 1, 2, 3])
|
||||
self.assertIsNone(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 4,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_allows_short_radix_hit_suffix_with_replicated_compute(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 3,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(owners, [0, 1, 2])
|
||||
self.assertIsNone(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 3,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_reports_prefix_misalignment(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 16,
|
||||
extend_prefix_len=1,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
),
|
||||
"prefix_not_page_aligned",
|
||||
)
|
||||
|
||||
def test_shared_allocator_can_allocate_pages_from_compute_owner_lanes(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
)
|
||||
|
||||
page_size = 64
|
||||
cp_size = 4
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=page_size * 16,
|
||||
extend_prefix_len=0,
|
||||
page_size=page_size,
|
||||
cp_size=cp_size,
|
||||
)
|
||||
allocator = CPSharedPagedTokenToKVPoolAllocator(
|
||||
logical_size=page_size * 32,
|
||||
physical_size=page_size * 8,
|
||||
page_size=page_size,
|
||||
dtype=torch.bfloat16,
|
||||
device="cpu",
|
||||
kvcache=None,
|
||||
need_sort=False,
|
||||
cp_size=cp_size,
|
||||
cp_rank=0,
|
||||
)
|
||||
|
||||
locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 16], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 16], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 16,
|
||||
page_compute_owners=owners,
|
||||
)
|
||||
|
||||
self.assertIsNotNone(locs)
|
||||
logical_pages = locs.view(-1, page_size)[:, 0] // page_size
|
||||
self.assertEqual(
|
||||
((logical_pages - 1) % cp_size).tolist(),
|
||||
owners,
|
||||
)
|
||||
self.assertEqual(allocator.available_size(), page_size * 16)
|
||||
|
||||
def test_compute_owner_lane_eviction_recovers_exhausted_owner_lane(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
|
||||
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
|
||||
|
||||
page_size = 64
|
||||
cp_size = 4
|
||||
allocator = CPSharedPagedTokenToKVPoolAllocator(
|
||||
logical_size=page_size * 16,
|
||||
physical_size=page_size * 4,
|
||||
page_size=page_size,
|
||||
dtype=torch.bfloat16,
|
||||
device="cpu",
|
||||
kvcache=None,
|
||||
need_sort=False,
|
||||
cp_size=cp_size,
|
||||
cp_rank=0,
|
||||
)
|
||||
lane0_locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
self.assertIsNone(
|
||||
allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
)
|
||||
|
||||
class FakeTreeCache:
|
||||
def __init__(self):
|
||||
self.evict_calls = []
|
||||
|
||||
def is_chunk_cache(self):
|
||||
return False
|
||||
|
||||
def evictable_size(self):
|
||||
return lane0_locs.numel()
|
||||
|
||||
def evict(self, params):
|
||||
self.evict_calls.append(params.num_tokens)
|
||||
allocator.free(lane0_locs)
|
||||
return EvictResult(num_tokens_evicted=lane0_locs.numel())
|
||||
|
||||
tree_cache = FakeTreeCache()
|
||||
_evict_for_compute_owner_lanes(
|
||||
tree_cache=tree_cache,
|
||||
allocator=allocator,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
|
||||
self.assertGreaterEqual(len(tree_cache.evict_calls), 1)
|
||||
locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
self.assertIsNotNone(locs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -612,5 +612,224 @@ class TestCpSharedKVLazyDebugLogging(unittest.TestCase):
|
||||
self.assertEqual(key_to_write.shape[0], 2)
|
||||
|
||||
|
||||
class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
|
||||
def test_paged_materialize_uses_tai_kernel_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def materialize_shared_pages(
|
||||
self,
|
||||
page_buffer,
|
||||
logical_pages,
|
||||
*,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
self.calls.append((page_buffer, logical_pages, cp_rank, cp_size))
|
||||
return (
|
||||
torch.full((logical_pages.numel() + 1, 3), 7, dtype=torch.uint8),
|
||||
torch.tensor([1, 0, 3], dtype=logical_pages.dtype),
|
||||
)
|
||||
|
||||
fake_tai = FakeTaiKernels()
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
page_buffer = torch.arange(0, 5 * 3, dtype=torch.uint8).view(5, 3)
|
||||
logical_pages = torch.tensor([1, 0, 3], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertEqual(len(fake_tai.calls), 1)
|
||||
self.assertIs(fake_tai.calls[0][0], page_buffer)
|
||||
self.assertIs(fake_tai.calls[0][1], logical_pages)
|
||||
self.assertEqual(fake_tai.calls[0][2:], (1, 2))
|
||||
self.assertEqual(dense_pages.tolist(), [1, 0, 3])
|
||||
self.assertEqual(int(dense_page_buffer.sum().item()), 7 * 4 * 3)
|
||||
|
||||
def test_token_materialize_uses_tai_kernel_for_slot_remap_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def __init__(self):
|
||||
self.page_inverse_calls = []
|
||||
self.remap_calls = []
|
||||
self.token_calls = []
|
||||
|
||||
def build_slot_page_inverse(self, slot_logical_pages, logical_page_capacity):
|
||||
self.page_inverse_calls.append((slot_logical_pages, logical_page_capacity))
|
||||
return torch.tensor([0, 1, 2, -1, 3], dtype=torch.long)
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs(
|
||||
self,
|
||||
logical_locs,
|
||||
page_inverse,
|
||||
*,
|
||||
page_size,
|
||||
):
|
||||
self.remap_calls.append((logical_locs, page_inverse, page_size))
|
||||
return torch.tensor([4, 8, -1], dtype=logical_locs.dtype)
|
||||
|
||||
def materialize_shared_token_kv_pages(
|
||||
self,
|
||||
kv_cache,
|
||||
slot_logical_pages,
|
||||
*,
|
||||
page_size,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
self.token_calls.append(
|
||||
(kv_cache, slot_logical_pages, page_size, cp_rank, cp_size)
|
||||
)
|
||||
return torch.full((16, 1, 1), 3.0, dtype=kv_cache.dtype)
|
||||
|
||||
fake_tai = FakeTaiKernels()
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
kv_cache = torch.arange(0, 24, dtype=torch.float32).view(24, 1, 1)
|
||||
logical_locs = torch.tensor([4, 8, -1], dtype=torch.int64)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 4]], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(len(fake_tai.page_inverse_calls), 1)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
fake_tai.page_inverse_calls[0][0],
|
||||
remap_logical_pages.reshape(-1),
|
||||
)
|
||||
)
|
||||
self.assertEqual(len(fake_tai.remap_calls), 1)
|
||||
self.assertEqual(len(fake_tai.token_calls), 1)
|
||||
self.assertIs(fake_tai.token_calls[0][0], kv_cache)
|
||||
self.assertTrue(
|
||||
torch.equal(fake_tai.token_calls[0][1], remap_logical_pages.reshape(-1))
|
||||
)
|
||||
self.assertEqual(fake_tai.token_calls[0][2:], (4, 1, 2))
|
||||
self.assertEqual(dense_locs.tolist(), [4, 8, -1])
|
||||
self.assertEqual(float(dense_kv.sum().item()), 48.0)
|
||||
|
||||
def test_token_tai_path_skips_torch_slot_page_remap(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def build_slot_page_inverse(self, slot_logical_pages, logical_page_capacity):
|
||||
return torch.tensor([0, 1, 2, -1, 3], dtype=torch.long)
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs(
|
||||
self,
|
||||
logical_locs,
|
||||
page_inverse,
|
||||
*,
|
||||
page_size,
|
||||
):
|
||||
return torch.tensor([4, 8, -1], dtype=logical_locs.dtype)
|
||||
|
||||
def materialize_shared_token_kv_pages(
|
||||
self,
|
||||
kv_cache,
|
||||
slot_logical_pages,
|
||||
*,
|
||||
page_size,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
return torch.full((16, 1, 1), 3.0, dtype=kv_cache.dtype)
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
kv_cache = torch.arange(0, 24, dtype=torch.float32).view(24, 1, 1)
|
||||
logical_locs = torch.tensor([4, 8, -1], dtype=torch.int64)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 4]], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=FakeTaiKernels()
|
||||
), patch.object(
|
||||
runtime,
|
||||
"build_slot_page_remap",
|
||||
side_effect=AssertionError("tai token path must not run torch remap"),
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_locs.tolist(), [4, 8, -1])
|
||||
self.assertEqual(float(dense_kv.sum().item()), 48.0)
|
||||
|
||||
def test_tai_materialize_is_not_used_when_debug_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
page_buffer = torch.arange(0, 4 * 3, dtype=torch.uint8).view(4, 3)
|
||||
logical_pages = torch.tensor([1, 2], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=True
|
||||
), patch.object(
|
||||
runtime,
|
||||
"_load_tai_materialize_kernels",
|
||||
side_effect=AssertionError("tai path must stay off in debug mode"),
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_pages.tolist(), [1, 2])
|
||||
self.assertTrue(torch.equal(dense_page_buffer[1], page_buffer[1]))
|
||||
self.assertTrue(torch.equal(dense_page_buffer[2], page_buffer[2]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user