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:
laoyao0822
2026-05-02 07:07:28 +08:00
parent 2317952a01
commit 5769b63082
17 changed files with 2524 additions and 96 deletions
@@ -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()