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