From bef37d6de86ac842541a10f0e54bd446732f19c9 Mon Sep 17 00:00:00 2001 From: hlu1 <14827759+hlu1@users.noreply.github.com> Date: Fri, 7 Nov 2025 01:34:04 -0800 Subject: [PATCH] [Deepseek V3.2] Only skip Indexer logits computation when is_extend_without_speculative (#12816) Signed-off-by: Hao Lu <14827759+hlu1@users.noreply.github.com> --- .../srt/layers/attention/nsa/nsa_indexer.py | 11 ++++++---- .../srt/model_executor/forward_batch_info.py | 7 +++++++ python/sglang/srt/models/deepseek_v2.py | 20 ++++++------------- 3 files changed, 20 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 2b9cbcba9..c8903d279 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -410,6 +410,8 @@ class Indexer(CustomOp): metadata: BaseIndexerMetadata, return_indices: bool = True, ) -> Optional[torch.Tensor]: + assert forward_batch.forward_mode.is_extend_without_speculative() + # Fast path: only compute and store k cache, skip all q and weights ops key = self._get_k_bf16(x, positions, enable_dual_stream) k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) @@ -555,14 +557,15 @@ class Indexer(CustomOp): return None # Determine if should skip topk based on sequence length - should_skip = False - if not forward_batch.forward_mode.is_decode_or_idle(): + # We can only skip the logits computation if cuda graph is not involved + skip_logits_computation = False + if forward_batch.forward_mode.is_extend_without_speculative(): if forward_batch.seq_lens_cpu is not None: max_kv_len = forward_batch.seq_lens_cpu.max().item() - should_skip = max_kv_len <= self.index_topk + skip_logits_computation = max_kv_len <= self.index_topk # Optimization: fast path when skipping topk computation - if should_skip: + if skip_logits_computation: return self._forward_cuda_k_only( x, positions, diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 4f843726f..0c51b2b29 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -141,6 +141,13 @@ class ForwardMode(IntEnum): def is_split_prefill(self): return self == ForwardMode.SPLIT_PREFILL + def is_extend_without_speculative(self): + return ( + self.is_extend() + and not self.is_target_verify() + and not self.is_draft_extend() + ) + @total_ordering class CaptureHiddenMode(IntEnum): diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index a7ba1c82f..d9c4eca4d 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -308,14 +308,6 @@ def _get_sum_extend_prefix_lens(forward_batch): ) -def _is_extend_without_speculative(forward_batch): - return ( - forward_batch.forward_mode.is_extend() - and not forward_batch.forward_mode.is_target_verify() - and not forward_batch.forward_mode.is_draft_extend() - ) - - def _support_mha_one_shot(attn: DeepseekV2AttentionMLA, forward_batch, backend_name): attn_supported = backend_name in ["fa3", "flashinfer", "flashmla"] sum_seq_lens = ( @@ -334,7 +326,7 @@ def _handle_attention_backend( if ( not disable_ragged - and _is_extend_without_speculative(forward_batch) + and forward_batch.forward_mode.is_extend_without_speculative() and ( ( sum_extend_prefix_lens >= attn.chunked_prefix_cache_threshold @@ -377,7 +369,7 @@ def handle_attention_fa4(attn, forward_batch): def handle_attention_trtllm_mla(attn, forward_batch): sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) - if _is_extend_without_speculative(forward_batch) and ( + if forward_batch.forward_mode.is_extend_without_speculative() and ( not attn.disable_chunked_prefix_cache or sum_extend_prefix_lens == 0 ): return AttnForwardMethod.MHA_CHUNKED_KV @@ -386,7 +378,7 @@ def handle_attention_trtllm_mla(attn, forward_batch): def handle_attention_aiter(attn, forward_batch): - if _is_extend_without_speculative(forward_batch): + if forward_batch.forward_mode.is_extend_without_speculative(): if is_dp_attention_enabled(): if sum(forward_batch.extend_prefix_lens_cpu) == 0: return AttnForwardMethod.MHA @@ -411,7 +403,7 @@ def handle_attention_nsa(attn, forward_batch): if forward_batch.forward_mode.is_decode_or_idle(): return AttnForwardMethod.MLA - if _is_extend_without_speculative(forward_batch): + if forward_batch.forward_mode.is_extend_without_speculative(): assert forward_batch.seq_lens_cpu is not None max_kv_len = forward_batch.seq_lens_cpu.max().item() @@ -439,7 +431,7 @@ def handle_attention_triton(attn, forward_batch): return _dispatch_mla_subtype(attn, forward_batch) if ( - _is_extend_without_speculative(forward_batch) + forward_batch.forward_mode.is_extend_without_speculative() and sum(forward_batch.extend_prefix_lens_cpu) == 0 ): return AttnForwardMethod.MHA @@ -1517,7 +1509,7 @@ class DeepseekV2AttentionMLA(nn.Module): ) # NSA Indexer: cache quantized keys, auto-skip topk for sequences <= nsa_index_topk - if self.use_nsa and _is_extend_without_speculative(forward_batch): + if self.use_nsa: _ = self.indexer( x=hidden_states, q_lora=q_lora,