[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 = sum(accept_lens) - len(batch.reqs)
|
||||||
result.num_accepted_drafts_per_req_cpu = [x - 1 for x in accept_lens]
|
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 = []
|
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):
|
for i, req in enumerate(batch.reqs):
|
||||||
# -1 because prepare_for_decode pre-claimed the bonus slot.
|
# -1 because prepare_for_decode pre-claimed the bonus slot.
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ class GenerationBatchResult:
|
|||||||
copy_done: Optional[torch.cuda.Event] = None
|
copy_done: Optional[torch.cuda.Event] = None
|
||||||
delay_sample_func: Optional[callable] = None
|
delay_sample_func: Optional[callable] = None
|
||||||
future_indices: Optional[FutureIndices] = None
|
future_indices: Optional[FutureIndices] = None
|
||||||
|
speculative_num_draft_tokens: Optional[int] = None
|
||||||
|
|
||||||
# FIXME(lsyin): maybe move to a better place?
|
# FIXME(lsyin): maybe move to a better place?
|
||||||
# sync path: forward stream -> output processor
|
# sync path: forward stream -> output processor
|
||||||
|
|||||||
@@ -9,6 +9,8 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from sglang.srt.utils import log_info_on_rank0
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
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 "
|
"enable_dp_attention=True is not supported "
|
||||||
"(adaptive tier decisions are not synchronized across DP ranks)"
|
"(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:
|
if server_args.enable_multi_layer_eagle:
|
||||||
return (
|
return (
|
||||||
"enable_multi_layer_eagle=True is not supported "
|
"enable_multi_layer_eagle=True is not supported "
|
||||||
@@ -95,7 +92,22 @@ class AdaptiveSpeculativeParams:
|
|||||||
):
|
):
|
||||||
cfg = config or {}
|
cfg = config or {}
|
||||||
# TODO: Wider range of candidate_steps (once lazy init is supported).
|
# 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 (
|
assert (
|
||||||
len(self.candidate_steps) >= 2
|
len(self.candidate_steps) >= 2
|
||||||
), "candidate_steps must have at least 2 distinct values"
|
), "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.down_hysteresis = cfg.get("down_hysteresis", -0.25)
|
||||||
self.up_hysteresis = cfg.get("up_hysteresis", 0.0)
|
self.up_hysteresis = cfg.get("up_hysteresis", 0.0)
|
||||||
|
|
||||||
self.current_steps = min(
|
self.current_steps = initial_steps
|
||||||
self.candidate_steps,
|
|
||||||
key=lambda step: (abs(step - initial_steps), -step),
|
|
||||||
)
|
|
||||||
|
|
||||||
# Initialize EMA at current steps - 1 (neutral starting point)
|
# Initialize EMA at current steps - 1 (neutral starting point)
|
||||||
self.ema_accept_len = float(self.current_steps - 1)
|
self.ema_accept_len = float(self.current_steps - 1)
|
||||||
self._batch_count = 0
|
self._batch_count = 0
|
||||||
|
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
f"AdaptiveSpeculativeParams initialized: "
|
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:
|
def update(self, num_accepted_drafts_per_req: list[int]) -> bool:
|
||||||
@@ -171,9 +181,10 @@ class AdaptiveSpeculativeParams:
|
|||||||
|
|
||||||
if target != old_steps:
|
if target != old_steps:
|
||||||
self.current_steps = target
|
self.current_steps = target
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
f"Adaptive spec params updated: steps {old_steps} -> {target} "
|
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 True
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -32,3 +32,11 @@ class BaseSpecWorker(ABC):
|
|||||||
def clear_cache_pool(self):
|
def clear_cache_pool(self):
|
||||||
# TODO: move this abstract method to BaseTpWorker and call through self.model_runner
|
# TODO: move this abstract method to BaseTpWorker and call through self.model_runner
|
||||||
pass
|
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
|
num_needed_tokens = 0
|
||||||
for i, r in enumerate(batch.reqs):
|
for i, r in enumerate(batch.reqs):
|
||||||
cur = r.kv_allocated_len
|
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
|
cur_kv_lens[i] = cur
|
||||||
nxt_kv_lens[i] = nxt
|
nxt_kv_lens[i] = nxt
|
||||||
num_needed_tokens += nxt - cur
|
num_needed_tokens += nxt - cur
|
||||||
|
|||||||
@@ -73,6 +73,7 @@ from sglang.srt.utils import (
|
|||||||
is_cuda,
|
is_cuda,
|
||||||
is_musa,
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
|
log_info_on_rank0,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
@@ -278,30 +279,34 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
if self.speculative_num_steps > 1:
|
if self.speculative_num_steps > 1:
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
|
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.cuda_graph_runner = Device2DraftCudaGraphRunner[
|
||||||
self.target_worker.device
|
self.target_worker.device
|
||||||
](self)
|
](self)
|
||||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
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."
|
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
|
# Capture extend
|
||||||
if self.draft_extend_attn_backend and not _is_npu:
|
if self.draft_extend_attn_backend and not _is_npu:
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
|
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.cuda_graph_runner_for_draft_extend = EAGLEDraftExtendCudaGraphRunner(
|
||||||
self
|
self
|
||||||
)
|
)
|
||||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
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."
|
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):
|
def apply_runtime_state(self, state: SpecRuntimeState):
|
||||||
@@ -309,11 +314,12 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
if self.speculative_num_steps == state.speculative_num_steps:
|
if self.speculative_num_steps == state.speculative_num_steps:
|
||||||
return
|
return
|
||||||
|
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
"Switch adaptive runtime state: "
|
"Switch adaptive runtime state: "
|
||||||
f"steps {self.speculative_num_steps} -> {state.speculative_num_steps}, "
|
f"steps {self.speculative_num_steps} -> {state.speculative_num_steps}, "
|
||||||
f"draft_tokens {self.speculative_num_draft_tokens} -> "
|
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
|
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)
|
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"Built adaptive runtime state steps={speculative_num_steps}: "
|
||||||
f"elapsed={time.perf_counter() - tic:.2f}s, "
|
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
|
return state
|
||||||
@@ -403,9 +410,9 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
self.speculative_num_draft_tokens,
|
self.speculative_num_draft_tokens,
|
||||||
self.draft_attn_backend,
|
self.draft_attn_backend,
|
||||||
self.draft_extend_attn_backend,
|
self.draft_extend_attn_backend,
|
||||||
getattr(self.draft_model_runner, "draft_attn_backend", None),
|
self.draft_model_runner.draft_attn_backend,
|
||||||
getattr(self, "cuda_graph_runner", None),
|
self.cuda_graph_runner,
|
||||||
getattr(self, "cuda_graph_runner_for_draft_extend", None),
|
self.cuda_graph_runner_for_draft_extend,
|
||||||
sa.speculative_num_steps,
|
sa.speculative_num_steps,
|
||||||
sa.speculative_num_draft_tokens,
|
sa.speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
@@ -507,9 +514,8 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
batch.reqs, "set_spec_draft_extend_end_time", trace_only=True
|
batch.reqs, "set_spec_draft_extend_end_time", trace_only=True
|
||||||
)
|
)
|
||||||
|
|
||||||
controller = getattr(self, "adaptive_controller", None)
|
if self.adaptive_controller is not None:
|
||||||
if controller is not None:
|
self.adaptive_controller.on_verify_complete(
|
||||||
controller.on_verify_complete(
|
|
||||||
verify_output.num_accepted_drafts_per_req_cpu
|
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.schedule_batch import ModelWorkerBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
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.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch
|
||||||
from sglang.srt.server_args import ServerArgs
|
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.base_spec_worker import BaseDraftWorker, BaseSpecWorker
|
||||||
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
||||||
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
|
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
|
||||||
@@ -65,6 +70,7 @@ from sglang.srt.utils.common import (
|
|||||||
is_hip,
|
is_hip,
|
||||||
is_musa,
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
|
log_info_on_rank0,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||||
@@ -273,15 +279,17 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
if self.speculative_num_steps > 1:
|
if self.speculative_num_steps > 1:
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
f"Capture draft cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
|
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.cuda_graph_runner = Device2DraftCudaGraphRunner[
|
||||||
self.target_worker.device
|
self.target_worker.device
|
||||||
](self)
|
](self)
|
||||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
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."
|
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 = {
|
Device2ExtendCudaGraphRunner = {
|
||||||
@@ -313,15 +321,17 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
):
|
):
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB"
|
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.cuda_graph_runner_for_draft_extend = Device2ExtendCudaGraphRunner[
|
||||||
self.target_worker.device
|
self.target_worker.device
|
||||||
](self)
|
](self)
|
||||||
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
logger.info(
|
log_info_on_rank0(
|
||||||
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."
|
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):
|
def draft(self, model_worker_batch: ModelWorkerBatch):
|
||||||
@@ -678,6 +688,13 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
target_worker,
|
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
|
# Some dummy tensors
|
||||||
self.num_new_pages_per_topk = torch.empty(
|
self.num_new_pages_per_topk = torch.empty(
|
||||||
(), dtype=torch.int64, device=self.device
|
(), 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)
|
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
|
@property
|
||||||
def target_worker(self):
|
def target_worker(self):
|
||||||
return self._target_worker
|
return self._target_worker
|
||||||
@@ -754,8 +790,150 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
self.draft_worker._draft_extend_for_decode(
|
self.draft_worker._draft_extend_for_decode(
|
||||||
model_worker_batch, batch_output
|
model_worker_batch, batch_output
|
||||||
)
|
)
|
||||||
|
|
||||||
return 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):
|
def verify(self, batch: ModelWorkerBatch):
|
||||||
# Since batch.seq_lens is allocated in another stream, we need
|
# Since batch.seq_lens is allocated in another stream, we need
|
||||||
# record_stream() to prevent pytorch gc and reuse the gpu memory
|
# record_stream() to prevent pytorch gc and reuse the gpu memory
|
||||||
@@ -890,6 +1068,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=predict,
|
next_token_ids=predict,
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
|
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
|
||||||
next_draft_input=next_draft_input,
|
next_draft_input=next_draft_input,
|
||||||
accept_lens=accept_lens,
|
accept_lens=accept_lens,
|
||||||
routed_experts_output=forward_batch_output.routed_experts_output,
|
routed_experts_output=forward_batch_output.routed_experts_output,
|
||||||
|
|||||||
@@ -788,6 +788,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
next_token_ids=predict,
|
next_token_ids=predict,
|
||||||
can_run_cuda_graph=can_run_cuda_graph,
|
can_run_cuda_graph=can_run_cuda_graph,
|
||||||
|
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
|
||||||
next_draft_input=next_draft_input,
|
next_draft_input=next_draft_input,
|
||||||
accept_lens=accept_lens,
|
accept_lens=accept_lens,
|
||||||
routed_experts_output=forward_batch_output.routed_experts_output,
|
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.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.server_args import ServerArgs
|
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.eagle_worker import EAGLEWorker
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import draft_tp_context, load_token_map
|
from sglang.srt.speculative.spec_utils import draft_tp_context, load_token_map
|
||||||
@@ -48,6 +51,9 @@ class StandaloneWorker(EAGLEWorker):
|
|||||||
server_args.speculative_algorithm
|
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.
|
# 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
|
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.layers.moe.utils import speculative_moe_backend_context
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.server_args import ServerArgs
|
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_utils import TreeMaskMode
|
||||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
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.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
|
||||||
|
|
||||||
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
|
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
|
||||||
|
|
||||||
|
# TODO: Adaptive speculative
|
||||||
|
self.adaptive_controller: Optional[AdaptiveController] = None
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from types import SimpleNamespace
|
|||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.run_eval import run_eval
|
from sglang.test.run_eval import run_eval
|
||||||
@@ -59,7 +58,6 @@ class TestAdaptiveSpeculativeServer(CustomTestCase):
|
|||||||
cls.adaptive_config_path = f.name
|
cls.adaptive_config_path = f.name
|
||||||
|
|
||||||
try:
|
try:
|
||||||
with envs.SGLANG_ENABLE_SPEC_V2.override(False):
|
|
||||||
cls.process = popen_launch_server(
|
cls.process = popen_launch_server(
|
||||||
cls.model,
|
cls.model,
|
||||||
cls.base_url,
|
cls.base_url,
|
||||||
|
|||||||
@@ -7,14 +7,15 @@ register_cpu_ci(est_time=6, suite="stage-a-test-cpu")
|
|||||||
|
|
||||||
|
|
||||||
class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
||||||
def test_initial_steps_snap_to_nearest_candidate_preferring_larger_step(self):
|
def test_initial_steps_added_to_candidates_when_missing(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = AdaptiveSpeculativeParams(
|
||||||
initial_steps=2,
|
initial_steps=2,
|
||||||
config={"candidate_steps": [1, 3, 7]},
|
config={"candidate_steps": [1, 3, 7]},
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(params.current_steps, 3)
|
self.assertEqual(params.candidate_steps, [1, 2, 3, 7])
|
||||||
self.assertEqual(params.ema_accept_len, 2.0)
|
self.assertEqual(params.current_steps, 2)
|
||||||
|
self.assertEqual(params.ema_accept_len, 1.0)
|
||||||
|
|
||||||
def test_update_respects_warmup_and_interval(self):
|
def test_update_respects_warmup_and_interval(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = AdaptiveSpeculativeParams(
|
||||||
|
|||||||
Reference in New Issue
Block a user