Fix CP shared-KV bs=1 cache-hit zero-prefix corruption; spans are sole compose input

A bs=1 cache-hit prefill produced deterministically wrong output (e.g. "0.5,0.5,0.5").
Root cause: regression a24111a5f4 changed the bs=1 cache-hit call sites to pass
prefix_slot_spans=[] (get_or_build_batch_slot_spans want_prefix=False) together with
prefix_pages>0. In materialize_prefix_and_reuse_current_kv_page_slots and its index twin,
the guard `if prefix_slot_spans is not None:` let the empty list shadow the prefix_pages
fallback -> prefix_spans=[] -> the cached prefix dense slots were never materialized
(both IPC and local paths gate on `if prefix_spans:`) -> attention read zero KV over the
entire reused prefix. The indexer compose dropped its prefix the same way. bs=1 only;
bs>1 (non-empty per-request spans) and fresh prefill (no prefix) were unaffected.

Fix: make the canonical per-request slot-span list the sole description of the
prefix/current regions and fail loud on a missing list. Remove all four
span-reconstruction fallbacks: prefix_slot_span (singular, dead), prefix_pages->span
(the bug), current_slot_spans=None->span (dead), and the prefetcher's copy. Both twins
now require prefix_slot_spans + current_slot_spans (raise on None; [] = genuine
no-prefix). All call sites build canonical spans via want_prefix=True (cached per-forward,
so a24111a5f4's per-layer-CPU goal is preserved by the cache, not by skipping the build);
the MLA current-only path also uses want_prefix=True to keep want_prefix uniform and the
span cache thrash-free.

Tests: migrate unit tests off the removed prefix_pages param (drop it where spans were
already passed; supply canonical spans otherwise); fix two source-string/captured-kwarg
assertions; add test_materialize_prefix_requires_explicit_prefix_slot_spans (None->raise,
explicit spans -> non-zero prefix). test_cp_shared_kv_runtime.py: 158 passed on the
g0033 container.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-06-17 18:25:22 +00:00
co-authored by Claude Opus 4.8
parent b34e7cb932
commit e23168e7f5
7 changed files with 156 additions and 166 deletions
@@ -218,6 +218,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([0, 1, 2], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2], dtype=torch.int64),
dense_num_pages=4,
@@ -594,6 +596,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([0, 1, 2], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2], dtype=torch.int64),
dense_num_pages=4,
@@ -1292,7 +1296,10 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIn(
"materialize_prefix_and_reuse_current_kv_page_slots", branch_source
)
self.assertIn("prefix_pages=0", branch_source)
# current-only compose now passes the canonical (empty) prefix_slot_spans
# instead of the removed prefix_pages scalar fallback.
self.assertIn("prefix_slot_spans=prefix_slot_spans", branch_source)
self.assertNotIn("prefix_pages=", branch_source)
self.assertNotIn("kv_cache = current_kv_cache", branch_source)
def test_mla_partial_current_sync_uses_batch_prefix_slot_spans(self):
@@ -2192,7 +2199,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
)
)
@@ -2202,6 +2209,56 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertTrue(torch.equal(mixed_kv[12:14], current_kv))
self.assertEqual(mixed_locs.tolist(), [[4, 8, 12, 13, -1, -1]])
def test_materialize_prefix_requires_explicit_prefix_slot_spans(self):
# Regression (a24111a5f4): a bs=1 cache-hit passed prefix_slot_spans=[] +
# prefix_pages>0; the `is not None` guard let [] shadow the prefix_pages
# fallback so the cached prefix was never materialized (zero KV over the
# reused prefix -> garbage output). The compose now takes prefix spans as
# the sole canonical input: None must fail loud; an explicit span list must
# materialize a NON-ZERO prefix.
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
page_size = 4
layout = CpSharedKVLayout(page_size=page_size, cp_size=1, cp_rank=0)
kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1)
logical_locs = torch.tensor([[4, 8, 20, 21, 22, 23]], dtype=torch.int64)
current_locs = torch.tensor([20, 21], dtype=torch.int64)
current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1)
slot_remap = runtime.build_shared_token_kv_slot_remap(
kv_cache=kv_cache,
logical_locs=logical_locs,
remap_logical_pages=torch.tensor([[1, 2, 5]], dtype=torch.int64),
layout=layout,
page_size=page_size,
)
def _compose(prefix_slot_spans):
with patch.object(
runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce
):
return runtime.materialize_prefix_and_reuse_current_kv_page_slots(
kv_cache=kv_cache,
logical_locs=logical_locs,
current_kv_cache=current_kv,
current_locs=current_locs,
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_slot_spans=prefix_slot_spans,
)
# None -> fail loud (never silently materialize a zero prefix).
with self.assertRaises(ValueError):
_compose(None)
# The canonical bs=1 cache-hit spans materialize the prefix as NON-ZERO
# (the rows the regression left empty).
mixed_kv, _ = _compose([(0, 2)])
self.assertGreater(mixed_kv[4:12].abs().sum().item(), 0.0)
self.assertTrue(torch.equal(mixed_kv[4:8], kv_cache[4:8]))
self.assertTrue(torch.equal(mixed_kv[8:12], kv_cache[8:12]))
def test_batch_prefix_slot_span_covers_bounding_prefix_range(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -2956,7 +3013,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
)
@@ -3027,7 +3083,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
)
@@ -3084,7 +3139,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=2,
@@ -3249,7 +3303,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
)
)
@@ -3303,7 +3358,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=[],
current_slot_spans=[(0, 1)],
layer_id=0,
)
@@ -3386,6 +3441,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=page_inverse,
dense_num_pages=4,
@@ -3452,6 +3509,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64),
dense_num_pages=4,
@@ -3499,6 +3558,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64),
dense_num_pages=4,
@@ -3538,6 +3599,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64),
dense_num_pages=4,
@@ -3578,6 +3641,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
prefix_pages=2,
prefix_slot_spans=[(0, 2)],
current_slot_spans=[(2, 3)],
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64),
dense_num_pages=4,
@@ -3623,6 +3688,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=page_size, cp_size=cp_size, cp_rank=0),
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
current_slot_spans=[(1, 3)],
slot_logical_pages=torch.tensor([1, 2, 3], dtype=torch.int64),
page_inverse=torch.tensor([-1, 1, 2, 3], dtype=torch.int64),
dense_num_pages=padded_pages,
@@ -3699,7 +3766,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
layer_id=2,
)
)
@@ -3831,7 +3898,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
current_slot_spans=[(1, 2)],
layer_id=2,
)
)
@@ -3889,7 +3957,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_spans=[],
current_slot_spans=[(0, 1)],
layer_id=0,
)
@@ -3968,6 +4036,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=page_size, cp_size=1, cp_rank=0),
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
current_slot_spans=[(1, 2)],
slot_logical_pages=torch.tensor([1, 20], dtype=torch.int64),
page_inverse=torch.tensor(
[-1, 1] + [-1] * 18 + [2],
@@ -4036,6 +4106,8 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=page_size, cp_size=2, cp_rank=0),
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
current_slot_spans=[(1, 2)],
slot_logical_pages=torch.tensor([1, 20], dtype=torch.int64),
page_inverse=torch.tensor(
[-1, 1] + [-1] * 18 + [2],
@@ -5620,7 +5692,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=[],
layer_id=0,
@@ -5997,7 +6068,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=0,
@@ -6056,7 +6126,7 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=1,
prefix_slot_spans=[(0, 1)],
current_slot_spans=[(1, 3)],
layer_id=0,
)
@@ -6180,7 +6250,6 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=0,