Overlap CP shared KV prefix materialization for cached MLA prefill

Shared CP KV materialization remained on the critical path for cached
NSA/MLA prefill batches.  This change introduces a one-layer-ahead
prefetcher that materializes the cached prefix for the next layer on a
separate CUDA stream and consumes it when that layer reaches attention.
The prefetch path keeps the existing dense page-table semantics, defers
waiting until the prefetched buffer is actually consumed, and uses the
TAI optimized materialize/remap helpers when enabled before falling back
to the torch implementation.

The implementation is intentionally gated by environment variables and
keeps layer-2-only probe logging for functional confirmation without
making normal profiling noisy.

Constraint: Prefill CP shared KV must preserve existing page-table and dense KV semantics for NSA paged topk attention
Constraint: The production performance path requires SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE=1 and logging disabled
Rejected: Wait immediately after the producer layer attention | this truncated the overlap window and hid less work
Rejected: Torch-only prefetch materialize | it bypassed the optimized TAI materialize/remap path and could erase the expected win
Confidence: medium
Scope-risk: moderate
Directive: Do not evaluate Phase8 throughput with SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH=1; use it only to confirm create/start/consume_hit behavior
Tested: Local AST parse for modified Python files
Tested: Local git diff --check
Tested: Remote g0034 container AST parse for modified files under /sgl-workspace/sglang-tai
Tested: Remote g0034 container pytest target covering Phase8 log env, TAI range materialize, optimized slot inverse/remap, and existing token TAI path
Not-tested: Full prefill/decode/router throughput after the TAI prefetch-path fix
This commit is contained in:
laoyao0822
2026-05-03 03:09:59 +08:00
parent 5769b63082
commit bc23a81884
6 changed files with 1178 additions and 99 deletions
@@ -182,6 +182,208 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertTrue(torch.equal(dense_kv[20:24], kv_cache[12:16]))
self.assertEqual(float(dense_kv[4:8].abs().sum().item()), 0.0)
def test_materialize_local_token_kv_page_slots_into_matches_full_slots(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
materialize_local_token_kv_page_slots,
materialize_local_token_kv_page_slots_into,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
kv_cache = torch.arange(0, 32 * 2, dtype=torch.float32).view(32, 1, 2)
slot_logical_pages = torch.tensor([1, 2, 3, 4, 5, 6], dtype=torch.int64)
full = materialize_local_token_kv_page_slots(
kv_cache=kv_cache,
slot_logical_pages=slot_logical_pages,
layout=layout,
page_size=4,
)
ranged = kv_cache.new_zeros(full.shape)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=ranged,
slot_logical_pages=slot_logical_pages,
layout=layout,
page_size=4,
start_slot=0,
end_slot=3,
)
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=ranged,
slot_logical_pages=slot_logical_pages,
layout=layout,
page_size=4,
start_slot=3,
end_slot=slot_logical_pages.numel(),
)
self.assertTrue(torch.equal(ranged, full))
def test_slot_range_to_token_slice_preserves_dummy_page_offset(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
slot_range_to_token_slice,
)
self.assertEqual(slot_range_to_token_slice(4, 0, 2), slice(4, 12))
self.assertEqual(slot_range_to_token_slice(4, 2, 6), slice(12, 28))
def test_build_shared_token_kv_slot_remap_reuses_slot_layout(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_shared_token_kv_slot_remap,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0)
kv_cache = torch.zeros((20, 1, 2), dtype=torch.float32)
logical_locs = torch.tensor([[4, 12, -1], [16, 0, 20]], dtype=torch.int64)
remap_logical_pages = torch.tensor([[1, 3, 0], [4, 5, 6]], dtype=torch.int64)
remap = build_shared_token_kv_slot_remap(
kv_cache=kv_cache,
logical_locs=logical_locs,
remap_logical_pages=remap_logical_pages,
layout=layout,
page_size=4,
)
self.assertEqual(remap.slot_logical_pages.tolist(), [1, 3, 0, 4, 5, 6])
self.assertEqual(remap.dense_num_pages, 7)
self.assertEqual(remap.dense_locs.tolist(), [[4, 8, -1], [16, 0, 20]])
def test_mla_prefetch_log_env_defaults_to_off_and_can_enable(self):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
cp_shared_kv_mla_prefetch_log_enabled,
cp_shared_kv_mla_prefetch_should_log_layer,
cp_shared_kv_mla_prefetch_wait_after_attention_enabled,
)
envs.SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH.clear()
envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_WAIT_AFTER_ATTENTION.clear()
self.assertFalse(cp_shared_kv_mla_prefetch_log_enabled())
self.assertFalse(cp_shared_kv_mla_prefetch_wait_after_attention_enabled())
self.assertFalse(cp_shared_kv_mla_prefetch_should_log_layer(1))
self.assertTrue(cp_shared_kv_mla_prefetch_should_log_layer(2))
self.assertFalse(cp_shared_kv_mla_prefetch_should_log_layer(3))
with envs.SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH.override(True):
self.assertTrue(cp_shared_kv_mla_prefetch_log_enabled())
with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_WAIT_AFTER_ATTENTION.override(True):
self.assertTrue(cp_shared_kv_mla_prefetch_wait_after_attention_enabled())
def test_token_range_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.token_calls = []
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)
)
rows = (slot_logical_pages.numel() + 1) * page_size
return torch.arange(
rows * 2,
dtype=kv_cache.dtype,
device=kv_cache.device,
).view(rows, 1, 2)
fake_tai = FakeTaiKernels()
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
kv_cache = torch.zeros((32, 1, 2), dtype=torch.float32)
dense_kv_cache = torch.zeros((28, 1, 2), dtype=torch.float32)
slot_logical_pages = torch.tensor([1, 2, 3, 4, 5, 6], 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
):
runtime.materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_logical_pages,
layout=layout,
page_size=4,
start_slot=2,
end_slot=5,
)
self.assertEqual(len(fake_tai.token_calls), 1)
self.assertTrue(
torch.equal(
fake_tai.token_calls[0][1],
torch.tensor([3, 4, 5], dtype=torch.int64),
)
)
expected_tmp = torch.arange(32, dtype=torch.float32).view(16, 1, 2)
self.assertEqual(float(dense_kv_cache[:12].sum().item()), 0.0)
self.assertTrue(torch.equal(dense_kv_cache[12:24], expected_tmp[4:16]))
self.assertEqual(float(dense_kv_cache[24:].sum().item()), 0.0)
def test_slot_remap_helpers_use_tai_when_enabled(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
class FakeTaiKernels:
def __init__(self):
self.inverse_calls = []
self.remap_calls = []
def build_slot_page_inverse(self, slot_logical_pages, logical_page_capacity):
self.inverse_calls.append((slot_logical_pages, logical_page_capacity))
return torch.tensor([0, 1, 2, -1], 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, -1], dtype=logical_locs.dtype)
fake_tai = FakeTaiKernels()
slot_logical_pages = torch.tensor([1, 2, 3], dtype=torch.int64)
logical_locs = torch.tensor([4, 12], 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
):
page_inverse = runtime.build_slot_page_inverse_optimized(
slot_logical_pages,
logical_page_capacity=4,
)
dense_locs = runtime.remap_logical_locs_to_slot_dense_locs_optimized(
logical_locs,
page_inverse=page_inverse,
page_size=4,
)
self.assertEqual(len(fake_tai.inverse_calls), 1)
self.assertEqual(len(fake_tai.remap_calls), 1)
self.assertEqual(dense_locs.tolist(), [4, -1])
def test_materialize_local_paged_index_buffer(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_dense_page_remap,