Repair CP current-slot composition instead of disabling reuse

CP shared KV current-only and partial-current paths were previously unsafe because owner-local current rows could be filled into dense page slots without synchronizing the current slots across CP ranks. This preserves the fast path and fixes the missing synchronization at the actual contract boundary: prefix slots remain materialized by the existing IPC/collective path, while freshly produced current slots are composed from current forward tensors and reduced only over the current slot ranges.\n\nThe change also keeps page-tail compute padding explicit, restores current-only index compose, fixes TAI index descriptor dtype consistency, and records the investigation ledger to avoid repeating discarded hypotheses.\n\nConstraint: CP shared KV must preserve current reuse and TAI materialize fast paths; blanket disable is not an acceptable fix.\nConstraint: Prefix slots must not be reduced twice after IPC/prefix materialization.\nRejected: Disable current-only MLA/index reuse | hides the owner-local composition bug and regresses the intended fast path.\nRejected: Disable TAI materialize globally | avoids symptoms without proving the byte/layout contract.\nConfidence: medium\nScope-risk: broad\nDirective: Do not remove current-slot range synchronization unless replacing it with an equivalent owner-aware P2P/IPC gather contract.\nTested: Local py_compile for touched files.\nTested: Remote g0034 py_compile for touched runtime/test files.\nTested: Remote test_cp_shared_kv_runtime.py: 111 passed, 5 warnings, 2 subtests passed.\nTested: Remote targeted current-slot compose tests and direct current-only index compose script.\nNot-tested: Full ETE output quality and decode accept len after restart.\nNot-tested: Full test_nsa_cp_utils.py collection on remote due incomplete installed sgl_kernel import/fake-op environment.
This commit is contained in:
laoyao0822
2026-06-05 06:25:06 +08:00
parent e200091638
commit c8510938b9
8 changed files with 898 additions and 63 deletions
@@ -1030,6 +1030,43 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIn("[CP_SHARED_KV_FALLBACK][current_reuse]", joined)
self.assertIn("prefix_extend_seq_len_mismatch_req_0", joined)
def test_cp_shared_current_only_reuse_accepts_owner_local_current_rows(
self,
):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
forward_mode=_FakeExtendForwardMode(),
batch_size=1,
extend_prefix_lens_cpu=[0],
extend_seq_lens_cpu=[40387],
seq_lens_cpu=torch.tensor([40387], dtype=torch.int32),
out_cache_loc=torch.arange(40392, dtype=torch.int64),
cp_local_out_cache_loc=torch.arange(5056, dtype=torch.int64),
)
local_k = torch.empty((5056, 1), dtype=torch.float32)
local_rope = torch.empty((5056, 1), dtype=torch.float32)
with envs.SGLANG_CP_SHARED_KV_CURRENT_REUSE.override(True):
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
self.assertEqual(
runtime.current_extend_kv_rows_for_reuse(
forward_batch,
local_k,
local_rope,
),
5056,
)
self.assertIsNone(
runtime.current_extend_kv_rows_for_reuse(
forward_batch,
local_k[:5055],
local_rope,
)
)
def test_tai_index_mqa_prepare_fast_path_miss_logs_warning(self):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -1371,6 +1408,36 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
(0, 0),
)
def test_batch_current_slot_spans_follow_prefix_and_extend_pages(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
logical_pages = torch.tensor(
[
[1, 2, 5, 0],
[9, 11, 12, 13],
],
dtype=torch.int64,
)
self.assertEqual(
runtime.build_batch_current_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[8, 4],
extend_lens_cpu=[2, 7],
page_size=4,
),
[(2, 3), (5, 7)],
)
self.assertEqual(
runtime.build_batch_current_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[0, 0],
extend_lens_cpu=[8, 4],
page_size=4,
),
[(0, 2), (4, 5)],
)
def test_materialize_batch_prefix_span_and_reuse_current_kv_page_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
@@ -1406,6 +1473,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefix_lens_cpu=[8, 4],
page_size=page_size,
)
current_slot_spans = runtime.build_batch_current_slot_spans(
logical_pages=remap_logical_pages,
prefix_lens_cpu=[8, 4],
extend_lens_cpu=[2, 2],
page_size=page_size,
)
with patch.object(
runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce
@@ -1421,6 +1494,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
page_size=page_size,
prefix_pages=0,
prefix_slot_span=prefix_slot_span,
current_slot_spans=current_slot_spans,
)
)
@@ -1455,6 +1529,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefix_lens_cpu=[8, 4],
page_size=page_size,
)
current_slot_spans = runtime.build_batch_current_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[8, 4],
extend_lens_cpu=[2, 2],
page_size=page_size,
)
current_k = torch.tensor(
[[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12], [13, 14, 15, 16]],
dtype=torch.uint8,
@@ -1476,6 +1556,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_span=prefix_slot_span,
current_slot_spans=current_slot_spans,
layer_id=2,
)
)
@@ -1588,7 +1669,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertEqual(dense_locs.tolist(), [4, 8, 12])
self.assertEqual(list(dense_kv[4:16].flatten().tolist()), list(range(12)))
def test_materialize_prefix_current_token_kv_uses_ipc_without_all_reduce(self):
def test_materialize_prefix_current_token_kv_uses_ipc_and_reduces_current_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
@@ -1615,6 +1696,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
).view(8, 1, 1)
return True
range_calls = []
def record_range_reduce(buffer, cp_size, start_row, end_row, **kwargs):
range_calls.append((start_row, end_row, kwargs.get("nvtx_source")))
return buffer
with patch.object(
runtime,
"_try_tai_ipc_materialize_token_kv_page_slots_into",
@@ -1622,7 +1709,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
), patch.object(
runtime,
"_all_reduce_materialized_buffer_range",
side_effect=AssertionError("IPC path must not range all-reduce"),
side_effect=record_range_reduce,
):
mixed_kv, mixed_locs = (
runtime.materialize_prefix_and_reuse_current_kv_page_slots(
@@ -1641,6 +1728,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertEqual(mixed_kv[4].item(), 10)
self.assertEqual(mixed_kv[8].item(), 14)
self.assertTrue(torch.equal(mixed_kv[12:14], current_kv))
self.assertEqual(range_calls, [(12, 16, "mla.partial_current_sync.current")])
def test_mla_prefetch_consume_prefix_with_current_skips_suffix_materialize(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -1676,6 +1764,11 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
)
prefetcher.handles[1] = handle
prefetcher.pending_attention_handle = handle
range_calls = []
def record_range_reduce(buffer, cp_size, start_row, end_row, **kwargs):
range_calls.append((start_row, end_row, kwargs.get("nvtx_source")))
return buffer
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
@@ -1683,6 +1776,10 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetch,
"materialize_local_token_kv_page_slots_into",
side_effect=AssertionError("suffix materialize must not run"),
), patch.object(
prefetch,
"_all_reduce_materialized_buffer_range",
side_effect=record_range_reduce,
):
mixed_kv, mixed_locs = prefetcher.consume_prefix_with_current(
layer_id=1,
@@ -1701,6 +1798,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
expected_kv[12:14] = current_kv
self.assertTrue(torch.equal(mixed_kv, expected_kv))
self.assertEqual(mixed_locs.tolist(), [[4, 12], [13, 7], [-1, -1]])
self.assertEqual(range_calls, [(12, 16, "mla.prefetch_current")])
def test_mla_prefetch_attention_window_defers_pending_event_wait(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -1995,7 +2093,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
[[1, 2, 3], [4, 5, 6], [7, 8, 9]],
)
def test_materialize_prefix_current_index_uses_ipc_without_all_reduce(self):
def test_materialize_prefix_current_index_uses_ipc_and_reduces_current_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
@@ -2025,6 +2123,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
)
return True
range_calls = []
def record_range_reduce(buffer, cp_size, start_row, end_row, **kwargs):
range_calls.append((start_row, end_row, kwargs.get("nvtx_source")))
return buffer
with patch.object(
runtime,
"_try_tai_ipc_materialize_paged_buffer_page_slots_into",
@@ -2032,7 +2136,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
), patch.object(
runtime,
"_all_reduce_materialized_buffer_range",
side_effect=AssertionError("IPC path must not range all-reduce"),
side_effect=record_range_reduce,
):
dense_page_buffer, dense_pages = (
runtime.materialize_prefix_and_reuse_current_index_page_slots(
@@ -2053,6 +2157,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertEqual(dense_page_buffer[1].tolist(), list(range(page_bytes)))
self.assertTrue(torch.equal(dense_page_buffer[2, 0:4], current_k[0]))
self.assertTrue(torch.equal(dense_page_buffer[2, 4:8], current_k[1]))
self.assertEqual(range_calls, [(2, 3, "index.partial_current_sync.current")])
def test_index_prefetch_partial_current_compose_fills_current_page_slots(self):
from sglang.srt.environ import envs
@@ -2138,7 +2243,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
method_source = source[start:end]
self.assertIn("should_reuse_current_extend_kv(forward_batch)", method_source)
self.assertNotIn("is_current_only_extend_batch(forward_batch)", method_source)
self.assertNotIn("current_only_compact_unsupported", method_source)
def test_index_current_reuse_prepare_accepts_padded_out_cache_loc(self):
from pathlib import Path