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 cb83a13e9..fc3fd6c19 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 @@ -91,6 +91,7 @@ class DeepseekMLAForwardMixin: forward_batch: ForwardBatch, zero_allocator: BumpAllocator, llama_4_scaling: Optional[torch.Tensor] = None, + prev_topk_indices: Optional[torch.Tensor] = None, ): from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode @@ -182,18 +183,7 @@ class DeepseekMLAForwardMixin: q = self.q_b_proj(q)[0].view( -1, self.num_local_heads, self.qk_head_dim ) - topk_indices = self.indexer( - x=hidden_states, - q_lora=q_lora, - positions=positions, - forward_batch=forward_batch, - layer_id=self.layer_id, - ) - current_stream.wait_stream(self.alt_stream) - else: - 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: topk_indices = self.indexer( x=hidden_states, q_lora=q_lora, @@ -201,6 +191,23 @@ class DeepseekMLAForwardMixin: forward_batch=forward_batch, layer_id=self.layer_id, ) + else: + topk_indices = prev_topk_indices + current_stream.wait_stream(self.alt_stream) + else: + 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: + topk_indices = self.indexer( + x=hidden_states, + q_lora=q_lora, + positions=positions, + forward_batch=forward_batch, + layer_id=self.layer_id, + ) + else: + topk_indices = prev_topk_indices else: q = self.q_proj(hidden_states)[0].view( -1, self.num_local_heads, self.qk_head_dim @@ -557,7 +564,14 @@ class DeepseekMLAForwardMixin: ) output, _ = self.o_proj(attn_bmm_output) - return output + if self.next_skip_topk is None: + return output + + # Return topk_indices for the next layer when enabling index cache + if not self.next_skip_topk: + return output, None + else: + return output, topk_indices def _fuse_rope_for_trtllm_mla( self: DeepseekV2AttentionMLA, forward_batch: ForwardBatch diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 28029a0c7..f83db2d47 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -185,7 +185,7 @@ class DeepseekModelNextN(nn.Module): positions = cp_split_and_rebuild_position(forward_batch, positions) residual = None with get_global_expert_distribution_recorder().disable_this_region(): - hidden_states, residual = self.decoder( + hidden_states, residual, topk_indices = self.decoder( positions, hidden_states, forward_batch, diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index a0a298d22..87d751465 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1105,6 +1105,7 @@ class DeepseekV2AttentionMLA( prefix: str = "", alt_stream: Optional[torch.cuda.Stream] = None, skip_rope: bool = False, + is_nextn: bool = False, ) -> None: super().__init__() self.layer_id = layer_id @@ -1174,6 +1175,8 @@ class DeepseekV2AttentionMLA( prefix=add_prefix("kv_a_proj_with_mqa", prefix), ) + self.skip_topk = None + self.next_skip_topk = None if self.use_nsa: is_neox_style = not getattr(config, "indexer_rope_interleave", False) self.indexer = Indexer( @@ -1194,6 +1197,26 @@ class DeepseekV2AttentionMLA( layer_id=layer_id, alt_stream=alt_stream, ) + # Refer: https://arxiv.org/abs/2603.12201 for more details. + # 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 + 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.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.kv_b_proj = ColumnParallelLinear( self.kv_lora_rank, @@ -1325,9 +1348,14 @@ class DeepseekV2AttentionMLA( ) def op_core(self, state): - state.hidden_states_after_attn = self.forward_core( - state.pop("attn_intermediate_state") - ) + result = self.forward_core(state.pop("attn_intermediate_state")) + # forward_core may return (hidden_states, topk_indices) for NSA models + # with index cache enabled. In the TBO path, topk_indices is not + # propagated between layers, so we discard it here. + if isinstance(result, tuple): + state.hidden_states_after_attn = result[0] + else: + state.hidden_states_after_attn = result def forward( self, @@ -1337,6 +1365,7 @@ class DeepseekV2AttentionMLA( zero_allocator: BumpAllocator, layer_scatter_modes: LayerScatterModes = None, llama_4_scaling: Optional[torch.Tensor] = None, + prev_topk_indices: Optional[torch.Tensor] = None, ): s = self.forward_prepare( positions=positions, @@ -1345,6 +1374,7 @@ class DeepseekV2AttentionMLA( zero_allocator=zero_allocator, layer_scatter_modes=layer_scatter_modes, llama_4_scaling=llama_4_scaling, + prev_topk_indices=prev_topk_indices, ) return self.forward_core(s) @@ -1356,6 +1386,7 @@ class DeepseekV2AttentionMLA( zero_allocator: BumpAllocator, layer_scatter_modes: LayerScatterModes = None, llama_4_scaling: Optional[torch.Tensor] = None, + prev_topk_indices: Optional[torch.Tensor] = None, ): if self.attn_mha.kv_b_proj is None: self.attn_mha.kv_b_proj = self.kv_b_proj @@ -1395,7 +1426,12 @@ class DeepseekV2AttentionMLA( ) elif attn_forward_method == AttnForwardMethod.MLA: inner_state = self.forward_absorb_prepare( - positions, hidden_states, forward_batch, zero_allocator, llama_4_scaling + positions, + hidden_states, + forward_batch, + zero_allocator, + llama_4_scaling, + prev_topk_indices, ) elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE_ROCM: inner_state = self.forward_absorb_fused_mla_rope_prepare( @@ -1556,6 +1592,7 @@ class DeepseekV2DecoderLayer(nn.Module): reduce_results=False, prefix=add_prefix("self_attn", prefix), alt_stream=alt_stream, + is_nextn=is_nextn, ) if not hasattr(config, "q_lora_rank") and envs.SGLANG_USE_AG_AFTER_QLORA.get(): raise ValueError( @@ -1642,6 +1679,7 @@ class DeepseekV2DecoderLayer(nn.Module): zero_allocator: BumpAllocator, gemm_output_zero_allocator: BumpAllocator = None, llama_4_scaling: Optional[torch.Tensor] = None, + prev_topk_indices: Optional[torch.Tensor] = None, ) -> torch.Tensor: quant_format = ( "mxfp4" @@ -1684,7 +1722,12 @@ class DeepseekV2DecoderLayer(nn.Module): zero_allocator=zero_allocator, llama_4_scaling=llama_4_scaling, layer_scatter_modes=self.layer_scatter_modes, + prev_topk_indices=prev_topk_indices, ) + if isinstance(hidden_states, tuple): + hidden_states, topk_indices = hidden_states + else: + topk_indices = None hidden_states, residual = self.layer_communicator.prepare_mlp( hidden_states, residual, forward_batch @@ -1720,7 +1763,7 @@ class DeepseekV2DecoderLayer(nn.Module): hidden_states, residual, forward_batch ) - return hidden_states, residual + return hidden_states, residual, topk_indices def op_comm_prepare_attn( self, @@ -1997,6 +2040,7 @@ 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 for i in range(normal_start_layer, normal_end_layer): # NOTE: torch dynamo does not support graph break in context manager ctx = ( @@ -2014,7 +2058,7 @@ class DeepseekV2Model(nn.Module): else: aux_hidden_states.append(hidden_states + residual) layer = self.layers[i] - hidden_states, residual = layer( + hidden_states, residual, topk_indices = layer( positions, hidden_states, forward_batch, @@ -2022,6 +2066,7 @@ class DeepseekV2Model(nn.Module): zero_allocator, gemm_output_zero_allocator, llama_4_scaling, + prev_topk_indices=topk_indices, ) if normal_end_layer != self.end_layer: diff --git a/test/registered/8-gpu-models/test_deepseek_v32_indexcache.py b/test/registered/8-gpu-models/test_deepseek_v32_indexcache.py new file mode 100644 index 000000000..4b769ec57 --- /dev/null +++ b/test/registered/8-gpu-models/test_deepseek_v32_indexcache.py @@ -0,0 +1,117 @@ +import unittest +from types import SimpleNamespace + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + is_in_ci, + popen_launch_server, + write_github_step_summary, +) + +register_cuda_ci(est_time=360, suite="stage-c-test-8-gpu-h200") + +DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2" + + +class TestDeepseekV32IndexTopkPattern(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V32_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + other_args = [ + "--trust-remote-code", + "--tp", + "8", + "--enable-dp-attention", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + "--json-model-override-args", + '{"index_topk_pattern": "FFSFSSSFSSFFFSSSFFFSFSSSSSSFFSFFSFFSSFFFFFFSFFFFFSFFSSSSSSFSF"}', + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=other_args, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_a_gsm8k( + self, + ): # Append an "a" to make this test run first (alphabetically) to warm up the server + args = SimpleNamespace( + num_shots=20, + data_path=None, + num_questions=1400, + parallel=1400, + max_new_tokens=512, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_gsm8k (deepseek-v32)\n" f'{metrics["accuracy"]=:.3f}\n' + ) + self.assertGreater(metrics["accuracy"], 0.935) + + +class TestDeepseekV32IndexFreq(CustomTestCase): + @classmethod + def setUpClass(cls): + cls.model = DEEPSEEK_V32_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + other_args = [ + "--trust-remote-code", + "--tp", + "8", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + "--json-model-override-args", + '{"index_topk_freq": 4}', + ] + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=other_args, + ) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + def test_a_gsm8k( + self, + ): # Append an "a" to make this test run first (alphabetically) to warm up the server + args = SimpleNamespace( + num_shots=20, + data_path=None, + num_questions=1400, + parallel=1400, + max_new_tokens=512, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"{metrics=}") + + if is_in_ci(): + write_github_step_summary( + f"### test_gsm8k (deepseek-v32)\n" f'{metrics["accuracy"]=:.3f}\n' + ) + self.assertGreater(metrics["accuracy"], 0.935) + + +if __name__ == "__main__": + unittest.main()