diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 1ad52b123..d0d36d7fe 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -701,6 +701,7 @@ class TboForwardBatchPreparer: "split_index", # for split prefill "orig_seq_lens", # only used by qwen-1m, thus not care "return_pooled_hidden_states", + "reuse_mtp_topk_indices", # forward-level flag, inherited by both child batches ]: output_dict[key] = getattr(batch, key) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 7754eb6b5..75feabdb2 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -385,6 +385,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # For hidden states before normal return_hidden_states_before_norm: bool = False + # For NSA/DSA topk_indices reuse across forward calls (e.g., EAGLE draft) + topk_indices: Optional[torch.Tensor] = None + reuse_mtp_topk_indices: Optional[bool] = False + # === Forward-derived (built in init_new on the forward stream; FB-owned) === # Position information positions: torch.Tensor = None diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 42bfa9213..54affa654 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -245,7 +245,15 @@ class DeepseekMLAForwardMixin: q = self.q_b_proj(q)[0].view( -1, self.num_local_heads, self.qk_head_dim ) - if not self.skip_topk or prev_topk_indices is None: + # skip_topk (shared) layers carry no indexer weights in the + # checkpoint, so they must reuse the carried topk and never run + # the indexer. Do NOT widen this to `or prev_topk_indices is + # None` (the upstream gate): that recomputes with an + # uninitialized indexer whenever cross-layer propagation is + # unavailable (e.g. the TBO op path drops topk_indices), + # reintroducing the >index_topk garbling. The is_nextn clause is + # the sole intentional fallback (layer 78 has its own weights). + if not self.skip_topk or (self.is_nextn and prev_topk_indices is None): topk_indices = self.indexer( x=hidden_states, q_lora=q_lora, @@ -264,7 +272,12 @@ class DeepseekMLAForwardMixin: k_nope = k_nope.unsqueeze(1) q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) if q_lora is not None: - if not self.skip_topk or prev_topk_indices is None: + # See the skip_topk note above: shared layers have no + # indexer weights, so this gate must not fall back to + # computing when prev_topk_indices is None. + if not self.skip_topk or ( + self.is_nextn and prev_topk_indices is None + ): topk_indices = self.indexer( x=hidden_states, q_lora=q_lora, diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 517789a8d..2dbe37b58 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -230,7 +230,14 @@ class DeepseekModelNextN(nn.Module): forward_batch, residual, zero_allocator, + prev_topk_indices=( + forward_batch.topk_indices + if forward_batch.reuse_mtp_topk_indices + else None + ), ) + if forward_batch.reuse_mtp_topk_indices: + forward_batch.topk_indices = topk_indices if not forward_batch.forward_mode.is_idle(): if residual is not None: diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index a7800d4b8..c5fc1113f 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1469,6 +1469,7 @@ class DeepseekV2AttentionMLA( self.q_lora_rank = q_lora_rank self.kv_lora_rank = kv_lora_rank self.quant_config = quant_config + self.is_nextn = is_nextn attn_tp_rank = get_attention_tp_rank() attn_tp_size = get_attention_tp_size() self.use_dsa = is_deepseek_dsa(config) @@ -1556,12 +1557,33 @@ class DeepseekV2AttentionMLA( # skip_topk: when True, this layer will skip computation and reuse previous layer's topk indices. # next_skip_topk: when True, the next layer will skip computation and reuse this layer's topk indices. if is_nextn: - self.skip_topk = False - self.next_skip_topk = False + 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) - if self.index_topk_pattern is 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: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 85cc92c0b..690f9d518 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1851,6 +1851,20 @@ 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_pattern = getattr(hf_config, "index_topk_pattern", None) + if self.enable_two_batch_overlap and ( + index_topk_freq > 1 + or (index_topk_pattern is not None and "S" in index_topk_pattern) + ): + raise ValueError( + "--enable-two-batch-overlap is not supported with DSA " + "index-topk sharing (index_topk_freq > 1 or an " + "index_topk_pattern containing shared layers): the TBO op " + "path does not propagate topk indices across layers, so " + "shared layers would run sparse attention without indices." + ) + if not is_npu() and not is_xpu(): # CUDA or ROCm GPU if self.enable_dsa_prefill_context_parallel: logger.warning( diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 246da3046..856c2839f 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -897,6 +897,17 @@ class EAGLEWorker(TpModelWorker): # Forward multiple steps scores = None + # Reuse NSA/DSA topk_indices from the first draft forward step for + # subsequent steps, analogous to skip_topk in deepseek_v2.py layers. + # Only safe with topk == 1: select_top_k_tokens reorders candidate rows + # each step, which would desync the cached indices from their rows. + index_share_for_mtp_iteration = ( + getattr(self.model_config.hf_config, "index_share_for_mtp_iteration", False) + and self.topk == 1 + ) + if index_share_for_mtp_iteration: + forward_batch.reuse_mtp_topk_indices = True + forward_batch.topk_indices = None for i in range(self.speculative_num_steps): input_ids, hidden_states, scores, tree_info = select_top_k_tokens( i, topk_p, topk_index, hidden_states, scores, self.topk @@ -949,6 +960,9 @@ class EAGLEWorker(TpModelWorker): maybe_detect_inf(hidden_states, f"draft_forward step {i}: hidden_states") forward_batch.positions.add_(1) + if index_share_for_mtp_iteration: + forward_batch.topk_indices = None + forward_batch.reuse_mtp_topk_indices = False parent_list, top_scores_index, draft_tokens = organize_draft_results( score_list, token_list, parents_list, self.speculative_num_draft_tokens ) diff --git a/python/sglang/srt/utils/hf_transformers/config.py b/python/sglang/srt/utils/hf_transformers/config.py index f85a25f94..f66ea9a67 100644 --- a/python/sglang/srt/utils/hf_transformers/config.py +++ b/python/sglang/srt/utils/hf_transformers/config.py @@ -67,6 +67,26 @@ class HfModelConfigParser(ModelConfigParserBase): **kwargs, ) + if ( + config.architectures is not None + and config.architectures[0] == "GlmMoeDsaForCausalLM" + ): + # GlmMoeDsaConfig drops/clobbers raw checkpoint fields the DSA path + # needs, so re-read them from config.json and restore. Fixed upstream + # by https://github.com/huggingface/transformers/pull/46338; remove + # this block once SGLang requires transformers >= 5.10. + from transformers import PretrainedConfig + + raw_config, _ = PretrainedConfig.get_config_dict(model, revision=revision) + for key in ( + "qk_rope_head_dim", + "index_topk_freq", + ): + if key in raw_config: + setattr(config, key, raw_config[key]) + if hasattr(config, "qk_head_dim") and hasattr(config, "qk_nope_head_dim"): + config.qk_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim + if ( config.architectures is not None and config.architectures[0] == "Phi4MMForCausalLM"