Surface CP shared-KV fast-path misses

CP shared-KV performance debugging depends on seeing when the runtime leaves the intended TAI, IPC, prefetch, or current-reuse paths. This change makes those misses visible through standardized warning markers while keeping per-reason log limits to avoid per-layer log floods.\n\nThe warnings intentionally distinguish fallback from fail-fast: unsupported correctness-sensitive states still raise, while performance-path misses emit [CP_SHARED_KV_FALLBACK] with the concrete reason.\n\nConstraint: Production ETE debugging needs visible fallback evidence without enabling heavy debug mode, which can itself disable fast paths.\nRejected: Rely only on optional MLA prefetch debug logs | they are env-gated, layer-limited, and miss non-prefetch TAI/IPC/current-reuse fallbacks.\nRejected: Log every per-layer event without limits | would drown useful transfer/cache diagnostics under steady-state traffic.\nConfidence: high\nScope-risk: moderate\nDirective: Do not remove these fallback warnings unless an equivalent low-noise observability path exists for every fast-path miss.\nTested: local py_compile for touched files; local git diff --check for touched files; remote g0034 py_compile and pytest for test_nsa_cp_utils.py, test_cp_shared_kv_layout.py, test_cp_shared_kv_runtime.py passed before commit (156 passed, 5 warnings, 2 subtests passed).\nNot-tested: full ETE serving traffic after warning additions.
This commit is contained in:
laoyao0822
2026-06-02 06:47:40 +08:00
parent 6a25c312c7
commit 8be4a3a8b5
4 changed files with 489 additions and 20 deletions
@@ -885,7 +885,9 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
runtime._CURRENT_REUSE_FALLBACK_LOG_COUNTS.clear()
runtime._TAI_MATERIALIZE_FALLBACK_LOG_COUNTS.clear()
runtime._TAI_IPC_MATERIALIZE_FALLBACK_LOG_COUNTS.clear()
runtime._TAI_FUSED_MLA_STORE_FALLBACK_LOG_COUNTS.clear()
runtime._TAI_INDEX_MQA_PREPARE_FALLBACK_LOG_COUNTS.clear()
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
runtime._log_current_reuse_fallback(
@@ -903,6 +905,16 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
"store fallback %s",
3,
)
runtime._log_tai_ipc_materialize_fallback(
"ipc_reason",
"ipc fallback %s",
4,
)
runtime._log_tai_index_mqa_prepare_fallback(
"index_reason",
"index fallback %s",
5,
)
self.assertIn("[CP_SHARED_KV_FALLBACK][current_reuse]", logs.output[0])
self.assertIn("current_reason", logs.output[0])
@@ -913,6 +925,115 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIn("[CP_SHARED_KV_FALLBACK][tai_fused_mla_store]", logs.output[2])
self.assertIn("store_reason", logs.output[2])
self.assertIn("store fallback 3", logs.output[2])
self.assertIn("[CP_SHARED_KV_FALLBACK][tai_ipc_materialize]", logs.output[3])
self.assertIn("ipc_reason", logs.output[3])
self.assertIn("ipc fallback 4", logs.output[3])
self.assertIn("[CP_SHARED_KV_FALLBACK][tai_index_mqa_prepare]", logs.output[4])
self.assertIn("index_reason", logs.output[4])
self.assertIn("index fallback 5", logs.output[4])
def test_current_reuse_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
runtime._CURRENT_REUSE_FALLBACK_LOG_COUNTS.clear()
forward_batch = SimpleNamespace(
forward_mode=_FakeExtendForwardMode(),
batch_size=2,
extend_prefix_lens_cpu=[64, 64],
extend_seq_lens_cpu=[64, 64],
seq_lens_cpu=torch.tensor([128, 128], dtype=torch.int32),
out_cache_loc=torch.arange(128, dtype=torch.int64),
)
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))
joined = "\n".join(logs.output)
self.assertIn("[CP_SHARED_KV_FALLBACK][current_reuse]", joined)
self.assertIn("batch_size_not_one", joined)
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
runtime._TAI_INDEX_MQA_PREPARE_FALLBACK_LOG_COUNTS.clear()
with envs.SGLANG_CP_SHARED_KV_FUSED_INDEX_MQA_PREPARE.override(True):
with patch.object(
runtime, "_load_tai_index_mqa_prepare_kernel", return_value=None
):
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
self.assertIsNone(
runtime.try_tai_prepare_cp_mqa_index(
index_buffer=torch.empty((4, 32), dtype=torch.uint8),
page_indices=torch.tensor([1], dtype=torch.int64),
kv_len=64,
valid_q_count=1,
ke_start=0,
page_size=64,
index_head_dim=128,
)
)
joined = "\n".join(logs.output)
self.assertIn("[CP_SHARED_KV_FALLBACK][tai_index_mqa_prepare]", joined)
self.assertIn("kernel_missing", joined)
def test_tai_ipc_token_materialize_fast_path_miss_logs_warning(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
runtime._TAI_IPC_MATERIALIZE_FALLBACK_LOG_COUNTS.clear()
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=0)
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
self.assertFalse(
runtime._try_tai_ipc_materialize_token_kv_page_slots_into(
kv_cache=torch.empty((16, 1), dtype=torch.float32),
dense_kv_cache=torch.empty((16, 1), dtype=torch.float32),
slot_logical_pages=torch.tensor([1, 2], dtype=torch.int64),
layout=layout,
page_size=4,
start_slot=1,
end_slot=2,
)
)
joined = "\n".join(logs.output)
self.assertIn("[CP_SHARED_KV_FALLBACK][tai_ipc_materialize]", joined)
self.assertIn("start_slot_nonzero", joined)
def test_mla_prefetch_sync_compose_paths_log_warning_in_source(self):
from pathlib import Path
source = (
Path(__file__).resolve().parents[4]
/ "python/sglang/srt/layers/attention/nsa_backend.py"
).read_text()
self.assertIn("def _log_cp_shared_kv_mla_prefetch_fallback", source)
self.assertIn("[CP_SHARED_KV_FALLBACK][mla_prefetch]", source)
self.assertIn(
"MLA prefetch fast path did not provide partial-current", source
)
self.assertIn(
"MLA prefetch fast path did not provide full materialize", source
)
def test_index_prefetch_sync_compose_paths_log_warning_in_source(self):
from pathlib import Path
source = (
Path(__file__).resolve().parents[4]
/ "python/sglang/srt/layers/attention/nsa/nsa_indexer.py"
).read_text()
self.assertIn("[CP_SHARED_KV_FALLBACK][index_prefetch]", source)
self.assertIn("current_missing_prefetcher", source)
self.assertIn("current_consume_miss", source)
self.assertIn("missing_prefetcher", source)
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 (
@@ -1881,7 +2002,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
layout=layout,
page_size=4,
)
warning.assert_not_called()
self.assertFalse(
any(
"[CP_SHARED_KV_FALLBACK][slot_remap_cache]" in str(call)
for call in warning.call_args_list
)
)
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
remap_b = runtime.get_or_build_shared_token_kv_slot_remap(
@@ -1993,7 +2119,12 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
logical_pages=logical_pages_a,
layout=layout,
)
warning.assert_not_called()
self.assertFalse(
any(
"[CP_SHARED_KV_FALLBACK][slot_remap_cache]" in str(call)
for call in warning.call_args_list
)
)
with self.assertLogs(runtime.logger.name, level="WARNING") as logs:
remap_b = runtime.get_or_build_shared_paged_buffer_slot_remap(