From 5f79cf35110d6a0be828f266160b75d83a2a6276 Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Wed, 5 Aug 2026 11:40:28 -0700 Subject: [PATCH] [DCP] Match the replicated draft KV pool's page granularity to its allocator (#33348) --- .../srt/mem_cache/kv_cache_configurator.py | 56 +++++++++++-------- 1 file changed, 34 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index acd0f4137..221c73dcc 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -270,6 +270,19 @@ class KVCacheConfigurator: unified_memory_pool=pools.unified_memory_pool, ) + # Note(kpham-sgl): + # 1. A replicated draft indexes the allocator's virtual locs raw, so its pools + # span and page that space; the sharded target translates and stays per-rank. + # 2. A pool must page as its allocator does, or its last page falls short. + @property + def loc_space_scale(self) -> int: + dcp_size = self.server_args.dcp_size + return dcp_size if (self.is_draft_worker and dcp_size > 1) else 1 + + @property + def pool_page_size(self) -> int: + return get_schedule().page_size * self.loc_space_scale + def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes: max_total_num_tokens = config.max_total_num_tokens max_running_requests = config.max_running_requests @@ -281,13 +294,12 @@ class KVCacheConfigurator: # Draft pools are replicated, not DCP-sharded, yet consume the shared # allocator's virtual locs in [0, max_total * dcp_size) untranslated. - dcp_size = self.server_args.dcp_size - if self.is_draft_worker and dcp_size > 1: - max_total_num_tokens *= dcp_size - if full_max_total_num_tokens is not None: - full_max_total_num_tokens *= dcp_size - if swa_max_total_num_tokens is not None: - swa_max_total_num_tokens *= dcp_size + loc_scale = self.loc_space_scale + max_total_num_tokens *= loc_scale + if full_max_total_num_tokens is not None: + full_max_total_num_tokens *= loc_scale + if swa_max_total_num_tokens is not None: + swa_max_total_num_tokens *= loc_scale # DSV4 compressed-attention pool sizes. Draft worker reuses target's # full/swa sizes but does NOT own c4/c128/state pools (those live on @@ -1000,7 +1012,7 @@ class KVCacheConfigurator: c128_size=c128_max_total_num_tokens, c4_state_pool_size=c4_state_pool_size, c128_state_pool_size=c128_state_pool_size, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, swa_page_size=swa_page_size, sliding_window=self.model_config.window_size, dtype=self.kv_cache_dtype, @@ -1026,7 +1038,7 @@ class KVCacheConfigurator: PoolCls = current_platform.get_dsa_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1050,7 +1062,7 @@ class KVCacheConfigurator: PoolCls = current_platform.get_mla_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1067,7 +1079,7 @@ class KVCacheConfigurator: PoolCls = current_platform.get_mha_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1104,7 +1116,7 @@ class KVCacheConfigurator: token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=swa_max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1126,7 +1138,7 @@ class KVCacheConfigurator: token_to_kv_pool = NPUMLATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1146,7 +1158,7 @@ class KVCacheConfigurator: token_to_kv_pool = NPUMHATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1186,7 +1198,7 @@ class KVCacheConfigurator: PoolCls = DSATokenToKVPool token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1208,7 +1220,7 @@ class KVCacheConfigurator: def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPoolFP4( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1223,7 +1235,7 @@ class KVCacheConfigurator: def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPool( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1286,7 +1298,7 @@ class KVCacheConfigurator: token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=size_swa, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1311,7 +1323,7 @@ class KVCacheConfigurator: ) token_to_kv_pool = MiniMaxSparseKVPool( size=max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, # fp8 attn-GEMM mode opts the lightning-indexer cache into # fp8 too (fp8 indexer GEMMs); fp8 KV without the mode @@ -1368,7 +1380,7 @@ class KVCacheConfigurator: else mha_pool_class ) token_to_kv_pool = HybridLinearKVPool( - page_size=get_schedule().page_size, + page_size=self.pool_page_size, size=max_total_num_tokens, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1391,7 +1403,7 @@ class KVCacheConfigurator: def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MHATokenToKVPoolFP4( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, @@ -1424,7 +1436,7 @@ class KVCacheConfigurator: pool_kwargs["post_capture_active"] = self.post_capture_kv_active token_to_kv_pool = pool_cls( max_total_num_tokens, - page_size=get_schedule().page_size, + page_size=self.pool_page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim,