From 88c261c3f34cc5b05b35654e4defa6fe6c534bf1 Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Fri, 19 Jun 2026 23:39:09 +0800 Subject: [PATCH] Fix IndexCache PP topk handoff (#28532) --- python/sglang/srt/configs/model_config.py | 23 +++++ .../sglang/srt/managers/scheduler_pp_mixin.py | 7 ++ .../sglang/srt/model_executor/model_runner.py | 14 +++ .../runner/decode_cuda_graph_runner.py | 6 ++ .../model_executor/runner_utils/buffers.py | 5 ++ python/sglang/srt/models/deepseek_v2.py | 85 ++++++++++--------- python/sglang/srt/server_args.py | 2 +- 7 files changed, 99 insertions(+), 43 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index b51e1c3a5..be029153a 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -132,6 +132,29 @@ def get_dsa_index_topk(config: PretrainedConfig) -> int: return config.index_topk +def dsa_layer_skips_topk(config: PretrainedConfig, layer_id: int) -> bool: + """Return whether a DSA layer reuses the previous layer's top-k indices.""" + assert is_deepseek_dsa(config) + + pattern = getattr(config, "index_topk_pattern", None) + if pattern is not None: + return layer_id < len(pattern) and pattern[layer_id] == "S" + + freq = getattr(config, "index_topk_freq", 1) + if freq is None: + freq = 1 + assert freq > 0, f"index_topk_freq must be positive, got {freq}" + offset = getattr(config, "index_skip_topk_offset", None) + if offset is not None: + assert offset > 0, ( + "index_skip_topk_offset must be positive; offset <= 0 " + "marks layer 0 as skip_topk with no prior topk to reuse" + ) + return max(layer_id - offset + 1, 0) % freq != 0 + + return max(layer_id - 1, 0) % freq != 0 + + def get_dsa_index_n_heads(config: PretrainedConfig) -> int: assert is_deepseek_dsa(config) return config.index_n_heads diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 8babb55cf..242afde9b 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -650,6 +650,13 @@ class SchedulerPPMixin: device=self.device, ), } + pp_proxy_topk_size = model_runner.get_pp_proxy_topk_size() + if pp_proxy_topk_size is not None: + proxy_tensors["topk_indices"] = torch.zeros( + (current_seq_len, pp_proxy_topk_size), + dtype=torch.int32, + device=self.device, + ) pp_proxy = PPProxyTensors(proxy_tensors) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index fe0c94d70..aca970269 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -61,7 +61,9 @@ from sglang.srt.configs.model_config import ( AttentionArch, ModelConfig, ModelImpl, + dsa_layer_skips_topk, get_num_indexer_layers, + is_deepseek_dsa, ) from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS @@ -839,6 +841,17 @@ class ModelRunner(ModelRunnerKVCacheMixin): cpu_group=get_world_group().cpu_group, ) + def get_pp_proxy_topk_size(self) -> Optional[int]: + hf_config = self.model_config.hf_text_config + if ( + self.pp_size <= 1 + or self.pp_rank == 0 + or not is_deepseek_dsa(hf_config) + or not dsa_layer_skips_topk(hf_config, self.start_layer) + ): + return None + return getattr(hf_config, "index_topk", None) + def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None): """Allocate KV cache memory pools only (no backends or cuda graphs).""" if memory_pool_config is not None: @@ -2758,6 +2771,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): cache_loc_dtype=torch.int64, enable_mamba_track=False, hc_hidden_size=getattr(self.model_config, "hc_hidden_size", None), + pp_proxy_topk_size=self.get_pp_proxy_topk_size(), ) buffers.num_token_non_padded[...] = num_tokens diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index ccfe8135d..163c04805 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -182,6 +182,7 @@ def _allocate_decode_buffers( enable_mamba_track: bool, ne_token_table: Optional[torch.Tensor] = None, hc_hidden_size: Optional[int] = None, + pp_proxy_topk_size: Optional[int] = None, ) -> SimpleNamespace: """Allocate the FB-shared decode buffers as a namespace adopted by ``build_decode_registry(source=...)``.""" @@ -220,6 +221,10 @@ def _allocate_decode_buffers( pp_proxy_tensors["residual"] = torch.zeros( (max_bs, hidden_size), dtype=dtype ) + if pp_proxy_topk_size is not None: + pp_proxy_tensors["topk_indices"] = torch.zeros( + (max_num_token, pp_proxy_topk_size), dtype=torch.int32 + ) else: pp_proxy_tensors = None @@ -450,6 +455,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): hc_hidden_size=getattr( self.model_runner.model_config, "hc_hidden_size", None ), + pp_proxy_topk_size=self.model_runner.get_pp_proxy_topk_size(), ) self.buffers.share_buffers() # FB-shared slot registry adopting DecodeInputBuffers storage (same diff --git a/python/sglang/srt/model_executor/runner_utils/buffers.py b/python/sglang/srt/model_executor/runner_utils/buffers.py index f78e388e1..a9d0cbd35 100644 --- a/python/sglang/srt/model_executor/runner_utils/buffers.py +++ b/python/sglang/srt/model_executor/runner_utils/buffers.py @@ -107,6 +107,7 @@ class DecodeInputBuffers(ForwardInputBuffers): ne_token_table: Optional[torch.Tensor] = None, is_hybrid_swa: bool = False, hc_hidden_size: Optional[int] = None, + pp_proxy_topk_size: Optional[int] = None, ) -> DecodeInputBuffers: with torch.device(device): input_ids = torch.zeros((max_num_token,), dtype=torch.int64) @@ -149,6 +150,10 @@ class DecodeInputBuffers(ForwardInputBuffers): pp_proxy_tensors["residual"] = torch.zeros( (max_bs, hidden_size), dtype=dtype ) + if pp_proxy_topk_size is not None: + pp_proxy_tensors["topk_indices"] = torch.zeros( + (max_num_token, pp_proxy_topk_size), dtype=torch.int32 + ) else: pp_proxy_tensors = None diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 4468090ef..1787deeae 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -40,6 +40,7 @@ from sglang.srt.batch_overlap.two_batch_overlap import ( ) from sglang.srt.configs.model_config import ( compute_mla_mscale_scaling, + dsa_layer_skips_topk, get_dsa_index_head_dim, get_dsa_index_n_heads, get_dsa_index_topk, @@ -1604,40 +1605,8 @@ class DeepseekV2AttentionMLA( self.skip_topk = True self.next_skip_topk = True else: - self.index_topk_freq = getattr(config, "index_topk_freq", 1) - self.index_topk_pattern = getattr(config, "index_topk_pattern", None) - self.index_skip_topk_offset = getattr( - config, "index_skip_topk_offset", None - ) - if ( - self.index_topk_pattern is None - and self.index_skip_topk_offset is not None - ): - assert self.index_skip_topk_offset > 0, ( - "index_skip_topk_offset must be positive; offset <= 0 " - "marks layer 0 as skip_topk with no prior topk to reuse" - ) - self.skip_topk = ( - max(layer_id - self.index_skip_topk_offset + 1, 0) - % self.index_topk_freq - != 0 - ) - self.next_skip_topk = ( - max(layer_id - self.index_skip_topk_offset + 2, 0) - % self.index_topk_freq - != 0 - ) - elif self.index_topk_pattern is None: - self.skip_topk = max(layer_id - 1, 0) % self.index_topk_freq != 0 - self.next_skip_topk = layer_id % self.index_topk_freq != 0 - else: - self.skip_topk = self.index_topk_pattern[layer_id] == "S" - if layer_id < len(self.index_topk_pattern) - 1: - self.next_skip_topk = ( - self.index_topk_pattern[layer_id + 1] == "S" - ) - else: - self.next_skip_topk = False + self.skip_topk = dsa_layer_skips_topk(config, layer_id) + self.next_skip_topk = dsa_layer_skips_topk(config, layer_id + 1) self.kv_b_proj = ColumnParallelLinear( self.kv_lora_rank, @@ -2290,13 +2259,15 @@ class DeepseekV2Model(nn.Module): prefix: str = "", ) -> None: super().__init__() + self.config = config + self.use_dsa = is_deepseek_dsa(config) self.padding_id = config.pad_token_id self.vocab_size = config.vocab_size self.first_k_dense_replace = config.first_k_dense_replace self.pp_group = get_pp_group() self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() self.mla_enable_prefill_cp = ( - is_prefill_context_parallel_enabled() and not is_deepseek_dsa(config) + is_prefill_context_parallel_enabled() and not self.use_dsa ) if self.dsa_enable_prefill_cp or self.mla_enable_prefill_cp: self.cp_size = get_parallel().attn_cp_size @@ -2435,6 +2406,17 @@ class DeepseekV2Model(nn.Module): assert pp_proxy_tensors is not None hidden_states = pp_proxy_tensors["hidden_states"] residual = pp_proxy_tensors["residual"] + topk_indices = pp_proxy_tensors.tensors.get("topk_indices") + assert not ( + not forward_batch.forward_mode.is_idle() + and hidden_states.shape[0] != 0 + and self.use_dsa + and dsa_layer_skips_topk(self.config, self.start_layer) + and topk_indices is None + ), ( + f"PP stage starting at layer {self.start_layer} requires DSA " + "topk_indices from the previous stage." + ) device = hidden_states.device zero_allocator = BumpAllocator( buffer_size=total_num_layers * 2 * (2 if forward_batch.can_run_tbo else 1), @@ -2487,7 +2469,8 @@ class DeepseekV2Model(nn.Module): elif self.first_k_dense_replace < normal_start_layer: normal_end_layer = normal_start_layer = 0 aux_hidden_states = [] - topk_indices = None + if self.pp_group.is_first_rank: + topk_indices = None for i in range(normal_start_layer, normal_end_layer): # NOTE: torch dynamo does not support graph break in context manager ctx = ( @@ -2526,12 +2509,30 @@ class DeepseekV2Model(nn.Module): ) if not self.pp_group.is_last_rank: - return PPProxyTensors( - { - "hidden_states": hidden_states, - "residual": residual, - } - ) + proxy_tensors = { + "hidden_states": hidden_states, + "residual": residual, + } + if ( + self.use_dsa + and self.end_layer < self.config.num_hidden_layers + and dsa_layer_skips_topk(self.config, self.end_layer) + ): + if ( + not forward_batch.forward_mode.is_idle() + and hidden_states.shape[0] != 0 + ): + assert topk_indices is not None, ( + f"PP stage ending at layer {self.end_layer} must forward " + "DSA topk_indices because the next stage starts on a " + "skip-topk layer." + ) + if topk_indices is None: + topk_indices = hidden_states.new_empty( + (0, get_dsa_index_topk(self.config)), dtype=torch.int32 + ) + proxy_tensors["topk_indices"] = topk_indices + return PPProxyTensors(proxy_tensors) else: if not forward_batch.forward_mode.is_idle(): if residual is None: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 3b5cf53dc..fc1f61596 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2007,7 +2007,7 @@ class ServerArgs: self.attention_backend = "dsa" logger.info("Use dsa attention backend for DeepSeek with DSA.") - index_topk_freq = getattr(hf_config, "index_topk_freq", 1) + index_topk_freq = getattr(hf_config, "index_topk_freq", 1) or 1 index_topk_pattern = getattr(hf_config, "index_topk_pattern", None) if self.enable_two_batch_overlap and ( index_topk_freq > 1