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
@@ -1,11 +1,18 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.layers.attention.nsa.utils import (
NSAContextParallelMetadata,
_get_in_seq_last_token_owner_and_offset,
build_page_aligned_in_seq_split_list,
build_token_balanced_in_seq_split_list,
can_cp_split,
cp_split_and_rebuild_1d,
get_cp_shared_kv_local_out_cache_loc,
split_in_seq_cp_local_pair,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
@@ -93,6 +100,85 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertFalse(info.page_aligned)
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(512, 8))
def test_page_aligned_split_allows_radix_hit_suffix_with_one_page_per_rank(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=512,
extend_len=512,
extend_prefix_len=54464,
page_size=64,
cp_size=8,
)
self.assertTrue(info.page_aligned)
self.assertEqual(sum(split_list), 512)
self.assertEqual(split_list[:8], [64] * 8)
self.assertEqual(split_list[8:], [0] * 8)
self.assert_page_aligned_boundaries(
split_list, extend_prefix_len=54464, extend_len=512, page_size=64
)
def test_page_aligned_split_falls_back_when_radix_hit_suffix_has_zero_rank(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=256,
extend_len=256,
extend_prefix_len=54464,
page_size=64,
cp_size=8,
)
self.assertFalse(info.page_aligned)
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(256, 8))
def test_can_cp_split_uses_replicated_compute_for_short_radix_hit_suffix(self):
class Mode:
def is_context_parallel_extend(self):
return True
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
extend_seq_lens_cpu=[256],
extend_prefix_lens_cpu=[54464],
token_to_kv_pool=SimpleNamespace(page_size=64),
forward_mode=Mode(),
)
with (
patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
),
patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp",
return_value=True,
),
):
self.assertFalse(can_cp_split(256, 8, True, forward_batch))
def test_can_cp_split_keeps_cp_for_radix_hit_suffix_with_one_page_per_rank(self):
class Mode:
def is_context_parallel_extend(self):
return True
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
extend_seq_lens_cpu=[512],
extend_prefix_lens_cpu=[54464],
token_to_kv_pool=SimpleNamespace(page_size=64),
forward_mode=Mode(),
)
with (
patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
),
patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp",
return_value=True,
),
):
self.assertTrue(can_cp_split(512, 8, True, forward_batch))
def test_page_aligned_split_adds_padding_tokens_to_last_segment(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=1040,
@@ -152,6 +238,309 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
with self.assertRaisesRegex(RuntimeError, "local in-seq CP length mismatch"):
split_in_seq_cp_local_pair(torch.arange(9), 5, 5, name="q_fp8")
def test_cp_split_and_rebuild_1d_matches_in_seq_zigzag_order(self):
import torch
from types import SimpleNamespace
forward_batch = SimpleNamespace(
nsa_cp_metadata=NSAContextParallelMetadata(
split_list=[2, 2, 2, 2, 2, 2, 2, 2],
zigzag_index=[1, 6],
)
)
local_locs = cp_split_and_rebuild_1d(forward_batch, torch.arange(16))
self.assertEqual(local_locs.tolist(), [2, 3, 12, 13])
def test_local_out_cache_loc_requires_compute_owner_pages(self):
import torch
from types import SimpleNamespace
page_size = 4
# Segment order for cp_size=4, cp_rank=1 is segment 1 then 6.
# The logical page ids below deliberately encode the same owners through
# (logical_page - 1) % cp_size:
# segment 1 -> page 2 owner 1
# segment 6 -> page 6 owner 1
segment_pages = [1, 2, 3, 4, 8, 7, 6, 5]
out_cache_loc = torch.cat(
[
torch.arange(page * page_size, (page + 1) * page_size)
for page in segment_pages
]
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=CpSharedKVLayout(
page_size=page_size,
cp_size=4,
cp_rank=1,
),
nsa_cp_metadata=NSAContextParallelMetadata(
split_list=[page_size] * 8,
zigzag_index=[1, 6],
page_aligned=True,
page_size=page_size,
extend_prefix_len=0,
),
out_cache_loc=out_cache_loc,
)
local_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch)
self.assertIsNotNone(local_locs)
self.assertEqual(
local_locs.tolist(),
list(range(2 * page_size, 3 * page_size))
+ list(range(6 * page_size, 7 * page_size)),
)
def test_local_out_cache_loc_falls_back_when_owner_mismatch(self):
import torch
from types import SimpleNamespace
page_size = 4
out_cache_loc = torch.arange(page_size * 8, page_size * 16)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=CpSharedKVLayout(
page_size=page_size,
cp_size=4,
cp_rank=1,
),
nsa_cp_metadata=NSAContextParallelMetadata(
split_list=[page_size] * 8,
zigzag_index=[1, 6],
page_aligned=True,
page_size=page_size,
extend_prefix_len=0,
),
out_cache_loc=out_cache_loc,
)
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
def test_local_out_cache_loc_logs_every_fallback_event(self):
import torch
from types import SimpleNamespace
from sglang.srt.layers.attention.nsa import utils as nsa_utils
page_size = 4
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=CpSharedKVLayout(
page_size=page_size,
cp_size=4,
cp_rank=1,
),
nsa_cp_metadata=NSAContextParallelMetadata(
split_list=[page_size] * 8,
zigzag_index=[1, 6],
page_aligned=False,
page_size=page_size,
extend_prefix_len=0,
),
out_cache_loc=torch.arange(page_size * 8, page_size * 16),
)
with self.assertLogs(
"sglang.srt.layers.attention.nsa.utils", level="INFO"
) as cm:
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
self.assertEqual(len(cm.output), 2)
self.assertIn("CP shared KV direct-write fallback", cm.output[0])
self.assertIn("metadata is not page-aligned", cm.output[0])
self.assertIn("CP shared KV direct-write fallback", cm.output[1])
self.assertIn("metadata is not page-aligned", cm.output[1])
def test_indexer_direct_write_does_not_log_missing_metadata_for_non_cp_batch(self):
import torch
from types import SimpleNamespace
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
indexer.nsa_enable_prefill_cp = True
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
nsa_cp_metadata=None,
)
with self.assertNoLogs(
"sglang.srt.layers.attention.nsa.utils", level="INFO"
):
stored = Indexer._store_cp_shared_local_index_k_cache(
indexer,
forward_batch,
layer_id=0,
local_key=torch.empty(0),
act_quant=None,
)
self.assertFalse(stored)
def test_indexer_in_seq_cp_pair_materializes_index_once_for_prev_next(self):
import torch
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
logical_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32)
materialized_index = torch.tensor([11], dtype=torch.int32)
dense_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32)
materialize_calls = []
topk_calls = []
class Metadata:
def get_page_table_64(self):
return logical_pages
def fake_materialize(forward_batch, layer_id, logical_page_table):
materialize_calls.append((layer_id, logical_page_table))
return materialized_index, dense_pages
def fake_get_topk(
forward_batch,
layer_id,
q_fp8,
weights,
metadata,
kv_len,
actual_seq_q,
cp_index=None,
current_index_kv=None,
shared_index_buffer=None,
shared_block_tables=None,
):
topk_calls.append(
{
"kv_len": kv_len,
"actual_seq_q": actual_seq_q,
"shared_index_buffer": shared_index_buffer,
"shared_block_tables": shared_block_tables,
"current_index_kv": current_index_kv,
}
)
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
indexer._maybe_materialize_shared_index_buffer = fake_materialize
indexer._get_topk_ragged_with_cp = fake_get_topk
forward_batch = type(
"ForwardBatchStub",
(),
{
"nsa_cp_metadata": NSAContextParallelMetadata(
kv_len_prev=5,
kv_len_next=9,
actual_seq_q_prev=3,
actual_seq_q_next=2,
)
},
)()
q_fp8 = torch.arange(5 * 4, dtype=torch.float32).view(5, 4)
weights = torch.arange(5 * 2, dtype=torch.float32).view(5, 2)
result = Indexer._get_topk_in_seq_cp_pair(
indexer,
forward_batch,
layer_id=7,
q_fp8=q_fp8,
weights=weights,
metadata=Metadata(),
current_index_kv=None,
)
self.assertEqual(len(materialize_calls), 1)
self.assertIs(materialize_calls[0][1], logical_pages)
self.assertEqual(len(topk_calls), 2)
self.assertIs(topk_calls[0]["shared_index_buffer"], materialized_index)
self.assertIs(topk_calls[1]["shared_index_buffer"], materialized_index)
self.assertIs(topk_calls[0]["shared_block_tables"], dense_pages)
self.assertIs(topk_calls[1]["shared_block_tables"], dense_pages)
self.assertIsNone(topk_calls[0]["current_index_kv"])
self.assertEqual(topk_calls[0]["kv_len"], 5)
self.assertEqual(topk_calls[1]["kv_len"], 9)
self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]])
def test_indexer_in_seq_cp_pair_skips_materialize_when_current_index_reused(self):
import torch
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
current_index_kv = (torch.tensor([1]), torch.tensor([2]))
materialize_calls = []
topk_calls = []
class Metadata:
def get_page_table_64(self):
raise AssertionError("current index reuse should not read page table")
def fake_materialize(forward_batch, layer_id, logical_page_table):
materialize_calls.append((layer_id, logical_page_table))
raise AssertionError("current index reuse should not materialize")
def fake_get_topk(
forward_batch,
layer_id,
q_fp8,
weights,
metadata,
kv_len,
actual_seq_q,
cp_index=None,
current_index_kv=None,
shared_index_buffer=None,
shared_block_tables=None,
):
topk_calls.append(
{
"current_index_kv": current_index_kv,
"shared_index_buffer": shared_index_buffer,
"shared_block_tables": shared_block_tables,
}
)
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
indexer._maybe_materialize_shared_index_buffer = fake_materialize
indexer._get_topk_ragged_with_cp = fake_get_topk
forward_batch = type(
"ForwardBatchStub",
(),
{
"nsa_cp_metadata": NSAContextParallelMetadata(
kv_len_prev=5,
kv_len_next=9,
actual_seq_q_prev=3,
actual_seq_q_next=2,
)
},
)()
result = Indexer._get_topk_in_seq_cp_pair(
indexer,
forward_batch,
layer_id=7,
q_fp8=torch.empty(5, 4),
weights=torch.empty(5, 2),
metadata=Metadata(),
current_index_kv=current_index_kv,
)
self.assertEqual(materialize_calls, [])
self.assertEqual(len(topk_calls), 2)
self.assertIs(topk_calls[0]["current_index_kv"], current_index_kv)
self.assertIs(topk_calls[1]["current_index_kv"], current_index_kv)
self.assertIsNone(topk_calls[0]["shared_index_buffer"])
self.assertIsNone(topk_calls[1]["shared_block_tables"])
self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]])
if __name__ == "__main__":
unittest.main()
@@ -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()