[SPEC V2][2/N] feat: adaptive spec support spec v2 (#23336)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
shuwenn
2026-05-07 18:33:47 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent 35870d55ac
commit d9dddd4d7d
12 changed files with 303 additions and 73 deletions
@@ -386,8 +386,16 @@ class SchedulerOutputProcessorMixin:
result.num_accepted_drafts = sum(accept_lens) - len(batch.reqs)
result.num_accepted_drafts_per_req_cpu = [x - 1 for x in accept_lens]
# Feed the adaptive controller now that accept_lens is on CPU,
# instead of doing a synchronous GPU→CPU copy in the worker hot path.
# BaseSpecWorker provides a no-op default for non-adaptive workers.
self.model_worker.on_verify_complete_cpu(result.num_accepted_drafts_per_req_cpu)
predict_tokens = []
stride = self.draft_worker.speculative_num_draft_tokens
# In adaptive spec-v2, the worker state may already have switched when this
# delayed result is processed. Use the draft token count recorded on result.
stride = result.speculative_num_draft_tokens
assert stride is not None, "spec-v2 result missing speculative_num_draft_tokens"
for i, req in enumerate(batch.reqs):
# -1 because prepare_for_decode pre-claimed the bonus slot.
+1
View File
@@ -39,6 +39,7 @@ class GenerationBatchResult:
copy_done: Optional[torch.cuda.Event] = None
delay_sample_func: Optional[callable] = None
future_indices: Optional[FutureIndices] = None
speculative_num_draft_tokens: Optional[int] = None
# FIXME(lsyin): maybe move to a better place?
# sync path: forward stream -> output processor
@@ -9,6 +9,8 @@ import json
import logging
from typing import TYPE_CHECKING
from sglang.srt.utils import log_info_on_rank0
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -32,11 +34,6 @@ def adaptive_unsupported_reason(server_args: ServerArgs) -> str | None:
"enable_dp_attention=True is not supported "
"(adaptive tier decisions are not synchronized across DP ranks)"
)
if not server_args.disable_overlap_schedule:
return (
"the overlap scheduler (spec v2) is enabled "
"(adaptive is only implemented for EAGLEWorker v1)"
)
if server_args.enable_multi_layer_eagle:
return (
"enable_multi_layer_eagle=True is not supported "
@@ -95,7 +92,22 @@ class AdaptiveSpeculativeParams:
):
cfg = config or {}
# TODO: Wider range of candidate_steps (once lazy init is supported).
self.candidate_steps = sorted(set(cfg.get("candidate_steps", [1, 3, 7])))
candidates = set(cfg.get("candidate_steps", [1, 3, 7]))
# Ensure the worker's initial speculative_num_steps is itself a candidate.
# Otherwise AdaptiveController.register() would store the worker's pre-built
# runtime state under a key that _activate() never queries, leaking that
# state's draft attn backend and cuda graph buffers for the process lifetime.
if initial_steps not in candidates:
log_info_on_rank0(
logger,
f"Adding initial speculative_num_steps={initial_steps} to "
f"candidate_steps={sorted(candidates)} so the pre-built "
f"runtime state is reused.",
)
candidates.add(initial_steps)
self.candidate_steps = sorted(candidates)
assert (
len(self.candidate_steps) >= 2
), "candidate_steps must have at least 2 distinct values"
@@ -108,18 +120,16 @@ class AdaptiveSpeculativeParams:
self.down_hysteresis = cfg.get("down_hysteresis", -0.25)
self.up_hysteresis = cfg.get("up_hysteresis", 0.0)
self.current_steps = min(
self.candidate_steps,
key=lambda step: (abs(step - initial_steps), -step),
)
self.current_steps = initial_steps
# Initialize EMA at current steps - 1 (neutral starting point)
self.ema_accept_len = float(self.current_steps - 1)
self._batch_count = 0
logger.info(
log_info_on_rank0(
logger,
f"AdaptiveSpeculativeParams initialized: "
f"steps={self.current_steps}, candidate_steps={self.candidate_steps}"
f"steps={self.current_steps}, candidate_steps={self.candidate_steps}",
)
def update(self, num_accepted_drafts_per_req: list[int]) -> bool:
@@ -171,9 +181,10 @@ class AdaptiveSpeculativeParams:
if target != old_steps:
self.current_steps = target
logger.info(
log_info_on_rank0(
logger,
f"Adaptive spec params updated: steps {old_steps} -> {target} "
f"(ema_accept_len={self.ema_accept_len:.2f})"
f"(ema_accept_len={self.ema_accept_len:.2f})",
)
return True
return False
@@ -32,3 +32,11 @@ class BaseSpecWorker(ABC):
def clear_cache_pool(self):
# TODO: move this abstract method to BaseTpWorker and call through self.model_runner
pass
def on_verify_complete_cpu(self, num_accepted_drafts_per_req: list[int]) -> None:
"""Hook called after verify finishes and accept counts are on CPU.
Default no-op. Adaptive-aware workers override this to feed the
controller without forcing a GPU→CPU sync in the worker hot path.
"""
pass
@@ -124,7 +124,12 @@ class EagleDraftInputV2Mixin:
num_needed_tokens = 0
for i, r in enumerate(batch.reqs):
cur = r.kv_allocated_len
nxt = r.kv_committed_len + double_alloc
# max(cur, ...) clamps so adaptive downswitch (smaller alloc_len_per_decode)
# cannot make nxt < cur and corrupt allocator state. kv_committed_len lags
# batch.seq_lens by ~1 verify in overlap mode, so we react to adaptive
# switches one batch later than a seq_lens-based baseline; the 2*alloc
# over-allocation buffer absorbs that lag.
nxt = max(cur, r.kv_committed_len + double_alloc)
cur_kv_lens[i] = cur
nxt_kv_lens[i] = nxt
num_needed_tokens += nxt - cur
+24 -18
View File
@@ -73,6 +73,7 @@ from sglang.srt.utils import (
is_cuda,
is_musa,
is_npu,
log_info_on_rank0,
next_power_of_2,
)
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
@@ -278,30 +279,34 @@ class EAGLEWorker(TpModelWorker):
if self.speculative_num_steps > 1:
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
log_info_on_rank0(
logger,
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB",
)
self.cuda_graph_runner = Device2DraftCudaGraphRunner[
self.target_worker.device
](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
log_info_on_rank0(
logger,
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB.",
)
# Capture extend
if self.draft_extend_attn_backend and not _is_npu:
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
log_info_on_rank0(
logger,
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB",
)
self.cuda_graph_runner_for_draft_extend = EAGLEDraftExtendCudaGraphRunner(
self
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
log_info_on_rank0(
logger,
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB.",
)
def apply_runtime_state(self, state: SpecRuntimeState):
@@ -309,11 +314,12 @@ class EAGLEWorker(TpModelWorker):
if self.speculative_num_steps == state.speculative_num_steps:
return
logger.info(
log_info_on_rank0(
logger,
"Switch adaptive runtime state: "
f"steps {self.speculative_num_steps} -> {state.speculative_num_steps}, "
f"draft_tokens {self.speculative_num_draft_tokens} -> "
f"{state.speculative_num_draft_tokens}"
f"{state.speculative_num_draft_tokens}",
)
self.speculative_num_steps = state.speculative_num_steps
@@ -384,10 +390,11 @@ class EAGLEWorker(TpModelWorker):
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
log_info_on_rank0(
logger,
f"Built adaptive runtime state steps={speculative_num_steps}: "
f"elapsed={time.perf_counter() - tic:.2f}s, "
f"mem={(before_mem - after_mem):.2f}GB"
f"mem={(before_mem - after_mem):.2f}GB",
)
return state
@@ -403,9 +410,9 @@ class EAGLEWorker(TpModelWorker):
self.speculative_num_draft_tokens,
self.draft_attn_backend,
self.draft_extend_attn_backend,
getattr(self.draft_model_runner, "draft_attn_backend", None),
getattr(self, "cuda_graph_runner", None),
getattr(self, "cuda_graph_runner_for_draft_extend", None),
self.draft_model_runner.draft_attn_backend,
self.cuda_graph_runner,
self.cuda_graph_runner_for_draft_extend,
sa.speculative_num_steps,
sa.speculative_num_draft_tokens,
)
@@ -507,9 +514,8 @@ class EAGLEWorker(TpModelWorker):
batch.reqs, "set_spec_draft_extend_end_time", trace_only=True
)
controller = getattr(self, "adaptive_controller", None)
if controller is not None:
controller.on_verify_complete(
if self.adaptive_controller is not None:
self.adaptive_controller.on_verify_complete(
verify_output.num_accepted_drafts_per_req_cpu
)
@@ -30,8 +30,13 @@ from sglang.srt.managers.io_struct import (
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
SpecRuntimeState,
)
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
@@ -65,6 +70,7 @@ from sglang.srt.utils.common import (
is_hip,
is_musa,
is_npu,
log_info_on_rank0,
next_power_of_2,
)
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
@@ -273,15 +279,17 @@ class EagleDraftWorker(BaseDraftWorker):
if self.speculative_num_steps > 1:
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
log_info_on_rank0(
logger,
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB",
)
self.cuda_graph_runner = Device2DraftCudaGraphRunner[
self.target_worker.device
](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
log_info_on_rank0(
logger,
f"Capture draft cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB.",
)
Device2ExtendCudaGraphRunner = {
@@ -313,15 +321,17 @@ class EagleDraftWorker(BaseDraftWorker):
):
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
log_info_on_rank0(
logger,
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB",
)
self.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[
self.target_worker.device
](self)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB."
log_info_on_rank0(
logger,
f"Capture draft extend cuda graph end. Time elapsed: {time.perf_counter() - tic:.2f} s. mem usage={(before_mem - after_mem):.2f} GB. avail mem={after_mem:.2f} GB.",
)
def draft(self, model_worker_batch: ModelWorkerBatch):
@@ -678,6 +688,13 @@ class EAGLEWorkerV2(BaseSpecWorker):
target_worker,
)
# Adaptive speculative
self.adaptive_controller: Optional[AdaptiveController] = None
if server_args.speculative_adaptive:
self.adaptive_controller = AdaptiveController(
self, config_path=server_args.speculative_adaptive_config
)
# Some dummy tensors
self.num_new_pages_per_topk = torch.empty(
(), dtype=torch.int64, device=self.device
@@ -686,6 +703,25 @@ class EAGLEWorkerV2(BaseSpecWorker):
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
# Build adaptive runtime states (must be after draft worker is fully initialized)
if self.adaptive_controller is not None:
with self._draft_worker.draft_tp_context(
self._draft_worker.draft_runner.tp_group
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
self.adaptive_controller.register(
SpecRuntimeState(
speculative_num_steps=self.speculative_num_steps,
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
draft_attn_backend=self._draft_worker.draft_attn_backend,
cuda_graph_runner=self._draft_worker.cuda_graph_runner,
target_attn_backend=self._target_worker.model_runner.attn_backend,
target_graph_runner=self._target_worker.model_runner.graph_runner,
draft_extend_attn_backend=self._draft_worker.draft_extend_attn_backend,
cuda_graph_runner_for_draft_extend=self._draft_worker.cuda_graph_runner_for_draft_extend,
)
)
self.adaptive_controller.init_states()
@property
def target_worker(self):
return self._target_worker
@@ -754,8 +790,150 @@ class EAGLEWorkerV2(BaseSpecWorker):
self.draft_worker._draft_extend_for_decode(
model_worker_batch, batch_output
)
return batch_output
def on_verify_complete_cpu(self, accepted_draft_tokens: list[int]) -> None:
if self.adaptive_controller is not None:
self.adaptive_controller.on_verify_complete(accepted_draft_tokens)
# -- Adaptive speculative decoding protocol --
def build_adaptive_runtime_state(
self, speculative_num_steps: int, speculative_num_draft_tokens: int
) -> SpecRuntimeState:
"""Build a SpecRuntimeState for the given step configuration."""
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
with self._override_worker_state(
speculative_num_steps, speculative_num_draft_tokens
):
self._draft_worker.init_attention_backend()
self._draft_worker.init_cuda_graphs()
# Build target attention backend and CUDA graph runner
target_model_runner = self._target_worker.model_runner
backup_init = target_model_runner.init_new_workspace
try:
target_attn_backend = target_model_runner._get_attention_backend(
init_new_workspace=True
)
finally:
target_model_runner.init_new_workspace = backup_init
target_graph_runner = None
if not self.server_args.disable_cuda_graph:
target_graph_runner = CudaGraphRunner(
target_model_runner,
attn_backend=target_attn_backend,
speculative_num_steps=speculative_num_steps,
speculative_num_draft_tokens=speculative_num_draft_tokens,
)
state = SpecRuntimeState(
speculative_num_steps=speculative_num_steps,
speculative_num_draft_tokens=speculative_num_draft_tokens,
draft_attn_backend=self._draft_worker.draft_attn_backend,
cuda_graph_runner=self._draft_worker.cuda_graph_runner,
target_attn_backend=target_attn_backend,
target_graph_runner=target_graph_runner,
draft_extend_attn_backend=self._draft_worker.draft_extend_attn_backend,
cuda_graph_runner_for_draft_extend=self._draft_worker.cuda_graph_runner_for_draft_extend,
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
log_info_on_rank0(
logger,
f"Built adaptive runtime state steps={speculative_num_steps}: "
f"elapsed={time.perf_counter() - tic:.2f}s, "
f"mem={(before_mem - after_mem):.2f}GB",
)
return state
def apply_runtime_state(self, state: SpecRuntimeState) -> None:
"""Apply a pre-built runtime state to this worker."""
if self.speculative_num_steps == state.speculative_num_steps:
return
log_info_on_rank0(
logger,
"Switch adaptive runtime state: "
f"steps {self.speculative_num_steps} -> {state.speculative_num_steps}, "
f"draft_tokens {self.speculative_num_draft_tokens} -> "
f"{state.speculative_num_draft_tokens}",
)
# Top-level
self.speculative_num_steps = state.speculative_num_steps
self.speculative_num_draft_tokens = state.speculative_num_draft_tokens
# Draft side
dw = self._draft_worker
dw.speculative_num_steps = state.speculative_num_steps
dw.speculative_num_draft_tokens = state.speculative_num_draft_tokens
dw.draft_attn_backend = state.draft_attn_backend
dw.draft_runner.draft_attn_backend = state.draft_attn_backend
dw.cuda_graph_runner = state.cuda_graph_runner
dw.draft_extend_attn_backend = state.draft_extend_attn_backend
dw.cuda_graph_runner_for_draft_extend = state.cuda_graph_runner_for_draft_extend
# Target side
self._target_worker.model_runner.attn_backend = state.target_attn_backend
self._target_worker.model_runner.graph_runner = state.target_graph_runner
# Sync server_args
self.server_args.speculative_num_steps = state.speculative_num_steps
self.server_args.speculative_num_draft_tokens = (
state.speculative_num_draft_tokens
)
@contextlib.contextmanager
def _override_worker_state(
self, speculative_num_steps: int, speculative_num_draft_tokens: int
):
"""Temporarily override server_args and worker attributes for graph capture."""
sa = self.server_args
dw = self._draft_worker
backup = (
self.speculative_num_steps,
self.speculative_num_draft_tokens,
dw.speculative_num_steps,
dw.speculative_num_draft_tokens,
dw.draft_attn_backend,
dw.draft_extend_attn_backend,
dw.draft_runner.draft_attn_backend,
dw.cuda_graph_runner,
dw.cuda_graph_runner_for_draft_extend,
sa.speculative_num_steps,
sa.speculative_num_draft_tokens,
)
self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens = speculative_num_draft_tokens
dw.speculative_num_steps = speculative_num_steps
dw.speculative_num_draft_tokens = speculative_num_draft_tokens
sa.speculative_num_steps = speculative_num_steps
sa.speculative_num_draft_tokens = speculative_num_draft_tokens
try:
yield
finally:
(
self.speculative_num_steps,
self.speculative_num_draft_tokens,
dw.speculative_num_steps,
dw.speculative_num_draft_tokens,
dw.draft_attn_backend,
dw.draft_extend_attn_backend,
dw.draft_runner.draft_attn_backend,
dw.cuda_graph_runner,
dw.cuda_graph_runner_for_draft_extend,
sa.speculative_num_steps,
sa.speculative_num_draft_tokens,
) = backup
def verify(self, batch: ModelWorkerBatch):
# Since batch.seq_lens is allocated in another stream, we need
# record_stream() to prevent pytorch gc and reuse the gpu memory
@@ -890,6 +1068,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
@@ -788,6 +788,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
@@ -9,6 +9,9 @@ from sglang.srt.layers.moe.utils import (
)
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
)
from sglang.srt.speculative.eagle_worker import EAGLEWorker
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import draft_tp_context, load_token_map
@@ -48,6 +51,9 @@ class StandaloneWorker(EAGLEWorker):
server_args.speculative_algorithm
)
# TODO: Adaptive speculative
self.adaptive_controller: Optional[AdaptiveController] = None
# Override the context length of the draft model to be the same as the target model.
server_args.context_length = target_worker.model_runner.model_config.context_len
@@ -8,6 +8,9 @@ from sglang.srt.environ import envs
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
)
from sglang.srt.speculative.eagle_utils import TreeMaskMode
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -183,3 +186,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
# TODO: Adaptive speculative
self.adaptive_controller: Optional[AdaptiveController] = None