diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index d749efefb..b73f3d7d0 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -285,28 +285,8 @@ def _handle_eagle_family(server_args: "ServerArgs") -> None: "Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests." ) - # Spec v2 tree drafting supports topk > 1 with page_size == 1 and page_size > 1 - # (the latter via partial-page duplication; backend-gated below). spec_v1_reason = None - # mamba / linear-attn state models only support topk == 1 on spec v2. - # mamba2_cache_params exists iff the config carries such state; check the - # class descriptor so the property getter is not invoked. - text_config = server_args.get_model_config().hf_config.get_text_config() - is_mamba_state_model = hasattr(type(text_config), "mamba2_cache_params") if ( - server_args.speculative_eagle_topk is not None - and server_args.speculative_eagle_topk > 1 - and is_mamba_state_model - and not server_args.disable_overlap_schedule - ): - # Spec v2 topk > 1 is not supported for mamba/linear-attn state models - # (only topk == 1); fall back to v1 for those. page_size > 1 is supported - # on v2 (partial-page duplication), so it no longer forces v1. - server_args.disable_overlap_schedule = True - spec_v1_reason = ( - "spec v2 topk > 1 is not supported for mamba/linear-attn models" - ) - elif ( not envs.SGLANG_ENABLE_SPEC_V2.get() and not server_args.disable_overlap_schedule ): diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index 4a98de6cf..d20e50a97 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -38,6 +38,16 @@ class LightningAttentionBackend(MambaAttnBackendBase): def __init__(self, model_runner: ModelRunner): super().__init__(model_runner) + # seg_la processes draft tokens as a chain -- it has no parent-indices + # plumbing for tree-shaped drafts, so spec v2 tree verify (topk > 1) would + # commit wrong mamba states silently. Fail fast instead of mis-decoding. + if self.topk > 1: + raise NotImplementedError( + "Lightning (seg_la) linear-attention backend does not support " + f"speculative decoding with topk > 1 (got topk={self.topk}); " + "seg_la verifies a draft tree as a chain. Use " + "--speculative-eagle-topk 1." + ) # lightning attn does not need conv cache, but to keep the interface for mamba cache self.conv_states_shape = ( model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index afd732f2c..4911a1492 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -1485,6 +1485,7 @@ class HybridLinearKVPool(KVCache): device: str, mamba_pool: MambaPool, enable_memory_saver: bool = False, + enable_kv_cache_copy: bool = False, # TODO: refactor mla related args use_mla: bool = False, kv_lora_rank: int = None, @@ -1525,6 +1526,7 @@ class HybridLinearKVPool(KVCache): layer_num=self.full_layer_nums, device=device, enable_memory_saver=enable_memory_saver, + enable_kv_cache_copy=enable_kv_cache_copy, ) else: TokenToKVPoolClass = MLATokenToKVPool diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index 189b4a3b5..e57dfc003 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -664,6 +664,9 @@ class ModelRunnerKVCacheMixin: device=self.device, mamba_pool=self.req_to_token_pool.mamba_pool, enable_memory_saver=self.server_args.enable_memory_saver, + enable_kv_cache_copy=( + self.server_args.speculative_algorithm is not None + ), use_mla=self.use_mla_backend, start_layer=self.start_layer, **extra_args, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index cfab7dc12..8468c17b1 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -1276,9 +1276,7 @@ class EAGLEWorkerV2(BaseSpecWorker): or self.target_worker.model_runner.mamba2_config is not None or self.target_worker.model_runner.hybrid_lightning_config is not None ): - self._mamba_verify_update( - batch, verify_input, accept_lens, accept_index, bs - ) + self._mamba_verify_update(batch, accept_lens, accept_index, bs) if not batch.forward_mode.is_idle(): accept_tokens = predict[accept_index] @@ -1328,7 +1326,6 @@ class EAGLEWorkerV2(BaseSpecWorker): def _mamba_verify_update( self, batch: ScheduleBatch, - verify_input: EagleVerifyInput, accept_lens: torch.Tensor, accept_index: torch.Tensor, bs: int, @@ -1336,9 +1333,6 @@ class EAGLEWorkerV2(BaseSpecWorker): """Update mamba state for hybrid GDN models after verification.""" # `accept_lens` already includes the bonus token (drafts + 1 per req). if not batch.forward_mode.is_idle() and accept_index.numel() > 0: - if verify_input.topk != 1: - raise ValueError("Spec v2 currently only supports topk = 1.") - accepted_indices_offset = torch.arange( 0, bs * self.speculative_num_draft_tokens, @@ -1346,7 +1340,13 @@ class EAGLEWorkerV2(BaseSpecWorker): dtype=accept_lens.dtype, device=accept_lens.device, ) - last_correct_step_indices = accept_lens - 1 + req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device) + # Per-req tree step of the last accepted node, i.e. the step whose + # mamba state to commit; reduces to accept_lens - 1 for topk == 1. + last_correct_step_indices = ( + accept_index[req_idx, (accept_lens - 1).to(torch.int64)] + - accepted_indices_offset + ) if batch.mamba_track_indices is not None: # If after verify, the request's seq_lens has crossed a mamba track interval, @@ -1364,11 +1364,6 @@ class EAGLEWorkerV2(BaseSpecWorker): to_track_ith = torch.clamp( tracking_point - seq_lens_pre_verify - 1, min=0 ).to(torch.int64) - req_idx = torch.arange( - bs, - dtype=torch.int64, - device=accept_lens.device, - ) candidate_track_steps = ( accept_index[req_idx, to_track_ith] - accepted_indices_offset ) diff --git a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py index f86e8f955..6ddbd2661 100644 --- a/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py +++ b/python/sglang/test/kits/attention_unittest/attention_methods/mamba2_attention.py @@ -1003,22 +1003,18 @@ def expected_mamba2_verify_output_from_inputs( ) -> torch.Tensor: """Reference output for chain (topk=1) target-verify cases. - Mamba2's SSM kernel does not consume the tree mask: under any topk it - processes the per-request draft tokens linearly through the chunked-scan - recurrence, just like EXTEND. For `topk == 1` this matches the - chain semantics the EAGLE verifier expects, so the eager SSM - reference (`_pure_torch_mamba2_reference`) doubles as the verify - reference. For `topk > 1` the production kernel still processes - siblings as a chain — this is documented at the call site as - structurally unsupported rather than wired through a tree-aware - reference. + This reference (`_pure_torch_mamba2_reference`) is a chain recurrence. + For `topk == 1` it matches the chain semantics the EAGLE verifier + expects, so it doubles as the verify reference. For `topk > 1` the + production SSM kernel DOES follow the draft tree (it consumes the + parent-indices plumbing), but this test has no tree-aware reference to + compare against, so tree verify is skipped here rather than validated. """ if topk != 1: raise ValueError( - "Mamba2 tree verify (topk>1) is not exercised: the SSM kernel " - "ignores the tree mask and processes draft tokens linearly. " - "Wiring a parent-indices-aware reference here would not match " - "production behavior. Only chain (topk=1) is supported." + "Mamba2 tree verify (topk>1) is not exercised here: this " + "reference is chain-only. The production kernel supports tree " + "verify; a tree-aware reference is future work." ) del inputs # `state` is the (ssm_states, conv_states) snapshot captured before diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_target_verify_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_target_verify_runner.py index c543da2a3..29ea566cb 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_target_verify_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_target_verify_runner.py @@ -1196,19 +1196,18 @@ def run_mamba2_eagle_verify_case( atol: float = MAMBA2_ATOL, rtol: float = MAMBA2_RTOL, ): - """Mamba2 chain verify (eager). Mamba2's SSM kernel processes draft - tokens linearly regardless of the spec_info tree mask, so only - `topk == 1` is supported here. The EXTEND-style recurrence reference - (`_pure_torch_mamba2_reference`) doubles as the chain verify - reference across all chain spec kinds (eagle / frozen_kv_mtp / - dflash / ngram). Tree verify (topk > 1) is structurally blocked - (the kernel doesn't consume the parent-indices plumbing); see - `expected_mamba2_verify_output_from_inputs`.""" + """Mamba2 chain verify (eager). This test's reference + (`_pure_torch_mamba2_reference`) is a chain recurrence, so it can only + validate `topk == 1`; it doubles as the chain verify reference across + all chain spec kinds (eagle / frozen_kv_mtp / dflash / ngram). Tree + verify (topk > 1) is skipped only for lack of a tree-aware reference -- + the production SSM kernel does consume the parent-indices plumbing and + supports tree verify. See `expected_mamba2_verify_output_from_inputs`.""" if topk != 1: testcase.skipTest( - "Mamba2 tree verify (topk>1) is structurally unsupported — " - "the SSM kernel ignores tree masks; only chain (topk=1) is " - "exercised. See `expected_mamba2_verify_output_from_inputs`." + "Mamba2 tree verify (topk>1) skipped: this test has no tree-aware " + "reference. The production kernel supports tree verify. See " + "`expected_mamba2_verify_output_from_inputs`." ) fixture = build_mamba2_attention_fixture( testcase, diff --git a/test/registered/models_e2e/test_qwen3_next_models_mtp.py b/test/registered/models_e2e/test_qwen3_next_models_mtp.py index 6c8dc6110..f74d1f6af 100644 --- a/test/registered/models_e2e/test_qwen3_next_models_mtp.py +++ b/test/registered/models_e2e/test_qwen3_next_models_mtp.py @@ -14,6 +14,9 @@ QWEN3_NEXT_MODEL = "Qwen/Qwen3-Next-80B-A3B-Instruct" class TestQwen3NextMTPTopk( GSM8KMixin, KLDivergenceMixin, PrefixCacheBranchingMixin, DefaultServerBase ): + # topk > 1 (tree) MTP on a hybrid-GDN model, on spec v2: the tree-aware mamba + # state update lives in the spec v2 verify path, so mamba + topk > 1 no longer + # falls back to spec v1. model = QWEN3_NEXT_MODEL cache_chunk_size = 64 gsm8k_accuracy_thres = 0.93