Use runtime token widths for Triton speculative verification (#39859)

Co-authored-by: raghotham <853234+raghotham@users.noreply.github.com>
This commit is contained in:
Lianmin Zheng
2026-09-18 11:07:55 -07:00
committed by GitHub
co-authored by raghotham
parent 6bd1a0af1d
commit 248c202b46
10 changed files with 302 additions and 65 deletions
@@ -213,6 +213,7 @@ class TritonAttnBackend(AttentionBackend):
self.page_size = getattr(model_runner, "page_size", 1) or 1
self.kv_index_translator = model_runner.kv_index_translator
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
self.target_verify_num_tokens_per_req = model_runner.decode_num_tokens_per_req()
self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = get_spec().speculative_eagle_topk or 0
# Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1)
@@ -551,6 +552,12 @@ class TritonAttnBackend(AttentionBackend):
)
return kv_indptr, window_kv_indptr, window_kv_lens, num_kv_splits_lens
def _target_verify_num_tokens_per_req(self, spec_info: Optional[SpecInput]) -> int:
# Runtime metadata may vary by step; nonpositive means use capture width.
if spec_info is None or spec_info.num_tokens_per_req <= 0:
return self.target_verify_num_tokens_per_req
return spec_info.num_tokens_per_req
def _update_target_verify_buffers(
self,
bs: int,
@@ -559,19 +566,12 @@ class TritonAttnBackend(AttentionBackend):
req_pool_indices: torch.Tensor,
):
"""Fill all cuda-graph buffers for target_verify mode."""
# Prefer the spec_info's per-request query length (DSpark draft propose
# uses gamma < verify window); fall back to the configured verify window.
num_draft_tokens = self.num_draft_tokens
if (
spec_info is not None
and getattr(spec_info, "draft_token_num", None) is not None
):
num_draft_tokens = int(spec_info.draft_token_num)
num_tokens_per_req = self._target_verify_num_tokens_per_req(spec_info)
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * num_draft_tokens,
step=num_draft_tokens,
(1 + bs) * num_tokens_per_req,
step=num_tokens_per_req,
dtype=torch.int32,
device=self.device,
)
@@ -601,14 +601,11 @@ class TritonAttnBackend(AttentionBackend):
custom_mask = (
self._verify_mask.buffer if self._verify_mask is not None else None
)
if (
spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
):
if spec_info is not None and spec_info.custom_mask is not None:
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
else:
custom_mask = None
seq_mask_len = num_draft_tokens * (seq_lens + num_draft_tokens)
seq_mask_len = num_tokens_per_req * (seq_lens + num_tokens_per_req)
mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
return (
@@ -890,18 +887,11 @@ class TritonAttnBackend(AttentionBackend):
max_extend_len = None
elif forward_batch.forward_mode.is_target_verify():
bs = len(forward_batch.req_pool_indices)
# self.num_draft_tokens is the verify window (gamma + 1), while
# DSpark draft propose runs a gamma-token TARGET_VERIFY forward.
num_draft_tokens = self.num_draft_tokens
if (
spec_info is not None
and getattr(spec_info, "draft_token_num", None) is not None
):
num_draft_tokens = int(spec_info.draft_token_num)
num_tokens_per_req = self._target_verify_num_tokens_per_req(spec_info)
qo_indptr = torch.arange(
0,
(1 + bs) * num_draft_tokens,
step=num_draft_tokens,
(1 + bs) * num_tokens_per_req,
step=num_tokens_per_req,
dtype=torch.int32,
device=self.device,
)
@@ -938,13 +928,13 @@ class TritonAttnBackend(AttentionBackend):
)
custom_mask = spec_info.custom_mask
seq_mask_len = num_draft_tokens * (
forward_batch.seq_lens + num_draft_tokens
seq_mask_len = num_tokens_per_req * (
forward_batch.seq_lens + num_tokens_per_req
)
mask_indptr = self.mask_indptr
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0)
mask_indptr = mask_indptr[: bs + 1]
max_extend_len = num_draft_tokens
max_extend_len = num_tokens_per_req
num_kv_splits = None
attn_logits = None
attn_lse = None
@@ -1208,15 +1198,10 @@ class TritonAttnBackend(AttentionBackend):
self._verify_mask.buffer
if self._verify_mask is not None
and spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
and spec_info.custom_mask is not None
else None
)
max_extend_len = self.num_draft_tokens
if (
spec_info is not None
and getattr(spec_info, "draft_token_num", None) is not None
):
max_extend_len = int(spec_info.draft_token_num)
max_extend_len = self._target_verify_num_tokens_per_req(spec_info)
return ForwardMetadata(
attn_logits=None,
attn_lse=None,
@@ -2008,23 +1993,14 @@ class TritonAttnBackend(AttentionBackend):
# Compute window start positions (absolute position of first key in window)
# window_start_pos = seq_len - window_len
window_kv_lens = prefix_kv_indptr[1 : bs + 1] - prefix_kv_indptr[:bs]
# Handle TARGET_VERIFY mode where extend_prefix_lens might not be set
if forward_batch.extend_prefix_lens is not None:
window_start_pos = (
forward_batch.extend_prefix_lens[:bs] - window_kv_lens
)
elif forward_batch.forward_mode.is_target_verify():
window_start_pos = forward_batch.seq_lens[:bs] - window_kv_lens
else:
# Infer from spec_info: prefix_len = seq_len - draft_token_num
if forward_batch.spec_info is not None and hasattr(
forward_batch.spec_info, "draft_token_num"
):
extend_prefix_lens = (
forward_batch.seq_lens[:bs]
- forward_batch.spec_info.draft_token_num
)
window_start_pos = extend_prefix_lens - window_kv_lens
else:
window_start_pos = None
window_start_pos = None
else:
sliding_window_size = -1
prefix_kv_indptr = self.forward_metadata.kv_indptr
@@ -2047,29 +2023,23 @@ class TritonAttnBackend(AttentionBackend):
elif self.forward_metadata.out_cache_loc_full_physical is not None:
extend_kv_indices = self.forward_metadata.out_cache_loc_full_physical
# Handle cases where extend_seq_lens or extend_start_loc might not be set
# In speculative decoding, we can infer these from spec_info or compute them
# Capture batches may not have a spec_info, so use the attention
# metadata's resolved uniform verify width when extend lengths are absent.
if forward_batch.extend_seq_lens is None:
# TARGET_VERIFY mode: infer extend_seq_lens from spec_info
if forward_batch.spec_info is not None and hasattr(
forward_batch.spec_info, "draft_token_num"
):
draft_token_num = forward_batch.spec_info.draft_token_num
extend_seq_lens = torch.full(
(bs,), draft_token_num, dtype=torch.int32, device=self.device
)
else:
if not forward_batch.forward_mode.is_target_verify():
raise RuntimeError(
"extend_seq_lens is None but cannot infer from spec_info. "
"This should not happen in TARGET_VERIFY mode."
"extend_seq_lens is None outside TARGET_VERIFY mode."
)
extend_seq_lens = torch.full(
(bs,),
self.forward_metadata.max_extend_len,
dtype=torch.int32,
device=self.device,
)
else:
extend_seq_lens = forward_batch.extend_seq_lens
# Check extend_start_loc separately - it might be None even when extend_seq_lens is set
if forward_batch.extend_start_loc is None:
# Compute extend_start_loc from extend_seq_lens
# extend_start_loc[i] = sum(extend_seq_lens[0:i])
extend_start_loc = torch.cat(
[
torch.zeros(1, dtype=torch.int32, device=self.device),
+4 -2
View File
@@ -1,5 +1,5 @@
import logging
from dataclasses import dataclass
from dataclasses import dataclass, field
from typing import List, Optional
import torch
@@ -15,7 +15,9 @@ logger = logging.getLogger(__name__)
@dataclass
class EagleVerifyInput(SpecInput):
draft_token: torch.Tensor
custom_mask: torch.Tensor
# Keep this dataclass argument required despite SpecInput's None default;
# otherwise the required positions field would follow a defaulted field.
custom_mask: torch.Tensor = field()
positions: torch.Tensor
retrieve_index: torch.Tensor
retrieve_next_token: torch.Tensor
@@ -403,6 +403,10 @@ class SpecInput(ABC):
num_tokens_per_req: int = -1
num_tokens_for_logprob_per_req: int = -1
# Dataclasses assign fields before __post_init__ calls this base's __init__;
# assigning None there would overwrite the constructor's custom_mask.
custom_mask: Optional[torch.Tensor] = None
# DSA MTP IndexShare seed relay. Class-level defaults (same rationale as
# ragged_verify_layout) so scheduler/relay/attention code reads them
# uniformly on any SpecInput; only the EAGLE-family inputs override them.
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_context, get_parallel, get_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
@@ -227,6 +228,7 @@ class MockGDNModelRunner(ModelRunner):
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
@@ -230,6 +231,7 @@ class MockKDAModelRunner(ModelRunner):
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -30,6 +30,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
_parallel_override = get_parallel().override(attn_tp_size=1, attn_tp_rank=0)
_parallel_override.__enter__()
@@ -238,6 +239,7 @@ class MockLightningModelRunner(ModelRunner):
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -51,6 +51,7 @@ from sglang.srt.model_executor.forward_context import ( # noqa: E402
forward_context,
)
from sglang.srt.model_executor.model_runner import ModelRunner # noqa: E402
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm # noqa: E402
# Tiny dims chosen to be the minimum that satisfies MambaMixer2's TP/chunk asserts:
# - num_heads % tp_size == 0 (tp_size=1)
@@ -323,6 +324,7 @@ class MockMamba2ModelRunner(ModelRunner):
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
self.model_config = model_config
@@ -25,6 +25,7 @@ from sglang.srt.model_executor.forward_context import (
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
_parallel_override = get_parallel().override(attn_tp_size=1)
_parallel_override.__enter__()
@@ -246,6 +247,7 @@ class MockMLAModelRunner(ModelRunner):
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
speculative_num_draft_tokens = (
max(case.input_lens)
if case.forward_mode.is_target_verify()