diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 3a7f612fa..94d1e5a2e 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -212,6 +212,7 @@ class Envs: SGLANG_CP_SHARED_KV_ENABLE_MLA_PREFETCH = EnvBool(False) SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH = EnvBool(False) SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX = EnvBool(False) + SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS = EnvInt(1024) SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES = EnvInt(-1) SGLANG_CP_DRAFT_SHARED_KV = EnvBool(False) SGLANG_CP_DRAFT_SHARED_KV_DEBUG = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py index e2c05bc83..71c6a0f79 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py @@ -403,15 +403,19 @@ class CpSharedKVMlaPrefetcher: int(real_page_table.numel()), ) return None - min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages(layout.cp_size) + min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages( + layout.cp_size, page_size=page_size + ) if prefix_pages < min_prefix_pages: _prefetch_log( "create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " - "prefix_pages=%s min_prefix_pages=%s", + "prefix_pages=%s min_prefix_pages=%s prefix_len=%s page_size=%s", layout.cp_rank, layout.cp_size, prefix_pages, min_prefix_pages, + extend_prefix_len, + page_size, ) return None @@ -1061,15 +1065,19 @@ class CpSharedKVIndexPrefetcher: int(real_page_table.numel()), ) return None - min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages(layout.cp_size) + min_prefix_pages = cp_shared_kv_mla_prefetch_min_prefix_pages( + layout.cp_size, page_size=page_size + ) if prefix_pages < min_prefix_pages: _prefetch_log( "index_create_skip reason=prefix_below_min cp_rank=%s cp_size=%s " - "prefix_pages=%s min_prefix_pages=%s", + "prefix_pages=%s min_prefix_pages=%s prefix_len=%s page_size=%s", layout.cp_rank, layout.cp_size, prefix_pages, min_prefix_pages, + extend_prefix_len, + page_size, ) return None diff --git a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py index 64e1d32de..6d151318b 100644 --- a/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py +++ b/python/sglang/srt/layers/attention/nsa/cp_shared_kv_runtime.py @@ -20,6 +20,10 @@ _TAI_MATERIALIZE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _TAI_FUSED_MLA_STORE_FALLBACK_LOG_COUNTS: dict[str, int] = {} _SLOT_REMAP_CACHE_LOG_COUNTS: dict[str, int] = {} _MLA_PREFETCH_LOG_PROBE_LAYER = 2 +_MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS = max( + int(envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS.get()), + 0, +) _SORT_NVTX_ENABLED = envs.SGLANG_DEBUG_SORT_NVTX.get() _MATERIALIZE_NVTX_ENABLED = envs.SGLANG_CP_SHARED_KV_MATERIALIZE_NVTX.get() @@ -60,17 +64,33 @@ def cp_shared_kv_mla_prefetch_log_enabled() -> bool: return envs.SGLANG_CP_SHARED_KV_LOG_MLA_PREFETCH.get() -def cp_shared_kv_mla_prefetch_min_prefix_pages(cp_size: int) -> int: +def cp_shared_kv_mla_prefetch_min_prefix_pages( + cp_size: int, *, page_size: int | None = None +) -> int: """Minimum prefix pages required to enable Phase8 prefetch. - Negative env values mean "use cp_size" so the default skips tiny prefixes - that cannot cover all CP lanes. Set the env to 0 to disable the gate, or to - a positive absolute page count for workload-specific tuning. + Negative env values mean "use the dynamic default": at least one page per CP + lane and, when the runtime page size is known, at least 1K prefix tokens. + This keeps tiny cache-hit prefixes on the simpler synchronous materialize + path where prefix prefetch launch/collective overhead can dominate. Set the + env to 0 to disable the gate, or to a positive absolute page count for + workload-specific tuning. """ configured = envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES.get() if configured < 0: - return max(int(cp_size), 0) + min_pages = max(int(cp_size), 0) + if page_size is not None and int(page_size) > 0: + min_pages = max( + min_pages, + ( + _MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS + + int(page_size) + - 1 + ) + // int(page_size), + ) + return min_pages return max(int(configured), 0) diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index 2fe31e1c6..b50db01f5 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -718,23 +718,41 @@ 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_cp_size_and_can_override(self): + def test_mla_prefetch_min_prefix_pages_defaults_to_1k_tokens_and_can_override(self): from sglang.srt.environ import envs - from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import ( - cp_shared_kv_mla_prefetch_min_prefix_pages, - ) + 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(cp_shared_kv_mla_prefetch_min_prefix_pages(8), 8) + self.assertEqual(runtime._MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS, 1024) + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64), 16 + ) + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(32, page_size=64), 32 + ) with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES.override(0): - self.assertEqual(cp_shared_kv_mla_prefetch_min_prefix_pages(8), 0) + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64), 0 + ) with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES.override(16): - self.assertEqual(cp_shared_kv_mla_prefetch_min_prefix_pages(8), 16) + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64), + 16, + ) with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_PAGES.override(-2): - self.assertEqual(cp_shared_kv_mla_prefetch_min_prefix_pages(4), 4) + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(4, page_size=64), + 16, + ) + + with patch.object(runtime, "_MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS", 2048): + self.assertEqual( + runtime.cp_shared_kv_mla_prefetch_min_prefix_pages(8, page_size=64), + 32, + ) def test_fused_mla_store_uses_tai_kernel_when_enabled(self): from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime