Support topk > 1 tree drafting for mamba/hybrid-linear models on spec v2 (#27463)
This commit is contained in:
@@ -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
|
||||
):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
+10
-11
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user