Stabilize CP shared-KV prefetch around draft cache hits

Cache-hit EAGLE/NextN draft extends can enter the draft DeepEP MoE immediately after CP shared-KV attention. The partial current-reuse path is kept for target layers, but draft cache-hit suffixes now use full materialization until draft has an explicit same-layer reuse contract. Next-layer MLA/index prefetch is also gated by the actual model depth, so the single-layer draft model does not enqueue unused next-layer async work.

The temporary stage traces used to isolate the hang are removed. The retained draft current-reuse fallback is a bounded warning because it changes the runtime path intentionally.

Constraint: EAGLE/NextN has one executable draft layer and mirrors target KV state.

Rejected: Keep partial current reuse for draft cache-hit suffixes | reproduced hangs at draft layer0 before DeepEP MoE completion.

Rejected: Keep temporary stage traces | useful for diagnosis but too noisy for normal runs.

Confidence: medium

Scope-risk: moderate

Directive: Do not re-enable draft cache-hit partial current reuse without an explicit draft same-layer reuse contract and ETE validation with CP shared KV + HiCache + EAGLE.

Tested: py_compile on edited Python files; git diff --check; temp trace grep returned no matches.

Not-tested: Local targeted pytest is blocked by missing pybase64 in this environment; full ETE after log cleanup not run.
This commit is contained in:
laoyao0822
2026-05-29 00:33:41 +08:00
parent 26c792939d
commit c3fc3ff752
7 changed files with 808 additions and 71 deletions
@@ -25,6 +25,14 @@ for _name in ("flash_attn_varlen_func", "flash_attn_with_kvcache"):
if not hasattr(flash_attn_stub, _name):
setattr(flash_attn_stub, _name, lambda *args, **kwargs: None)
sgl_kernel_stub = sys.modules.setdefault(
"sgl_kernel", types.ModuleType("sgl_kernel")
)
if not hasattr(sgl_kernel_stub, "__path__"):
sgl_kernel_stub.__path__ = []
if not hasattr(sgl_kernel_stub, "flash_attn"):
sgl_kernel_stub.flash_attn = flash_attn_stub
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
@@ -34,6 +42,11 @@ def _identity_all_reduce(buffer, *args, **kwargs):
return buffer
class _FakeExtendForwardMode:
def is_extend_without_speculative(self):
return True
class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def test_mla_prefetch_materializes_and_reduces_on_prefetch_stream(
self,
@@ -396,10 +409,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
is_current_only_extend_batch,
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
forward_batch = SimpleNamespace(
forward_mode=ForwardMode.EXTEND,
forward_mode=_FakeExtendForwardMode(),
extend_prefix_lens_cpu=[0, 0],
extend_seq_lens_cpu=[3, 5],
seq_lens_cpu=torch.tensor([3, 5], dtype=torch.int32),
@@ -414,6 +426,299 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
forward_batch.seq_lens_cpu = torch.tensor([4, 5], dtype=torch.int32)
self.assertFalse(is_current_only_extend_batch(forward_batch))
def test_can_reuse_current_extend_kv_allows_partial_cache_hit_single_batch(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
can_reuse_current_extend_kv,
)
forward_batch = SimpleNamespace(
forward_mode=_FakeExtendForwardMode(),
batch_size=1,
extend_seq_lens_cpu=[128],
seq_lens_cpu=torch.tensor([40384 + 128], dtype=torch.int32),
out_cache_loc=torch.arange(128, dtype=torch.int64),
)
self.assertTrue(can_reuse_current_extend_kv(forward_batch))
forward_batch.batch_size = 2
self.assertFalse(can_reuse_current_extend_kv(forward_batch))
forward_batch.batch_size = 1
forward_batch.out_cache_loc = torch.arange(127, dtype=torch.int64)
self.assertFalse(can_reuse_current_extend_kv(forward_batch))
def test_should_reuse_current_extend_kv_disables_draft_cache_hit_suffix(self):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
class DraftSpecInfo:
def is_draft_input(self):
return True
class TargetSpecInfo:
def is_draft_input(self):
return False
runtime._CURRENT_REUSE_FALLBACK_LOG_COUNTS.clear()
forward_batch = SimpleNamespace(
forward_mode=_FakeExtendForwardMode(),
batch_size=1,
extend_prefix_lens_cpu=[40384],
extend_seq_lens_cpu=[56],
seq_lens_cpu=torch.tensor([40384 + 56], dtype=torch.int32),
out_cache_loc=torch.arange(56, dtype=torch.int64),
spec_info=DraftSpecInfo(),
)
with envs.SGLANG_CP_SHARED_KV_CURRENT_REUSE.override(True):
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
self.assertFalse(runtime.should_reuse_current_extend_kv(forward_batch))
self.assertTrue(
any(
"draft_partial_current_reuse" in message
for message in logs.output
)
)
forward_batch.spec_info = TargetSpecInfo()
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
forward_batch.spec_info = DraftSpecInfo()
forward_batch.extend_prefix_lens_cpu = [0]
forward_batch.seq_lens_cpu = torch.tensor([56], dtype=torch.int32)
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
def test_current_loc_remap_fast_path_args_only_for_current_only_extend(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
current_loc_remap_fast_path_args,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
forward_batch = SimpleNamespace(
forward_mode=_FakeExtendForwardMode(),
batch_size=1,
extend_prefix_lens_cpu=[0],
extend_seq_lens_cpu=[128],
seq_lens_cpu=torch.tensor([128], dtype=torch.int32),
out_cache_loc=torch.arange(128, dtype=torch.int64),
token_to_kv_pool=SimpleNamespace(page_size=64, size=4096),
cp_shared_kv_layout=CpSharedKVLayout(page_size=64, cp_size=8, cp_rank=0),
)
self.assertEqual(current_loc_remap_fast_path_args(forward_batch), (64, 505))
forward_batch.extend_prefix_lens_cpu = [40389]
forward_batch.seq_lens_cpu = torch.tensor([40389 + 128], dtype=torch.int32)
self.assertEqual(current_loc_remap_fast_path_args(forward_batch), (None, None))
def test_merge_materialized_and_current_kv_remaps_only_current_locs(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
merge_materialized_and_current_kv,
)
materialized_kv = torch.arange(0, 8, dtype=torch.float32).view(8, 1, 1)
current_kv = torch.arange(100, 103, dtype=torch.float32).view(3, 1, 1)
logical_locs = torch.tensor([[4, 20, -1], [21, 7, 99]], dtype=torch.int32)
materialized_locs = torch.tensor([[4, -1, -1], [-1, 7, -1]], dtype=torch.int32)
current_locs = torch.tensor([20, 21, 22], dtype=torch.int64)
mixed_kv, mixed_locs, current_mask = merge_materialized_and_current_kv(
materialized_kv_cache=materialized_kv,
materialized_dense_locs=materialized_locs,
current_kv_cache=current_kv,
logical_locs=logical_locs,
current_locs=current_locs,
)
self.assertTrue(torch.equal(mixed_kv[:8], materialized_kv))
self.assertTrue(torch.equal(mixed_kv[8:], current_kv))
self.assertEqual(
current_mask.tolist(),
[[False, True, False], [True, False, False]],
)
self.assertEqual(mixed_locs.tolist(), [[4, 8, -1], [9, 7, -1]])
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
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeCurrentStream:
def __init__(self):
self.events = []
def wait_event(self, event):
self.events.append(event)
current_stream = FakeCurrentStream()
fake_event = object()
dense_kv = torch.arange(0, 16, dtype=torch.float32).view(16, 1, 1)
current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1)
page_inverse = torch.tensor([0, 1, 2, -1, -1, 3], dtype=torch.int64)
prefetcher = prefetch.CpSharedKVMlaPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
slot_logical_pages=torch.tensor([1, 2, 5], dtype=torch.int64),
page_inverse=page_inverse,
dense_num_pages=4,
stream=object(),
)
handle = prefetch.CpSharedKVMlaPrefetchHandle(
layer_id=1,
dense_kv_cache=dense_kv,
prefix_rows=slice(4, 12),
event=fake_event,
)
prefetcher.handles[1] = handle
prefetcher.pending_attention_handle = handle
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
), patch.object(
prefetch,
"materialize_local_token_kv_page_slots_into",
side_effect=AssertionError("suffix materialize must not run"),
):
mixed_kv, mixed_locs = prefetcher.consume_prefix_with_current(
layer_id=1,
kv_cache=torch.zeros((64, 1, 1), dtype=torch.float32),
logical_locs=torch.tensor([[4, 20], [21, 7]], dtype=torch.int32),
current_kv_cache=current_kv,
current_locs=torch.tensor([20, 21], dtype=torch.int64),
)
self.assertEqual(current_stream.events, [fake_event])
self.assertEqual(prefetcher.handles, {})
self.assertIsNone(prefetcher.pending_attention_handle)
self.assertTrue(torch.equal(mixed_kv[:16], dense_kv))
self.assertTrue(torch.equal(mixed_kv[16:], current_kv))
self.assertEqual(mixed_locs.tolist(), [[4, 16], [17, 7]])
def test_mla_prefetch_attention_window_waits_on_pending_event(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeCurrentStream:
def __init__(self):
self.events = []
def wait_event(self, event):
self.events.append(event)
current_stream = FakeCurrentStream()
fake_event = object()
prefetcher = prefetch.CpSharedKVMlaPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
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,
stream=object(),
)
handle = prefetch.CpSharedKVMlaPrefetchHandle(
layer_id=1,
dense_kv_cache=torch.zeros((16, 1, 1), dtype=torch.float32),
prefix_rows=slice(4, 12),
event=fake_event,
)
prefetcher.handles[1] = handle
prefetcher.pending_attention_handle = handle
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
):
prefetcher.wait_attention_window()
self.assertEqual(current_stream.events, [fake_event])
self.assertIsNone(prefetcher.pending_attention_handle)
self.assertIs(prefetcher.handles[1], handle)
def test_mla_prefetch_attention_window_launches_pending_reduce_before_wait(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeCurrentStream:
def __init__(self):
self.events = []
def wait_event(self, event):
self.events.append(event)
current_stream = FakeCurrentStream()
fake_event = object()
prefetcher = prefetch.CpSharedKVMlaPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
page_size=4,
prefix_pages=2,
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,
stream=object(),
)
handle = prefetch.CpSharedKVMlaPrefetchHandle(
layer_id=1,
dense_kv_cache=torch.zeros((16, 1, 1), dtype=torch.float32),
prefix_rows=slice(4, 12),
event=None,
)
prefetcher.handles[1] = handle
prefetcher.pending_attention_handle = handle
def finish_reduce():
handle.event = fake_event
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
), patch.object(
prefetcher, "launch_pending_reduce", side_effect=finish_reduce
) as launch_pending_reduce:
prefetcher.wait_attention_window()
launch_pending_reduce.assert_called_once_with()
self.assertEqual(current_stream.events, [fake_event])
self.assertIsNone(prefetcher.pending_attention_handle)
def test_index_prefetch_attention_window_waits_on_pending_event(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
class FakeCurrentStream:
def __init__(self):
self.events = []
def wait_event(self, event):
self.events.append(event)
current_stream = FakeCurrentStream()
fake_event = object()
prefetcher = prefetch.CpSharedKVIndexPrefetcher(
layout=CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0),
prefix_pages=2,
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,
stream=object(),
)
handle = prefetch.CpSharedKVIndexPrefetchHandle(
layer_id=1,
dense_page_buffer=torch.zeros((4, 3), dtype=torch.uint8),
prefix_rows=slice(1, 3),
event=fake_event,
)
prefetcher.handles[1] = handle
prefetcher.pending_attention_handle = handle
with patch.object(
prefetch.torch.cuda, "current_stream", return_value=current_stream
):
prefetcher.wait_attention_window()
self.assertEqual(current_stream.events, [fake_event])
self.assertIsNone(prefetcher.pending_attention_handle)
self.assertIs(prefetcher.handles[1], handle)
def test_materialize_local_token_kv_pages(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_dense_page_remap,
@@ -718,14 +1023,17 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
with envs.SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH.override(True):
self.assertTrue(cp_shared_kv_mla_prefetch_log_enabled())
def test_mla_prefetch_min_prefix_pages_defaults_to_1k_tokens_and_can_override(self):
def test_mla_prefetch_min_prefix_pages_uses_cached_token_default_and_can_override(self):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES.clear()
self.assertEqual(runtime._MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS, 1024)
default_tokens = envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS.get()
self.assertEqual(runtime._MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS, default_tokens)
expected_pages = (default_tokens + 63) // 64
self.assertEqual(
runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64), 16
runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64),
max(8, expected_pages),
)
self.assertEqual(
runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(32, page_size=64), 32
@@ -1823,6 +2131,50 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertEqual(fake_prefetcher.calls, [(12, token_to_kv_pool)])
def test_index_prefetch_skips_when_current_layer_is_last(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
cp_shared_kv_should_prefetch_next_layer,
)
class FakePrefetcher:
def __init__(self):
self.calls = []
def start_next_layer_prefix(self, *, next_layer_id, token_to_kv_pool):
self.calls.append((next_layer_id, token_to_kv_pool))
token_to_kv_pool = object()
fake_prefetcher = FakePrefetcher()
forward_batch = SimpleNamespace(
token_to_kv_pool=token_to_kv_pool,
cp_shared_kv_index_prefetcher=fake_prefetcher,
cp_shared_kv_num_model_layers=12,
)
indexer = object.__new__(nsa_indexer.Indexer)
self.assertFalse(cp_shared_kv_should_prefetch_next_layer(forward_batch, 11))
indexer._maybe_start_next_layer_index_prefetch(forward_batch, layer_id=11)
self.assertEqual(fake_prefetcher.calls, [])
def test_index_prefetch_skips_eagle_draft_next_layer(self):
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
cp_shared_kv_is_draft_input,
cp_shared_kv_should_prefetch_next_layer,
)
class FakeSpecInfo:
def is_draft_input(self):
return True
forward_batch = SimpleNamespace(
spec_info=FakeSpecInfo(),
)
self.assertTrue(cp_shared_kv_is_draft_input(forward_batch))
self.assertFalse(cp_shared_kv_should_prefetch_next_layer(forward_batch, 0))
def test_index_prefetch_consume_miss_logs_fallback_after_first_layer(self):
from sglang.srt.layers.attention.nsa import nsa_indexer