[Spec] Single-source num_tokens_per_req derivation and access (#31013)
This commit is contained in:
@@ -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 = (
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user