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