Reduce CP shared KV overhead without changing ownership semantics

The shared-KV path now keeps more CP metadata on-device and reuses
physical out-cache locations across MLA and NSA index writes, so each
layer avoids repeating logical-to-physical remaps. The in-seq CP
all-gather rerange path now delegates to tai-kernel when available and
falls back to the existing torch split/cat path with an explicit log.

This also extends the Phase8 prefetch machinery to cover shared KV
materialization metadata and keeps debug/fallback behavior gated so the
fast path is not polluted by diagnostic checks.

Constraint: Custom CP kernels must live in tai-kernel and be imported lazily from SGLang
Constraint: Decode does not use CP; these changes target NSA prefill CP in-seq-split shared KV
Rejected: Recompute physical local cache locations separately for MLA and index writes | repeats the same remap work every layer
Rejected: Keep the in-seq rerange Triton code inline in SGLang | duplicates kernel ownership and blocks tai-kernel reuse
Confidence: medium
Scope-risk: moderate
Directive: Keep CP collective ordering identical across ranks; do not add rank-local fallback decisions inside shared KV materialize paths
Tested: Remote g0034 container py_compile for modified SGLang/tai-kernel files; remote pytest test/registered/unit/layers/test_nsa_cp_utils.py passed with 24 tests
Not-tested: Full multi-node GLM5 prefill/decode throughput after the final commit boundary
This commit is contained in:
laoyao0822
2026-05-06 05:27:43 +08:00
parent 5e5ac5e2e7
commit 43ad2fe52d
10 changed files with 1152 additions and 46 deletions
@@ -11,6 +11,7 @@ from sglang.srt.layers.attention.nsa.utils import (
can_cp_split,
cp_split_and_rebuild_1d,
get_cp_shared_kv_local_out_cache_loc,
get_cp_shared_kv_local_physical_out_cache_loc,
split_in_seq_cp_local_pair,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
@@ -335,6 +336,45 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
+ list(range(6 * page_size, 7 * page_size)),
)
def test_local_physical_out_cache_loc_is_cached(self):
import torch
from types import SimpleNamespace
page_size = 4
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,
)
physical_locs = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch)
second_read = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch)
self.assertIs(physical_locs, second_read)
self.assertEqual(
physical_locs.tolist(),
list(range(1 * page_size, 2 * page_size))
+ list(range(2 * page_size, 3 * page_size)),
)
def test_local_out_cache_loc_falls_back_when_owner_mismatch(self):
import torch
from types import SimpleNamespace
@@ -1214,6 +1214,254 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertTrue(torch.equal(dense_page_buffer[1], page_buffer[1]))
self.assertTrue(torch.equal(dense_page_buffer[2], page_buffer[2]))
def test_materialize_local_paged_buffer_page_slots_into_matches_full_slots(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=64, cp_size=2, cp_rank=0)
page_buffer = torch.arange(0, 6 * 3, dtype=torch.float32).view(6, 3)
slot_logical_pages = torch.tensor([1, 2, 3, 4, 0, -1], dtype=torch.int64)
full = runtime.materialize_local_paged_buffer_page_slots(
page_buffer=page_buffer,
slot_logical_pages=slot_logical_pages,
layout=layout,
)
split = page_buffer.new_full(full.shape, -7)
split[0].zero_()
runtime.materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=split,
slot_logical_pages=slot_logical_pages,
layout=layout,
start_slot=0,
end_slot=3,
)
runtime.materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=split,
slot_logical_pages=slot_logical_pages,
layout=layout,
start_slot=3,
end_slot=slot_logical_pages.numel(),
)
self.assertTrue(torch.equal(split, full))
def test_remap_logical_pages_to_slot_dense_pages_preserves_sentinels(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
slot_logical_pages = torch.tensor([1, 2, 0, 4, -1, 5], dtype=torch.int64)
page_inverse = runtime.build_slot_page_inverse(
slot_logical_pages,
logical_page_capacity=8,
)
logical_pages = torch.tensor([[0, 1, 5, -1, 3, 9]], dtype=torch.int32)
dense_pages = runtime.remap_logical_pages_to_slot_dense_pages(
logical_pages,
page_inverse=page_inverse,
)
self.assertEqual(dense_pages.tolist(), [[0, 1, 6, -1, -1, -1]])
def test_index_materialize_uses_prefetched_buffer_before_fallback(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
class FakePool:
def __init__(self):
self.index_buffer = torch.arange(0, 12, dtype=torch.float32).view(4, 3)
def get_index_k_with_scale_buffer(self, layer_id):
self.layer_id = layer_id
return self.index_buffer
class FakePrefetcher:
def __init__(self):
self.calls = []
self.dense_buffer = torch.full((3, 3), 5.0)
self.dense_pages = torch.tensor([[1, 2]], dtype=torch.int32)
def consume(self, *, layer_id, page_buffer, logical_pages):
self.calls.append((layer_id, page_buffer, logical_pages))
return self.dense_buffer, self.dense_pages
fake_pool = FakePool()
fake_prefetcher = FakePrefetcher()
forward_batch = SimpleNamespace(
token_to_kv_pool=fake_pool,
uses_cp_shared_kv=True,
cp_shared_kv_layout=object(),
cp_shared_kv_index_prefetcher=fake_prefetcher,
)
logical_pages = torch.tensor([[1, 2]], dtype=torch.int32)
indexer = object.__new__(nsa_indexer.Indexer)
with patch.object(
nsa_indexer,
"materialize_shared_paged_buffer",
side_effect=AssertionError("prefetch hit must bypass full materialize"),
):
dense_buffer, dense_pages = indexer._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id=7,
logical_page_table=logical_pages,
)
self.assertIs(dense_buffer, fake_prefetcher.dense_buffer)
self.assertIs(dense_pages, fake_prefetcher.dense_pages)
self.assertEqual(fake_prefetcher.calls[0][0], 7)
self.assertIs(fake_prefetcher.calls[0][1], fake_pool.index_buffer)
self.assertIs(fake_prefetcher.calls[0][2], logical_pages)
def test_index_prefetch_start_targets_next_layer(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
class FakePrefetcher:
def __init__(self):
self.calls = []
def start_next_layer_prefix(self, *, next_layer_id, token_to_kv_pool):
self.calls.append((next_layer_id, token_to_kv_pool))
token_to_kv_pool = object()
fake_prefetcher = FakePrefetcher()
forward_batch = SimpleNamespace(
token_to_kv_pool=token_to_kv_pool,
cp_shared_kv_index_prefetcher=fake_prefetcher,
)
indexer = object.__new__(nsa_indexer.Indexer)
indexer._maybe_start_next_layer_index_prefetch(forward_batch, layer_id=11)
self.assertEqual(fake_prefetcher.calls, [(12, token_to_kv_pool)])
def test_index_prefetch_consume_miss_logs_fallback_after_first_layer(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
class FakePool:
start_layer = 0
def __init__(self):
self.index_buffer = torch.arange(0, 12, dtype=torch.float32).view(4, 3)
def get_index_k_with_scale_buffer(self, layer_id):
return self.index_buffer
class FakeLayout:
cp_rank = 3
class MissingPrefetcher:
def consume(self, *, layer_id, page_buffer, logical_pages):
return None
fallback_buffer = torch.full((3, 3), 9.0)
fallback_pages = torch.tensor([[1, 2]], dtype=torch.int32)
forward_batch = SimpleNamespace(
token_to_kv_pool=FakePool(),
uses_cp_shared_kv=True,
cp_shared_kv_layout=FakeLayout(),
cp_shared_kv_index_prefetcher=MissingPrefetcher(),
)
logical_pages = torch.tensor([[1, 2]], dtype=torch.int32)
indexer = object.__new__(nsa_indexer.Indexer)
with patch.object(
nsa_indexer,
"materialize_shared_paged_buffer",
return_value=(fallback_buffer, fallback_pages),
), patch(
"sglang.srt.layers.attention.nsa.nsa_indexer.logger",
create=True,
) as logger:
dense_buffer, dense_pages = indexer._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id=7,
logical_page_table=logical_pages,
)
self.assertIs(dense_buffer, fallback_buffer)
self.assertIs(dense_pages, fallback_pages)
logger.info.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
logger.info.call_args.args[0],
)
self.assertIn("consume_miss", logger.info.call_args.args[1])
def test_index_prefetch_first_layer_miss_does_not_log_fallback(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
class FakePool:
start_layer = 0
def __init__(self):
self.index_buffer = torch.arange(0, 12, dtype=torch.float32).view(4, 3)
def get_index_k_with_scale_buffer(self, layer_id):
return self.index_buffer
class FakeLayout:
cp_rank = 0
class MissingPrefetcher:
def consume(self, *, layer_id, page_buffer, logical_pages):
return None
fallback_buffer = torch.full((3, 3), 9.0)
fallback_pages = torch.tensor([[1, 2]], dtype=torch.int32)
forward_batch = SimpleNamespace(
token_to_kv_pool=FakePool(),
uses_cp_shared_kv=True,
cp_shared_kv_layout=FakeLayout(),
cp_shared_kv_index_prefetcher=MissingPrefetcher(),
)
logical_pages = torch.tensor([[1, 2]], dtype=torch.int32)
indexer = object.__new__(nsa_indexer.Indexer)
with patch.object(
nsa_indexer,
"materialize_shared_paged_buffer",
return_value=(fallback_buffer, fallback_pages),
), patch(
"sglang.srt.layers.attention.nsa.nsa_indexer.logger",
create=True,
) as logger:
indexer._maybe_materialize_shared_index_buffer(
forward_batch,
layer_id=0,
logical_page_table=logical_pages,
)
logger.info.assert_not_called()
def test_index_prefetch_create_skip_logs_fallback_when_enabled(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
with patch.object(
prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True
), patch.object(
prefetch, "cp_shared_kv_debug_enabled", return_value=False
), patch.object(
prefetch.torch.cuda, "is_available", return_value=False
), patch.object(
prefetch.logger, "info"
) as logger:
result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
forward_batch=SimpleNamespace(),
metadata=SimpleNamespace(),
topk_transform_is_paged=True,
)
self.assertIsNone(result)
logger.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
logger.call_args.args[0],
)
self.assertIn("cuda_unavailable_or_stream_capturing", logger.call_args.args[1])
if __name__ == "__main__":
unittest.main()