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