diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 1008aa3ca..ad98804b0 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -3885,6 +3885,62 @@ def build_batch_current_slot_spans( return _merge_slot_spans(spans) +def get_or_build_batch_slot_spans( + forward_batch: Any, + *, + logical_pages: torch.Tensor, + prefix_lens_cpu, + extend_lens_cpu=None, + page_size: int, + want_prefix: bool, +) -> tuple[list[tuple[int, int]], list[tuple[int, int]]]: + """Cache exact flattened page-table spans for batched prefix/current pages. + + The spans are batch metadata: they depend on the logical page table and the + per-request prefix/extend lengths, not on layer contents. Reusing them + across layers avoids rebuilding Python descriptor lists in the shared-KV + compose path while preserving per-request row boundaries. + """ + + prefix_lens_key = tuple(int(x) for x in prefix_lens_cpu) + extend_lens_key = ( + None + if extend_lens_cpu is None + else tuple(int(x) for x in extend_lens_cpu) + ) + key = ( + _tensor_identity_key(logical_pages), + prefix_lens_key, + extend_lens_key, + int(page_size), + bool(want_prefix), + ) + cached_key = getattr(forward_batch, "cp_shared_kv_batch_slot_spans_key", None) + cached = getattr(forward_batch, "cp_shared_kv_batch_slot_spans", None) + if cached is not None and cached_key == key: + return cached + + prefix_slot_spans = ( + build_batch_prefix_slot_spans( + logical_pages=logical_pages, + prefix_lens_cpu=prefix_lens_cpu, + page_size=page_size, + ) + if want_prefix + else [] + ) + current_slot_spans = build_batch_current_slot_spans( + logical_pages=logical_pages, + prefix_lens_cpu=prefix_lens_cpu, + extend_lens_cpu=extend_lens_cpu, + page_size=page_size, + ) + cached = (prefix_slot_spans, current_slot_spans) + forward_batch.cp_shared_kv_batch_slot_spans_key = key + forward_batch.cp_shared_kv_batch_slot_spans = cached + return cached + + def current_loc_remap_fast_path_args( forward_batch, ) -> tuple[int | None, int | None]: diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 0de5e697c..031929f50 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -31,7 +31,6 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( log_cp_draft_shared_kv_debug, materialize_prefix_and_reuse_current_index_page_slots, materialize_shared_paged_buffer, - maybe_build_current_page_writer_ranks, should_reuse_current_extend_kv, tensor_debug_checksum, tensor_debug_summary, @@ -825,13 +824,6 @@ class Indexer(MultiPlatformOp): prefix_slot_spans=prefix_slot_spans, current_slot_spans=current_slot_spans, layer_id=layer_id, - current_page_writer_ranks=maybe_build_current_page_writer_ranks( - forward_batch=forward_batch, - prefix_lens_cpu=prefix_lens_cpu, - extend_lens_cpu=extend_lens_cpu, - page_size=page_size, - layout=layout, - ), ) ) if ( diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index 5adc4b04f..b49220db6 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -2304,6 +2304,40 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase): [(0, 2), (4, 5)], ) + def test_get_or_build_batch_slot_spans_caches_exact_request_spans(self): + from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime + + forward_batch = SimpleNamespace() + logical_pages = torch.tensor( + [ + [1, 2, 5, 0], + [9, 11, 12, 13], + ], + dtype=torch.int64, + ) + + prefix_spans, current_spans = runtime.get_or_build_batch_slot_spans( + forward_batch, + logical_pages=logical_pages, + prefix_lens_cpu=[8, 4], + extend_lens_cpu=[2, 7], + page_size=4, + want_prefix=True, + ) + + self.assertEqual(prefix_spans, [(0, 2), (4, 5)]) + self.assertEqual(current_spans, [(2, 3), (5, 7)]) + cached = runtime.get_or_build_batch_slot_spans( + forward_batch, + logical_pages=logical_pages, + prefix_lens_cpu=[8, 4], + extend_lens_cpu=[2, 7], + page_size=4, + want_prefix=True, + ) + self.assertIs(cached[0], prefix_spans) + self.assertIs(cached[1], current_spans) + def test_mla_prefetch_create_batch_uses_exact_prefix_and_current_spans(self): from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch