[SPEC V2][2/N] feat: adaptive spec support spec v2 (#23336)
Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user