[Spec] Single-source num_tokens_per_req derivation and access (#31013)

This commit is contained in:
Liangsheng Yin
2026-07-14 18:41:08 -07:00
committed by GitHub
parent ca0ee3f1a8
commit 76dc427806
23 changed files with 116 additions and 70 deletions
@@ -30,6 +30,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
from sglang.srt.runtime_context import get_server_args
from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.utils import get_compiler_backend
if TYPE_CHECKING:
@@ -191,10 +192,15 @@ class FlashAttentionBackend(AttentionBackend):
self.speculative_num_draft_tokens is not None
and model_runner.is_draft_worker
):
self.speculative_num_draft_tokens = SpeculativeAlgorithm.from_string(
model_runner.server_args.speculative_algorithm
).get_num_tokens_per_req_for_target_verify(
int(self.speculative_num_draft_tokens), is_draft_worker=True
# Static verify width; NOTE: overwrites the config-named attr in place.
self.speculative_num_draft_tokens = resolve_num_tokens_per_req(
phase="target_verify",
server_args=model_runner.server_args,
spec_algorithm=SpeculativeAlgorithm.from_string(
model_runner.server_args.speculative_algorithm
),
is_draft_worker=True,
num_draft_tokens=int(self.speculative_num_draft_tokens),
)
self.speculative_step_id = speculative_step_id
@@ -2289,7 +2295,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = num_tokens // bs
num_tokens_per_req = spec_info.num_tokens_per_req
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs
]
@@ -2693,9 +2699,7 @@ class FlashAttentionBackend(AttentionBackend):
device=device,
)
else:
default_extend = getattr(
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
)
default_extend = spec_info.num_tokens_per_req
extend_seq_lens = torch.full(
(bs,), default_extend, dtype=torch.int32, device=device
)
@@ -2704,9 +2708,7 @@ class FlashAttentionBackend(AttentionBackend):
if extend_seq_lens_cpu:
metadata.max_seq_len_q = int(max(extend_seq_lens_cpu))
else:
metadata.max_seq_len_q = getattr(
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
)
metadata.max_seq_len_q = spec_info.num_tokens_per_req
metadata.cu_seqlens_q[1:].copy_(
torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32)
@@ -1937,7 +1937,9 @@ class FlashInferIndicesUpdaterPrefill:
# host-known qo/kv layout from the caller. Assert rather than silently
# fall back to plan()'s blocking D2H on the replay hot-path.
paged_plan_kwargs = {}
num_tokens_per_req = getattr(spec_info, "num_tokens_per_req", None)
num_tokens_per_req = (
spec_info.num_tokens_per_req if spec_info is not None else None
)
uses_fast_prefill = (
hasattr(wrapper_paged.begin_forward, "func")
and wrapper_paged.begin_forward.func is fast_prefill_plan
@@ -487,7 +487,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
)
self.target_verify_metadata[bs] = metadata
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = num_tokens // bs
num_tokens_per_req = spec_info.num_tokens_per_req
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs
]
@@ -581,6 +581,7 @@ class CPUGraphRunner:
self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL
# Static capture width: CPU graphs are decode-only.
self.num_tokens_per_req = 1
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
@@ -163,6 +163,7 @@ from sglang.srt.server_args import ( # noqa: F401 (re-export)
set_global_server_args_for_scheduler,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.state_capturer.base import TopkCaptureOutput
from sglang.srt.state_capturer.indexer_topk import (
create_indexer_capturer,
@@ -622,10 +623,12 @@ class ModelRunner:
) -> int:
"""Logits rows per decode batch slot."""
if self.spec_algorithm.is_speculative():
if num_draft_tokens is None:
num_draft_tokens = self.server_args.speculative_num_draft_tokens
return self.spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, self.is_draft_worker
return resolve_num_tokens_per_req(
phase="target_verify",
server_args=self.server_args,
spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker,
num_draft_tokens=num_draft_tokens,
)
dllm_config = DllmConfig.from_server_args(self.server_args)
return dllm_config.block_size if dllm_config is not None else 1
@@ -269,12 +269,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> DecodeGraphCapture:
role = "draft" if model_runner.is_draft_worker else "target"
if model_runner.spec_algorithm.is_speculative():
capture_name = f"{role} verify"
num_tokens_per_req = (
model_runner.spec_algorithm.get_num_tokens_per_req_for_target_verify(
model_runner.server_args.speculative_num_draft_tokens,
model_runner.is_draft_worker,
)
)
num_tokens_per_req = model_runner.decode_num_tokens_per_req()
else:
capture_name = f"{role} decode"
num_tokens_per_req = 1
@@ -360,11 +360,7 @@ class BaseRunner(ABC):
if not mr.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen")
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
)
)
num_tokens_per_req = mr.decode_num_tokens_per_req()
if mr.server_args.enable_return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL
@@ -254,6 +254,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL
# Static capture width.
self.num_tokens_per_req = model_runner.decode_num_tokens_per_req(
num_draft_tokens=self.speculative_num_draft_tokens
)
@@ -83,10 +83,8 @@ class EagerRunner(BaseRunner):
),
)
else:
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, mr.is_draft_worker
)
num_tokens_per_req = mr.decode_num_tokens_per_req(
num_draft_tokens=num_draft_tokens
)
else:
dllm_config = DllmConfig.from_server_args(sa)
@@ -150,11 +148,7 @@ class EagerRunner(BaseRunner):
mr = self.model_runner
num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative():
num_tokens_per_req = (
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify(
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
)
)
num_tokens_per_req = mr.decode_num_tokens_per_req()
return (
self._alloc_dummy_decode_buffers(
self._eager_max_bs, num_tokens_per_req=num_tokens_per_req
@@ -283,6 +283,7 @@ class EagleDraftWorkerBase(ABC):
)
# Get a forward batch
# Actual width of the next draft-decode forward: topk tokens per req.
draft_input.num_tokens_per_req = topk
draft_input.num_tokens_for_logprob_per_req = topk
capture_mode = (
+2 -4
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional, Tuple
from typing import TYPE_CHECKING, Optional
import torch
@@ -48,9 +48,7 @@ class DFlashVerifyInput(SpecInput):
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
if self.num_tokens_per_req == -1:
self.num_tokens_per_req = int(self.draft_token_num)
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
self.num_tokens_for_logprob_per_req = int(self.draft_token_num)
def prepare_for_verify(
self,
@@ -2,7 +2,7 @@
import contextlib
from dataclasses import dataclass
from typing import Optional, Tuple
from typing import Optional
import torch
@@ -63,10 +63,9 @@ class DFlashDraftInputV2(SpecInput):
def __post_init__(self):
super().__init__(spec_input_type=SpecInputType.DFLASH_DRAFT)
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
# Spec v2 draft state itself does not change token accounting.
return (1, 1)
self.num_tokens_per_req = 1
self.num_tokens_for_logprob_per_req = 1
def _ensure_prepare_length_buffers(
self, bs: int, device: torch.device | str
@@ -39,6 +39,7 @@ from sglang.srt.runtime_context import get_flags
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.utils import (
require_attn_tp_gather,
require_gathered_buffer,
@@ -145,7 +146,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Bucket sizes
self.capture_bs, _ = get_batch_sizes_to_capture(model_runner)
self.num_tokens_per_req = self.topk
# Static capture width.
self.num_tokens_per_req = resolve_num_tokens_per_req(
phase="draft_decode", server_args=model_runner.server_args
)
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req
@@ -38,7 +38,10 @@ from sglang.srt.model_executor.runner_backend_utils import (
from sglang.srt.runtime_context import get_flags
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.speculative.spec_utils import fast_topk
from sglang.srt.speculative.spec_utils import (
fast_topk,
resolve_num_tokens_per_req,
)
from sglang.srt.utils import (
is_hip,
require_attn_tp_gather,
@@ -131,9 +134,11 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_bs, _ = get_batch_sizes_to_capture(model_runner)
# Size cuda-graph buffers by num_draft_tokens (full tree width), not
# num_steps + 1, or topk > 1 draft-extend overflows them.
self.num_tokens_per_req = model_runner.server_args.speculative_num_draft_tokens
# Static capture width: full tree width (num_draft_tokens), not
# num_steps + 1 -- topk > 1 draft-extend overflows the buffers.
self.num_tokens_per_req = resolve_num_tokens_per_req(
phase="draft_extend", server_args=model_runner.server_args
)
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req
+8 -9
View File
@@ -41,6 +41,14 @@ class EagleVerifyInput(SpecInput):
super().__init__(SpecInputType.EAGLE_VERIFY)
if self.num_tokens_per_req < 0:
self.num_tokens_per_req = self.draft_token_num
self.num_tokens_for_logprob_per_req = self.draft_token_num
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
# Keep this override on draft_token_num: eagle_worker_v2.verify()
# re-stamps num_tokens_per_req = num_steps + 1, which diverges from
# the real verify width for topk > 1 trees, and the DP-attention
# global-token scaling must follow the actual tree width.
return self.draft_token_num, self.draft_token_num
@property
def max_tree_depth(self) -> int:
@@ -55,9 +63,6 @@ class EagleVerifyInput(SpecInput):
irregular tree (no fixed per-level branching)."""
return self.topk
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
@classmethod
def create_idle_input(
cls, topk: int, spec_steps: int, num_verify_tokens: int, device: str
@@ -183,9 +188,6 @@ class EagleDraftInput(SpecInput):
def __post_init__(self):
super().__init__(SpecInputType.EAGLE_DRAFT)
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req
@classmethod
def create_idle_input(
cls,
@@ -345,9 +347,6 @@ class EagleDraftExtendInput(SpecInput):
def __post_init__(self):
super().__init__(SpecInputType.EAGLE_DRAFT_EXTEND)
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req
@classmethod
def create_idle_input(
cls,
@@ -1510,6 +1510,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
verify_input: EagleVerifyInput = batch.spec_info
record_stream_for_v2_verify(batch, verify_input, fwd_stream)
# Actual-width stamp; equals the real width (draft_token_num) only on
# the topk-1 chain. Tree consumers must read draft_token_num -- see
# EagleVerifyInput.get_spec_adjust_token_coefficient.
verify_input.num_tokens_per_req = self.speculative_num_steps + 1
bs = len(batch.seq_lens)
@@ -34,6 +34,7 @@ from sglang.srt.model_executor.runner_backend_utils import (
)
from sglang.srt.runtime_context import get_flags
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.utils import (
require_attn_tp_gather,
require_gathered_buffer,
@@ -113,7 +114,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.LAST
self.num_tokens_per_req = self.topk
# Static capture width.
self.num_tokens_per_req = resolve_num_tokens_per_req(
phase="draft_decode", server_args=model_runner.server_args
)
self.capture_bs, _ = get_batch_sizes_to_capture(
model_runner, self.num_tokens_per_req
)
@@ -269,6 +273,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
bonus_tokens=bonus_tokens,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
# Actual width of the next draft-decode forward: topk tokens per req.
spec_info.num_tokens_per_req = self.topk
spec_info.num_tokens_for_logprob_per_req = self.topk
spec_info.positions = positions
@@ -420,6 +420,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
# gates SWA eviction timing and the SWA prefix-lock release.
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
# Actual width of the next draft-decode forward: topk tokens per req.
spec_info.num_tokens_per_req = self.topk
spec_info.num_tokens_for_logprob_per_req = self.topk
spec_info.positions = self._position_for_batch(batch)
@@ -62,7 +62,10 @@ from sglang.srt.model_executor.runner_backend_utils import (
from sglang.srt.runtime_context import get_flags
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
from sglang.srt.speculative.spec_utils import fast_topk
from sglang.srt.speculative.spec_utils import (
fast_topk,
resolve_num_tokens_per_req,
)
from sglang.srt.utils import (
get_available_gpu_memory,
require_attn_tp_gather,
@@ -159,7 +162,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fixed window: every step extends each request by the same number of
# tokens, which lets all steps share one buffer set.
self.num_tokens_per_req = self.speculative_num_draft_tokens
self.num_tokens_per_req = resolve_num_tokens_per_req(
phase="draft_extend", server_args=model_runner.server_args
)
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req
self.extend_seq_lens_cpu = [self.num_tokens_per_req] * self.max_bs
@@ -681,6 +686,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
num_correct_drafts=buffers.num_correct_drafts[:bs],
num_accept_tokens=buffers.num_accept_tokens[:bs],
)
# Actual width of the captured forward == static width by construction.
spec_info.num_tokens_per_req = self.num_tokens_per_req
spec_info.num_tokens_for_logprob_per_req = 1
spec_info.positions = buffers.positions[:padded_num_tokens]
@@ -515,6 +515,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# Batch 2: Draft extend
draft_extend_input = EagleDraftExtendInput(
hidden_states=batch_result.logits_output.hidden_states,
# Actual width: the multi-layer chain fills num_steps + 1 rows/req.
num_tokens_per_req=self.speculative_num_steps + 1,
num_tokens_for_logprob_per_req=1,
)
+3 -4
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import Optional, Tuple
from typing import Optional
import torch
@@ -33,6 +33,8 @@ class NgramVerifyInput(SpecInput):
self.retrieve_next_token = retrieve_next_token
self.retrieve_next_sibling = retrieve_next_sibling
self.draft_token_num = draft_token_num
self.num_tokens_per_req = draft_token_num
self.num_tokens_for_logprob_per_req = draft_token_num
self.grammar = grammar
# Inputs for V2 overlap worker
@@ -57,9 +59,6 @@ class NgramVerifyInput(SpecInput):
# Irregular tree: per-level branching follows the corpus matches.
return -1
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
def generate_attn_arg_prefill(
self,
req_pool_indices: torch.Tensor,
+8 -3
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import warnings
from abc import ABC, abstractmethod
from abc import ABC
from enum import Enum, IntEnum, auto
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union
@@ -314,6 +314,12 @@ class SpecInput(ABC):
# assignment, so an init-time default would clobber the passed layout.
ragged_verify_layout: Optional[RaggedVerifyLayout] = None
# Uniform per-request token width of this forward (and its logits-row
# counterpart). Doubles as the DP-attention global_num_tokens multiplier
# (ragged forwards carry 1 there). -1 = not set by this flow.
num_tokens_per_req: int = -1
num_tokens_for_logprob_per_req: int = -1
# 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.
@@ -344,9 +350,8 @@ class SpecInput(ABC):
SpecInputType.NGRAM_VERIFY,
}
@abstractmethod
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
pass
return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req
def get_spec_adjusted_global_num_tokens(
self, batch: ScheduleBatch
+27 -1
View File
@@ -5,7 +5,7 @@ import logging
import os
import time
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple
import torch
from huggingface_hub import snapshot_download
@@ -87,6 +87,32 @@ if _is_cpu:
logger = logging.getLogger(__name__)
def resolve_num_tokens_per_req(
*,
phase: Literal["draft_decode", "draft_extend", "target_verify"],
server_args: ServerArgs,
spec_algorithm=None,
is_draft_worker: bool = False,
num_draft_tokens: Optional[int] = None,
) -> int:
"""Single static derivation point for a spec phase's per-request token
width (sizes capture shapes / buffers); the per-forward dynamic width
lives on ``SpecInput.num_tokens_per_req``. Draft phases are
EAGLE-family-only; "target_verify" is algorithm-generic via the hook.
"""
if phase == "draft_decode":
return server_args.speculative_eagle_topk
if phase == "draft_extend":
return server_args.speculative_num_draft_tokens
if phase == "target_verify":
if num_draft_tokens is None:
num_draft_tokens = server_args.speculative_num_draft_tokens
return spec_algorithm.get_num_tokens_per_req_for_target_verify(
num_draft_tokens, is_draft_worker
)
raise ValueError(f"Unknown speculative phase: {phase}")
def fast_sample(probs: torch.Tensor, num_samples: int = 1):
sample_index = torch.multinomial(probs, num_samples=num_samples)
sample_p = probs.gather(1, sample_index)