[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.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 = (
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user