[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.runtime_context import get_server_args
from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry 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_info import SpecInput, SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
from sglang.srt.utils import get_compiler_backend from sglang.srt.utils import get_compiler_backend
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -191,10 +192,15 @@ class FlashAttentionBackend(AttentionBackend):
self.speculative_num_draft_tokens is not None self.speculative_num_draft_tokens is not None
and model_runner.is_draft_worker and model_runner.is_draft_worker
): ):
self.speculative_num_draft_tokens = SpeculativeAlgorithm.from_string( # 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 model_runner.server_args.speculative_algorithm
).get_num_tokens_per_req_for_target_verify( ),
int(self.speculative_num_draft_tokens), is_draft_worker=True is_draft_worker=True,
num_draft_tokens=int(self.speculative_num_draft_tokens),
) )
self.speculative_step_id = speculative_step_id self.speculative_step_id = speculative_step_id
@@ -2289,7 +2295,7 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_spec_metadata = metadata_swa metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend_v2(): 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"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
@@ -2693,9 +2699,7 @@ class FlashAttentionBackend(AttentionBackend):
device=device, device=device,
) )
else: else:
default_extend = getattr( default_extend = spec_info.num_tokens_per_req
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
)
extend_seq_lens = torch.full( extend_seq_lens = torch.full(
(bs,), default_extend, dtype=torch.int32, device=device (bs,), default_extend, dtype=torch.int32, device=device
) )
@@ -2704,9 +2708,7 @@ class FlashAttentionBackend(AttentionBackend):
if extend_seq_lens_cpu: if extend_seq_lens_cpu:
metadata.max_seq_len_q = int(max(extend_seq_lens_cpu)) metadata.max_seq_len_q = int(max(extend_seq_lens_cpu))
else: else:
metadata.max_seq_len_q = getattr( metadata.max_seq_len_q = spec_info.num_tokens_per_req
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
)
metadata.cu_seqlens_q[1:].copy_( metadata.cu_seqlens_q[1:].copy_(
torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32) 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 # host-known qo/kv layout from the caller. Assert rather than silently
# fall back to plan()'s blocking D2H on the replay hot-path. # fall back to plan()'s blocking D2H on the replay hot-path.
paged_plan_kwargs = {} 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 = ( uses_fast_prefill = (
hasattr(wrapper_paged.begin_forward, "func") hasattr(wrapper_paged.begin_forward, "func")
and wrapper_paged.begin_forward.func is fast_prefill_plan and wrapper_paged.begin_forward.func is fast_prefill_plan
@@ -487,7 +487,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
) )
self.target_verify_metadata[bs] = metadata self.target_verify_metadata[bs] = metadata
elif forward_mode.is_draft_extend_v2(): 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"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
@@ -581,6 +581,7 @@ class CPUGraphRunner:
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL self.capture_hidden_mode = CaptureHiddenMode.NULL
# Static capture width: CPU graphs are decode-only.
self.num_tokens_per_req = 1 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 # 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, set_global_server_args_for_scheduler,
) )
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm 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.base import TopkCaptureOutput
from sglang.srt.state_capturer.indexer_topk import ( from sglang.srt.state_capturer.indexer_topk import (
create_indexer_capturer, create_indexer_capturer,
@@ -622,10 +623,12 @@ class ModelRunner:
) -> int: ) -> int:
"""Logits rows per decode batch slot.""" """Logits rows per decode batch slot."""
if self.spec_algorithm.is_speculative(): if self.spec_algorithm.is_speculative():
if num_draft_tokens is None: return resolve_num_tokens_per_req(
num_draft_tokens = self.server_args.speculative_num_draft_tokens phase="target_verify",
return self.spec_algorithm.get_num_tokens_per_req_for_target_verify( server_args=self.server_args,
num_draft_tokens, self.is_draft_worker 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) dllm_config = DllmConfig.from_server_args(self.server_args)
return dllm_config.block_size if dllm_config is not None else 1 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" role = "draft" if model_runner.is_draft_worker else "target"
if model_runner.spec_algorithm.is_speculative(): if model_runner.spec_algorithm.is_speculative():
capture_name = f"{role} verify" capture_name = f"{role} verify"
num_tokens_per_req = ( num_tokens_per_req = model_runner.decode_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,
)
)
else: else:
capture_name = f"{role} decode" capture_name = f"{role} decode"
num_tokens_per_req = 1 num_tokens_per_req = 1
@@ -360,11 +360,7 @@ class BaseRunner(ABC):
if not mr.spec_algorithm.supports_target_verify_for_draft(): if not mr.spec_algorithm.supports_target_verify_for_draft():
raise RuntimeError("This should not happen") raise RuntimeError("This should not happen")
capture_forward_mode = ForwardMode.TARGET_VERIFY capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_req = ( num_tokens_per_req = mr.decode_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
)
)
if mr.server_args.enable_return_hidden_states: if mr.server_args.enable_return_hidden_states:
capture_hidden_mode = CaptureHiddenMode.FULL capture_hidden_mode = CaptureHiddenMode.FULL
@@ -254,6 +254,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------ # --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL self.capture_hidden_mode = CaptureHiddenMode.NULL
# Static capture width.
self.num_tokens_per_req = model_runner.decode_num_tokens_per_req( self.num_tokens_per_req = model_runner.decode_num_tokens_per_req(
num_draft_tokens=self.speculative_num_draft_tokens num_draft_tokens=self.speculative_num_draft_tokens
) )
@@ -83,10 +83,8 @@ class EagerRunner(BaseRunner):
), ),
) )
else: else:
num_tokens_per_req = ( num_tokens_per_req = mr.decode_num_tokens_per_req(
mr.spec_algorithm.get_num_tokens_per_req_for_target_verify( num_draft_tokens=num_draft_tokens
num_draft_tokens, mr.is_draft_worker
)
) )
else: else:
dllm_config = DllmConfig.from_server_args(sa) dllm_config = DllmConfig.from_server_args(sa)
@@ -150,11 +148,7 @@ class EagerRunner(BaseRunner):
mr = self.model_runner mr = self.model_runner
num_tokens_per_req = 1 num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative(): if mr.spec_algorithm.is_speculative():
num_tokens_per_req = ( num_tokens_per_req = mr.decode_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
)
)
return ( return (
self._alloc_dummy_decode_buffers( self._alloc_dummy_decode_buffers(
self._eager_max_bs, num_tokens_per_req=num_tokens_per_req self._eager_max_bs, num_tokens_per_req=num_tokens_per_req
@@ -283,6 +283,7 @@ class EagleDraftWorkerBase(ABC):
) )
# Get a forward batch # 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_per_req = topk
draft_input.num_tokens_for_logprob_per_req = topk draft_input.num_tokens_for_logprob_per_req = topk
capture_mode = ( capture_mode = (
+2 -4
View File
@@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional, Tuple from typing import TYPE_CHECKING, Optional
import torch import torch
@@ -48,9 +48,7 @@ class DFlashVerifyInput(SpecInput):
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY) super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
if self.num_tokens_per_req == -1: if self.num_tokens_per_req == -1:
self.num_tokens_per_req = int(self.draft_token_num) self.num_tokens_per_req = int(self.draft_token_num)
self.num_tokens_for_logprob_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
def prepare_for_verify( def prepare_for_verify(
self, self,
@@ -2,7 +2,7 @@
import contextlib import contextlib
from dataclasses import dataclass from dataclasses import dataclass
from typing import Optional, Tuple from typing import Optional
import torch import torch
@@ -63,10 +63,9 @@ class DFlashDraftInputV2(SpecInput):
def __post_init__(self): def __post_init__(self):
super().__init__(spec_input_type=SpecInputType.DFLASH_DRAFT) 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. # 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( def _ensure_prepare_length_buffers(
self, bs: int, device: torch.device | str 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.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.speculative.eagle_info import EagleDraftInput 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.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 ( from sglang.srt.utils import (
require_attn_tp_gather, require_attn_tp_gather,
require_gathered_buffer, require_gathered_buffer,
@@ -145,7 +146,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Bucket sizes # Bucket sizes
self.capture_bs, _ = get_batch_sizes_to_capture(model_runner) 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_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req 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.runtime_context import get_flags
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput 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.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 ( from sglang.srt.utils import (
is_hip, is_hip,
require_attn_tp_gather, require_attn_tp_gather,
@@ -131,9 +134,11 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_bs, _ = get_batch_sizes_to_capture(model_runner) self.capture_bs, _ = get_batch_sizes_to_capture(model_runner)
# Size cuda-graph buffers by num_draft_tokens (full tree width), not # Static capture width: full tree width (num_draft_tokens), not
# num_steps + 1, or topk > 1 draft-extend overflows them. # num_steps + 1 -- topk > 1 draft-extend overflows the buffers.
self.num_tokens_per_req = model_runner.server_args.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_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req 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) super().__init__(SpecInputType.EAGLE_VERIFY)
if self.num_tokens_per_req < 0: if self.num_tokens_per_req < 0:
self.num_tokens_per_req = self.draft_token_num 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 @property
def max_tree_depth(self) -> int: def max_tree_depth(self) -> int:
@@ -55,9 +63,6 @@ class EagleVerifyInput(SpecInput):
irregular tree (no fixed per-level branching).""" irregular tree (no fixed per-level branching)."""
return self.topk return self.topk
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
@classmethod @classmethod
def create_idle_input( def create_idle_input(
cls, topk: int, spec_steps: int, num_verify_tokens: int, device: str cls, topk: int, spec_steps: int, num_verify_tokens: int, device: str
@@ -183,9 +188,6 @@ class EagleDraftInput(SpecInput):
def __post_init__(self): def __post_init__(self):
super().__init__(SpecInputType.EAGLE_DRAFT) 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 @classmethod
def create_idle_input( def create_idle_input(
cls, cls,
@@ -345,9 +347,6 @@ class EagleDraftExtendInput(SpecInput):
def __post_init__(self): def __post_init__(self):
super().__init__(SpecInputType.EAGLE_DRAFT_EXTEND) 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 @classmethod
def create_idle_input( def create_idle_input(
cls, cls,
@@ -1510,6 +1510,9 @@ class EAGLEWorkerV2(BaseSpecWorker):
verify_input: EagleVerifyInput = batch.spec_info verify_input: EagleVerifyInput = batch.spec_info
record_stream_for_v2_verify(batch, verify_input, fwd_stream) 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 verify_input.num_tokens_per_req = self.speculative_num_steps + 1
bs = len(batch.seq_lens) 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.runtime_context import get_flags
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput 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 ( from sglang.srt.utils import (
require_attn_tp_gather, require_attn_tp_gather,
require_gathered_buffer, require_gathered_buffer,
@@ -113,7 +114,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.LAST 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( self.capture_bs, _ = get_batch_sizes_to_capture(
model_runner, self.num_tokens_per_req model_runner, self.num_tokens_per_req
) )
@@ -269,6 +273,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
bonus_tokens=bonus_tokens, bonus_tokens=bonus_tokens,
capture_hidden_mode=CaptureHiddenMode.LAST, 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_per_req = self.topk
spec_info.num_tokens_for_logprob_per_req = self.topk spec_info.num_tokens_for_logprob_per_req = self.topk
spec_info.positions = positions spec_info.positions = positions
@@ -420,6 +420,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
# gates SWA eviction timing and the SWA prefix-lock release. # gates SWA eviction timing and the SWA prefix-lock release.
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST 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_per_req = self.topk
spec_info.num_tokens_for_logprob_per_req = self.topk spec_info.num_tokens_for_logprob_per_req = self.topk
spec_info.positions = self._position_for_batch(batch) 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.runtime_context import get_flags
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput 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.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 ( from sglang.srt.utils import (
get_available_gpu_memory, get_available_gpu_memory,
require_attn_tp_gather, require_attn_tp_gather,
@@ -159,7 +162,9 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fixed window: every step extends each request by the same number of # Fixed window: every step extends each request by the same number of
# tokens, which lets all steps share one buffer set. # 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_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_req 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 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_correct_drafts=buffers.num_correct_drafts[:bs],
num_accept_tokens=buffers.num_accept_tokens[: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_per_req = self.num_tokens_per_req
spec_info.num_tokens_for_logprob_per_req = 1 spec_info.num_tokens_for_logprob_per_req = 1
spec_info.positions = buffers.positions[:padded_num_tokens] spec_info.positions = buffers.positions[:padded_num_tokens]
@@ -515,6 +515,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# Batch 2: Draft extend # Batch 2: Draft extend
draft_extend_input = EagleDraftExtendInput( draft_extend_input = EagleDraftExtendInput(
hidden_states=batch_result.logits_output.hidden_states, 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_per_req=self.speculative_num_steps + 1,
num_tokens_for_logprob_per_req=1, num_tokens_for_logprob_per_req=1,
) )
+3 -4
View File
@@ -1,6 +1,6 @@
from __future__ import annotations from __future__ import annotations
from typing import Optional, Tuple from typing import Optional
import torch import torch
@@ -33,6 +33,8 @@ class NgramVerifyInput(SpecInput):
self.retrieve_next_token = retrieve_next_token self.retrieve_next_token = retrieve_next_token
self.retrieve_next_sibling = retrieve_next_sibling self.retrieve_next_sibling = retrieve_next_sibling
self.draft_token_num = draft_token_num 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 self.grammar = grammar
# Inputs for V2 overlap worker # Inputs for V2 overlap worker
@@ -57,9 +59,6 @@ class NgramVerifyInput(SpecInput):
# Irregular tree: per-level branching follows the corpus matches. # Irregular tree: per-level branching follows the corpus matches.
return -1 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( def generate_attn_arg_prefill(
self, self,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
+8 -3
View File
@@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
import warnings import warnings
from abc import ABC, abstractmethod from abc import ABC
from enum import Enum, IntEnum, auto from enum import Enum, IntEnum, auto
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union 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. # assignment, so an init-time default would clobber the passed layout.
ragged_verify_layout: Optional[RaggedVerifyLayout] = None 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 # DSA MTP IndexShare seed relay. Class-level defaults (same rationale as
# ragged_verify_layout) so scheduler/relay/attention code reads them # ragged_verify_layout) so scheduler/relay/attention code reads them
# uniformly on any SpecInput; only the EAGLE-family inputs override them. # uniformly on any SpecInput; only the EAGLE-family inputs override them.
@@ -344,9 +350,8 @@ class SpecInput(ABC):
SpecInputType.NGRAM_VERIFY, SpecInputType.NGRAM_VERIFY,
} }
@abstractmethod
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]: 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( def get_spec_adjusted_global_num_tokens(
self, batch: ScheduleBatch self, batch: ScheduleBatch
+27 -1
View File
@@ -5,7 +5,7 @@ import logging
import os import os
import time import time
from contextlib import contextmanager 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 import torch
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
@@ -87,6 +87,32 @@ if _is_cpu:
logger = logging.getLogger(__name__) 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): def fast_sample(probs: torch.Tensor, num_samples: int = 1):
sample_index = torch.multinomial(probs, num_samples=num_samples) sample_index = torch.multinomial(probs, num_samples=num_samples)
sample_p = probs.gather(1, sample_index) sample_p = probs.gather(1, sample_index)