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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user