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