From 76dc427806d92e5d2360bd0f687266e434397008 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Tue, 14 Jul 2026 18:41:08 -0700 Subject: [PATCH] [Spec] Single-source `num_tokens_per_req` derivation and access (#31013) --- .../attention/flashattention_backend.py | 24 ++++++++-------- .../layers/attention/flashinfer_backend.py | 4 ++- .../layers/attention/trtllm_mha_backend.py | 2 +- .../srt/model_executor/cpu_graph_runner.py | 1 + .../sglang/srt/model_executor/model_runner.py | 11 +++++--- .../cuda_graph_setup.py | 7 +---- .../srt/model_executor/runner/base_runner.py | 6 +--- .../runner/decode_cuda_graph_runner.py | 1 + .../srt/model_executor/runner/eager_runner.py | 12 ++------ .../srt/speculative/base_spec_worker.py | 1 + python/sglang/srt/speculative/dflash_info.py | 6 ++-- .../sglang/srt/speculative/dflash_info_v2.py | 7 ++--- .../eagle_draft_cuda_graph_runner.py | 6 +++- .../eagle_draft_extend_cuda_graph_runner.py | 13 ++++++--- python/sglang/srt/speculative/eagle_info.py | 17 ++++++----- .../sglang/srt/speculative/eagle_worker_v2.py | 3 ++ .../frozen_kv_mtp_cuda_graph_runner.py | 7 ++++- .../speculative/frozen_kv_mtp_worker_v2.py | 1 + ...er_eagle_draft_extend_cuda_graph_runner.py | 10 +++++-- .../multi_layer_eagle_worker_v2.py | 1 + python/sglang/srt/speculative/ngram_info.py | 7 ++--- python/sglang/srt/speculative/spec_info.py | 11 ++++++-- python/sglang/srt/speculative/spec_utils.py | 28 ++++++++++++++++++- 23 files changed, 116 insertions(+), 70 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index ec88e1ac5..fe22e6296 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -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) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 97b288ac3..2cffc0ae4 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -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 diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index c3e3346b1..5375a5d8b 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -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 ] diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 2d3b41f91..74a1433c2 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index cb81353c2..eae69fd72 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index 331f6a225..b0657f985 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -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 diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index e6e8b8a4a..78560fb76 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 983b8648c..0f28aa29b 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -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 ) diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 9e08454c5..f6f4afd26 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index 453039de2..06c51fa36 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -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 = ( diff --git a/python/sglang/srt/speculative/dflash_info.py b/python/sglang/srt/speculative/dflash_info.py index 762e5d164..6fb382ecb 100644 --- a/python/sglang/srt/speculative/dflash_info.py +++ b/python/sglang/srt/speculative/dflash_info.py @@ -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, diff --git a/python/sglang/srt/speculative/dflash_info_v2.py b/python/sglang/srt/speculative/dflash_info_v2.py index 2a9185987..360c80ce1 100644 --- a/python/sglang/srt/speculative/dflash_info_v2.py +++ b/python/sglang/srt/speculative/dflash_info_v2.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 47912bc70..554a69b2d 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index b456503f1..c3bcf9185 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index f6ff20d05..dd4d69933 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -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, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 93260868e..d059cc2c4 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 8c08874aa..2ed69aa58 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index f7a54a5d2..2dcd76cb6 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -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) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 0e8d2d233..e7461021d 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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] diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 85c60483f..68699eca6 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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, ) diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index 5c96872e9..b866aaadf 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -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, diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 1e8f1988d..5fe6f1066 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -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 diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index f83514673..c7ecdf87c 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -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)