[SPEC][1/N] feat: add adaptive speculative_num_steps for EAGLE topk=1 (#21599)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
shuwenn
2026-04-20 14:25:04 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent dbcf7459b5
commit b65799cf83
13 changed files with 1296 additions and 33 deletions
+263
View File
@@ -0,0 +1,263 @@
"""Benchmark adaptive speculative decoding against static baselines.
Run the same workload against one adaptive server and one or more static
servers, then compare throughput, latency, and acceptance length.
Workloads:
- low: steady-state low-acceptance generation
- high: steady-state high-acceptance generation
- transition: alternating low/high acceptance shifts to stress runtime switching
"""
import argparse
import time
from concurrent.futures import ThreadPoolExecutor
import requests
HIGH_PROMPTS = [
"Output exactly 256 new lines. Every line must be 1. Do not add numbering, punctuation, or commentary.",
"Output exactly 256 new lines. Every line must be READY. Do not add numbering, punctuation, or commentary.",
]
LOW_PROMPTS = [
"Compose a poem in the style of Emily Dickinson about quantum entanglement. Make it emotionally resonant.",
"Write 100 two-sentence biographies of eccentric inventors with unique names, hometowns, and inventions.",
"Write a long travel diary from a botanist visiting a chain of floating islands. Every paragraph should introduce new flora, customs, weather, and political tensions.",
"Write 80 newspaper headlines and subheads from 80 different alternate-history worlds. Each headline must introduce a different place, conflict, and technology.",
]
WORKLOADS = {
"low": [
("low", LOW_PROMPTS),
],
"high": [
("high", HIGH_PROMPTS),
],
"transition": [
("low_1", LOW_PROMPTS),
("high_1", HIGH_PROMPTS),
("low_2", LOW_PROMPTS),
("high_2", HIGH_PROMPTS),
],
}
def build_phase_plan(workload: str, num_requests: int):
return [
(phase_name, prompts, num_requests)
for phase_name, prompts in WORKLOADS[workload]
]
def send_request(base_url: str, prompt: str, max_tokens: int = 256):
start = time.perf_counter()
try:
resp = requests.post(
f"{base_url}/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": max_tokens,
},
"return_logprob": False,
},
timeout=max(120, max_tokens),
)
resp.raise_for_status()
data = resp.json()
except Exception as e:
return {"error": str(e), "latency": time.perf_counter() - start}
latency = time.perf_counter() - start
meta = data.get("meta_info", {})
completion_tokens = meta.get("completion_tokens", 0)
spec_verify_ct = meta.get("spec_verify_ct", 0)
accept_len = (
completion_tokens / spec_verify_ct if spec_verify_ct > 0 else float("nan")
)
return {
"latency": latency,
"completion_tokens": completion_tokens,
"spec_verify_ct": spec_verify_ct,
"accept_length": accept_len,
}
def run_phase(
base_url: str,
prompts,
phase_name: str,
num_requests: int,
max_tokens: int,
concurrency: int,
):
expanded = (prompts * ((num_requests + len(prompts) - 1) // len(prompts)))[
:num_requests
]
print(
f"\n--- Phase: {phase_name} ({num_requests} requests, concurrency={concurrency}) ---"
)
start = time.perf_counter()
with ThreadPoolExecutor(max_workers=concurrency) as pool:
futures = [pool.submit(send_request, base_url, p, max_tokens) for p in expanded]
results = [f.result() for f in futures]
elapsed = time.perf_counter() - start
errors = [r for r in results if "error" in r]
ok = [r for r in results if "error" not in r]
if not ok:
print(f" All {len(errors)} requests failed!")
return {"phase": phase_name, "error": True}
total_tokens = sum(r["completion_tokens"] for r in ok)
total_verify = sum(r["spec_verify_ct"] for r in ok)
avg_latency = sum(r["latency"] for r in ok) / len(ok)
throughput = total_tokens / elapsed
avg_accept_len = total_tokens / total_verify if total_verify > 0 else float("nan")
stats = {
"phase": phase_name,
"num_requests": len(ok),
"num_errors": len(errors),
"total_tokens": total_tokens,
"elapsed_s": round(elapsed, 2),
"throughput_tok_s": round(throughput, 2),
"avg_latency_s": round(avg_latency, 3),
"avg_accept_length": round(avg_accept_len, 3),
}
print(
f" Throughput: {throughput:.1f} tok/s | "
f"Avg latency: {avg_latency:.3f}s | "
f"Avg accept_len: {avg_accept_len:.2f} | "
f"Errors: {len(errors)}"
)
return stats
def summarize_phases(phase_stats):
ok_stats = [s for s in phase_stats if not s.get("error")]
if not ok_stats:
return {"error": True}
total_tokens = sum(s["total_tokens"] for s in ok_stats)
total_elapsed = sum(s["elapsed_s"] for s in ok_stats)
total_requests = sum(s["num_requests"] for s in ok_stats)
weighted_latency = sum(s["avg_latency_s"] * s["num_requests"] for s in ok_stats)
weighted_accept = sum(s["avg_accept_length"] * s["num_requests"] for s in ok_stats)
return {
"num_requests": total_requests,
"total_tokens": total_tokens,
"elapsed_s": round(total_elapsed, 2),
"throughput_tok_s": round(total_tokens / total_elapsed, 2),
"avg_latency_s": round(weighted_latency / total_requests, 3),
"avg_accept_length": round(weighted_accept / total_requests, 3),
}
def main():
parser = argparse.ArgumentParser(
description="Benchmark one workload for adaptive-vs-static speculative decoding"
)
parser.add_argument("--host", type=str, default="127.0.0.1")
parser.add_argument("--port", type=int, default=30000)
parser.add_argument(
"--workload",
choices=sorted(WORKLOADS),
default="transition",
help="Workload preset to run.",
)
parser.add_argument(
"--requests",
type=int,
default=8,
help="Requests per phase.",
)
parser.add_argument("--max-tokens", type=int, default=256)
parser.add_argument(
"--concurrency",
type=int,
default=2,
help="Concurrent requests.",
)
parser.add_argument(
"--warmup", type=int, default=2, help="Warmup requests before the benchmark."
)
args = parser.parse_args()
if args.requests < 1:
parser.error("--requests must be >= 1")
if args.concurrency < 1:
parser.error("--concurrency must be >= 1")
if args.warmup < 0:
parser.error("--warmup must be >= 0")
base_url = f"http://{args.host}:{args.port}"
print(f"Server: {base_url}")
print(f"Workload: {args.workload}")
phase_plan = build_phase_plan(args.workload, args.requests)
if args.warmup > 0:
print(f"\nWarming up with {args.warmup} requests...")
warmup_prompts = phase_plan[0][1]
run_phase(
base_url,
warmup_prompts,
"warmup",
args.warmup,
args.max_tokens,
args.concurrency,
)
phase_stats = []
for phase_name, prompts, num_requests in phase_plan:
phase_stats.append(
run_phase(
base_url,
prompts,
phase_name,
num_requests,
args.max_tokens,
args.concurrency,
)
)
overall = summarize_phases(phase_stats)
print("\n" + "=" * 70)
print("SUMMARY")
print("=" * 70)
print(f"{'Phase':<10} {'Throughput':>12} {'Avg Latency':>12} {'Accept Len':>12}")
print("-" * 50)
for stats in phase_stats:
if stats.get("error"):
print(f"{stats['phase']:<10} {'ERROR':>12}")
continue
print(
f"{stats['phase']:<10} "
f"{stats['throughput_tok_s']:>10.1f}/s "
f"{stats['avg_latency_s']:>10.3f}s "
f"{stats['avg_accept_length']:>11.2f}"
)
if not overall.get("error"):
print("-" * 50)
print(
f"{'OVERALL':<10} "
f"{overall['throughput_tok_s']:>10.1f}/s "
f"{overall['avg_latency_s']:>10.3f}s "
f"{overall['avg_accept_length']:>11.2f}"
)
if __name__ == "__main__":
main()
@@ -0,0 +1,156 @@
# Adaptive Speculative Decoding
Adaptive speculative decoding lets SGLang adjust `speculative_num_steps/speculative_num_draft_tokens` at runtime instead of keeping a single fixed value for the whole server lifetime.
It is designed for workloads whose accept length changes over time, where one static step count is rarely optimal.
## Current support
- Only `--speculative-algorithm EAGLE`
- Only `--speculative-eagle-topk 1`
- If either condition is not met, SGLang falls back to static speculative settings
## Why adaptive steps help
`speculative_num_steps` controls how many draft-model autoregressive steps run in each speculative round. In practice, the best value depends on the current workload.
- If `num_steps` is too small, the draft model could have produced more accepted tokens, but the round stops too early.
- If `num_steps` is too large, the draft model produces many candidate tokens that the target model rejects, so extra draft work is wasted.
Real traffic often moves between high-acceptance and low-acceptance phases, so one fixed step count is usually a compromise. Adaptive mode tries to follow the workload instead of hard-coding a single global `num_steps`.
## Design overview
The adaptive mechanism has three pieces:
- `AdaptiveSpeculativeParams`: the EMA-based policy
- `SpecRuntimeState`: the per-tier runtime state bundle
- `AdaptiveController`: the coordinator that chooses a tier and activates the matching runtime state
At startup, SGLang pre-builds one runtime state per candidate tier. By default, the candidate tiers are `candidate_steps = [1, 3, 7]`.
```text
┌──────────────────────────────────────────────────────────┐
│ SpecRuntimeState │
│ │
│ speculative_num_steps / speculative_num_draft_tokens │
│ │
│ ┌────────────────┐ ┌────────────────┐ ┌──────────────┐ │
│ │ Draft stage │ │ Verify stage │ │ Extend stage │ │
│ │ │ │ │ │ │ │
│ │ attn_backend │ │ attn_backend │ │ attn_backend │ │
│ │ cuda_graph │ │ cuda_graph │ │ cuda_graph │ │
│ └────────────────┘ └────────────────┘ └──────────────┘ │
└──────────────────────────────────────────────────────────┘
```
This matters because `CudaGraphRunner` is shape-dependent. Each candidate tier owns its own graph and backend state, so runtime switching is a reference swap, not an online graph recapture.
## Runtime flow
The adaptive update happens after verify and affects the next round, not the current one:
```text
┌─────────────────────────────────────────────────────────────────────┐
│ EAGLEWorker.forward_batch_generation() — decode path │
│ │
│ ① draft(batch) │
│ │ draft model multi-step generation with current tier │
│ v │
│ ② verify(batch, spec_info) │
│ │ target model tree verification │
│ │ → produces accept_length_per_req │
│ v │
│ ③ forward_draft_extend_after_decode(batch) │
│ │ draft model KV-cache catch-up │
│ v │
│ ④ adaptive_controller.on_verify_complete(accept_lengths) │
│ │ │
│ │ update EMA, apply warmup / interval / hysteresis gates │
│ │ if tier changed, select a pre-built state from pool │
│ v │
│ worker.apply_runtime_state(state) │
│ │
│ Tier switch happens after the current round completes. │
│ Backends and CUDA graphs are never swapped mid-round. │
└─────────────────────────────────────────────────────────────────────┘
```
## How the policy decides
After each verify pass, SGLang reads the accepted draft length per request, computes the batch average, smooths it with an exponential moving average (EMA), and switches among the pre-built candidate tiers `[1, 3, 7]` by default.
The decision logic is intentionally conservative:
- `warmup_batches` skips the first few batches
- `update_interval` avoids switching every batch
- `down_hysteresis` and `up_hysteresis` reduce oscillation
Conceptually, the policy probes one step beyond the observed acceptance:
```text
target_steps ≈ clamp(round(ema_accept_len) + 1, min(candidate_steps), max(candidate_steps))
```
So if recent requests consistently accept more drafted tokens, the policy tends to move up. If they start rejecting earlier, it tends to move down.
## Usage
`--speculative-adaptive-config` is optional, but the speculative setup still needs to be valid for adaptive mode.
```bash
python3 -m sglang.launch_server \
--model meta-llama/Llama-2-7b-chat-hf \
--speculative-algorithm EAGLE \
--speculative-draft-model-path lmsys/sglang-EAGLE-llama2-chat-7B \
--speculative-eagle-topk 1 \
--speculative-num-steps 3 \
--speculative-num-draft-tokens 4 \
--speculative-adaptive
```
If you want to override the defaults, add `--speculative-adaptive-config /path/to/adaptive_spec.json`.
Example config:
```json
{
"candidate_steps": [1, 3, 7],
"ema_alpha": 0.2,
"warmup_batches": 10,
"update_interval": 5
}
```
## Config file reference
The config file is optional. Any omitted keys use defaults.
| Key | Default | Meaning |
|---|---|---|
| `candidate_steps` | `[1, 3, 7]` | Discrete `speculative_num_steps` tiers that adaptive mode can switch between |
| `ema_alpha` | `0.2` | EMA smoothing factor for accepted draft length |
| `update_interval` | `5` | Recompute interval, in verify batches, after warmup |
| `warmup_batches` | `10` | Number of verify batches to observe before switching |
| `down_hysteresis` | `-0.25` | Extra margin before moving to a smaller step |
| `up_hysteresis` | `0.0` | Extra margin before moving to a larger step |
The initial `--speculative-num-steps` is snapped to the nearest value in `candidate_steps`.
## Monitoring
You can inspect the active tier and acceptance metric via `/server_info`:
```bash
curl -s http://127.0.0.1:30000/server_info | jq '.internal_states[0] | {speculative_num_steps, avg_spec_accept_length}'
```
- `speculative_num_steps` is the current active tier
- `avg_spec_accept_length` helps explain whether the server is likely to move up or down
## Tuning tips
- Start with the default candidate tiers `[1, 3, 7]`
- Use fewer tiers if you want lower startup and graph-memory overhead
- Increase `ema_alpha` to react faster, or lower it for more stability
- Increase `warmup_batches` or `update_interval` if tier switching is too noisy
- If your workload is already stable and one static setting is well tuned, adaptive mode may not help much
@@ -23,6 +23,7 @@ SGLang provides several speculative decoding options, including EAGLE-2/EAGLE-3,
- **Best speed/quality (recommended)**: Use **EAGLE-3** with `--speculative-algorithm EAGLE3`.
- **Strong default / broad compatibility**: Use **EAGLE-2** with `--speculative-algorithm EAGLE`.
- **Workload acceptance changes over time**: Use [**Adaptive speculative decoding**](adaptive_speculative_decoding.md) on top of **EAGLE** with `--speculative-eagle-topk 1`.
- **Lower `lm_head` overhead for EAGLE-2**: Enable **FR-Spec** with `--speculative-token-map`.
- **Model is MTP-enabled**: Use **MTP via speculative decoding** (often with small `speculative_num_steps/topk/num_draft_tokens`, see the example section).
- **You have a smaller draft LLM**: Use **STANDALONE** (`--speculative-algorithm STANDALONE`).
@@ -75,6 +76,7 @@ To enable EAGLE speculative decoding the following parameters are relevant:
These parameters are mostly the same for EAGLE-2 and EAGLE-3. `--speculative-token-map` is ignored for EAGLE-3 models.
For `--speculative-num-steps`, `--speculative-eagle-topk`, and `--speculative-num-draft-tokens`: leave all three unset to use auto-tuning, or set all three explicitly when tuning.
If you use EAGLE with `--speculative-eagle-topk 1` and your acceptance rate varies across requests, see [Adaptive Speculative Decoding](adaptive_speculative_decoding.md).
You can find the best combinations of these parameters with [bench_speculative.py](https://github.com/sgl-project/sglang/blob/main/scripts/playground/bench_speculative.py).
+1
View File
@@ -45,6 +45,7 @@ Its core features include:
advanced_features/hyperparameter_tuning.md
advanced_features/attention_backend.md
advanced_features/speculative_decoding.ipynb
advanced_features/adaptive_speculative_decoding.md
advanced_features/structured_outputs.ipynb
advanced_features/structured_outputs_for_reasoning_models.ipynb
advanced_features/tool_parser.ipynb
@@ -512,7 +512,14 @@ def set_global_graph_memory_pool(val):
class CudaGraphRunner:
"""A CudaGraphRunner runs the forward pass of a model with cuda graph and torch.compile."""
def __init__(self, model_runner: ModelRunner):
def __init__(
self,
model_runner: ModelRunner,
*,
attn_backend=None,
speculative_num_steps: Optional[int] = None,
speculative_num_draft_tokens: Optional[int] = None,
):
# Parse args
self.model_runner = model_runner
self.device = model_runner.device
@@ -551,6 +558,17 @@ class CudaGraphRunner:
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
self.is_dllm = self.dllm_config is not None
self.attn_backend = attn_backend or model_runner.attn_backend
self.speculative_num_steps = (
model_runner.server_args.speculative_num_steps
if speculative_num_steps is None
else speculative_num_steps
)
self.speculative_num_draft_tokens = (
model_runner.server_args.speculative_num_draft_tokens
if speculative_num_draft_tokens is None
else speculative_num_draft_tokens
)
self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL
@@ -561,9 +579,7 @@ class CudaGraphRunner:
if not self.model_runner.spec_algorithm.is_dflash():
raise RuntimeError("This should not happen")
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
self.num_tokens_per_bs = (
self.model_runner.server_args.speculative_num_draft_tokens
)
self.num_tokens_per_bs = self.speculative_num_draft_tokens
elif self.is_dllm:
self.capture_forward_mode = ForwardMode.DLLM_EXTEND
self.num_tokens_per_bs = self.dllm_config.block_size
@@ -583,14 +599,12 @@ class CudaGraphRunner:
# Attention backend
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
self.model_runner.attn_backend.init_cuda_graph_state(
self.max_bs, self.max_num_token
)
self.attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
# Init PDMux if needed
self.maybe_init_pdmux()
self.seq_len_fill_value = (
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
self.attn_backend.get_cuda_graph_seq_len_fill_value()
if self.dllm_config is None
else self.dllm_config.block_size
)
@@ -964,7 +978,7 @@ class CudaGraphRunner:
)
if stream_idx is None:
attn_backend = self.model_runner.attn_backend
attn_backend = self.attn_backend
else:
assert self.enable_pdmux
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
@@ -1170,7 +1184,7 @@ class CudaGraphRunner:
stream_idx = get_current_stream_idx()
attn_backend = self.model_runner.decode_attn_backend_group[stream_idx]
else:
attn_backend = self.model_runner.attn_backend
attn_backend = self.attn_backend
attn_backend.init_forward_metadata_replay_cuda_graph(
bs,
buffers.req_pool_indices[:bs],
@@ -1270,9 +1284,9 @@ class CudaGraphRunner:
retrive_next_token=None,
retrive_next_sibling=None,
retrive_cum_len=None,
spec_steps=self.model_runner.server_args.speculative_num_steps,
spec_steps=self.speculative_num_steps,
topk=self.model_runner.server_args.speculative_eagle_topk,
draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None,
seq_lens_cpu=None,
+30
View File
@@ -509,6 +509,8 @@ class ServerArgs:
speculative_moe_runner_backend: Optional[str] = None
speculative_moe_a2a_backend: Optional[str] = None
speculative_draft_model_quantization: Optional[str] = None
speculative_adaptive: bool = False
speculative_adaptive_config: Optional[str] = None
# Speculative decoding (ngram)
speculative_ngram_min_bfs_breadth: int = 1
@@ -3455,6 +3457,22 @@ class ServerArgs:
"Currently ngram speculative decoding does not support dp attention."
)
if self.speculative_adaptive:
if self.speculative_algorithm not in ("EAGLE", "EAGLE3"):
logger.warning(
"speculative_adaptive is only supported with EAGLE/EAGLE3 and topk=1. "
f"Current algorithm={self.speculative_algorithm}. "
"Falling back to static params."
)
self.speculative_adaptive = False
elif self.speculative_eagle_topk != 1:
logger.warning(
"speculative_adaptive is only supported with topk=1. "
f"Current topk={self.speculative_eagle_topk}. "
"Falling back to static params."
)
self.speculative_adaptive = False
def _handle_load_format(self):
if (
self.load_format == "auto" or self.load_format == "gguf"
@@ -5290,6 +5308,18 @@ class ServerArgs:
default=ServerArgs.speculative_ngram_external_corpus_max_tokens,
help="Fail startup if the tokenized external ngram corpus exceeds this many tokens. Tune this based on your CPU memory budget.",
)
parser.add_argument(
"--speculative-adaptive",
action="store_true",
help="Enable adaptive speculative decoding that dynamically adjusts num_steps based on acceptance rate.",
default=ServerArgs.speculative_adaptive,
)
parser.add_argument(
"--speculative-adaptive-config",
type=str,
help="Path to a JSON config file for adaptive speculative decoding tuning knobs ",
default=ServerArgs.speculative_adaptive_config,
)
# Multi-layer Eagle speculative decoding
parser.add_argument(
@@ -0,0 +1,121 @@
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Protocol
from sglang.srt.speculative.adaptive_spec_params import (
AdaptiveSpeculativeParams,
load_adaptive_config,
)
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
)
from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
EAGLEDraftExtendCudaGraphRunner,
)
logger = logging.getLogger(__name__)
@dataclass
class SpecRuntimeState:
"""A complete set of runtime resources bound to a specific speculative
decoding configuration.
Each decode round runs three stages — draft, verify, extend — and every
stage has shape-dependent resources (attention backends and CUDA graphs)
that must match the current configuration. Switching adaptive steps
means swapping the entire state atomically.
"""
# -- Configuration (determines shapes for all stages) --
speculative_num_steps: int
speculative_num_draft_tokens: int
# -- Draft stage: draft model multi-step autoregressive generation --
draft_attn_backend: "AttentionBackend | None"
cuda_graph_runner: "EAGLEDraftCudaGraphRunner | None"
# -- Verify stage: target model one-pass tree verification --
target_attn_backend: "AttentionBackend"
target_graph_runner: "CudaGraphRunner | CPUGraphRunner | None"
# -- Extend stage: draft model KV cache catch-up after verify --
draft_extend_attn_backend: "AttentionBackend | None"
cuda_graph_runner_for_draft_extend: "EAGLEDraftExtendCudaGraphRunner | None"
class AdaptiveSpecWorker(Protocol):
"""Protocol that a worker must implement to use AdaptiveController."""
speculative_num_steps: int
def build_adaptive_runtime_state(
self, speculative_num_steps: int, speculative_num_draft_tokens: int
) -> SpecRuntimeState: ...
def apply_runtime_state(self, state: SpecRuntimeState) -> None: ...
class AdaptiveController:
"""Facade that owns adaptive decision-making and runtime state switching.
Works with any worker that implements ``AdaptiveSpecWorker`` protocol:
- ``build_adaptive_runtime_state(steps, draft_tokens)`` → runtime state
- ``apply_runtime_state(state)`` → apply it to the worker
The worker only needs to:
1. Call ``register()`` for the initial state, then ``init_states()``
once during startup.
2. Call ``on_verify_complete(accept_lengths)`` after each decode verify.
"""
def __init__(self, worker: AdaptiveSpecWorker, config_path: str | None = None):
self.worker = worker
cfg = load_adaptive_config(config_path)
self.params = AdaptiveSpeculativeParams(
initial_steps=worker.speculative_num_steps,
config=cfg,
)
self._states: dict[int, SpecRuntimeState] = {}
@property
def candidate_steps(self) -> list[int]:
return self.params.candidate_steps
def register(self, state: SpecRuntimeState, steps: int | None = None) -> None:
"""Register a pre-built runtime state.
*steps* defaults to ``state.speculative_num_steps`` when not given.
"""
key = steps if steps is not None else state.speculative_num_steps
self._states[key] = state
def init_states(self) -> None:
"""Build and register runtime states for all candidate steps."""
for steps in self.params.candidate_steps:
if steps in self._states:
continue
state = self.worker.build_adaptive_runtime_state(
speculative_num_steps=steps,
speculative_num_draft_tokens=steps + 1,
)
self._states[steps] = state
self._activate(self.params.current_steps)
def on_verify_complete(self, accept_lengths: list[int]) -> None:
"""Feed verify results; switch runtime state if EMA warrants it."""
if self.params.update(accept_lengths):
self._activate(self.params.current_steps)
def _activate(self, speculative_num_steps: int) -> None:
state = self._states.get(speculative_num_steps)
if state is None:
raise ValueError(
f"Missing adaptive runtime state for steps={speculative_num_steps}"
)
self.worker.apply_runtime_state(state)
@@ -0,0 +1,133 @@
"""Adaptive speculative decoding parameters.
Adjusts speculative_num_steps at runtime based on observed acceptance lengths.
"""
import json
import logging
logger = logging.getLogger(__name__)
def load_adaptive_config(path: str | None) -> dict[str, object]:
"""Load adaptive speculative config from a JSON file.
The file may contain any subset of the following keys:
ema_alpha, update_interval, warmup_batches,
down_hysteresis, up_hysteresis, candidate_steps
Returns an empty dict when *path* is ``None``.
"""
if path is None:
return {}
with open(path) as f:
cfg = json.load(f)
if not isinstance(cfg, dict):
raise ValueError(
"speculative_adaptive_config must be a JSON object, "
f"got {type(cfg).__name__}"
)
return cfg
class AdaptiveSpeculativeParams:
"""Tracks acceptance rate via EMA and adapts num_steps accordingly.
The core idea: if drafts are consistently accepted, try more steps;
if drafts are consistently rejected early, reduce steps to avoid waste.
Formula: target_steps = clamp(round(ema_accept_len) + 1, min_steps, max_steps)
- Probes one step beyond observed acceptance
- EMA smoothing prevents oscillation
- Only updates every `update_interval` batches for stability
"""
def __init__(
self,
initial_steps: int,
config: dict[str, object] | None = None,
):
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])))
assert (
len(self.candidate_steps) >= 2
), "candidate_steps must have at least 2 distinct values"
self.min_steps = self.candidate_steps[0]
self.max_steps = self.candidate_steps[-1]
self.ema_alpha = cfg.get("ema_alpha", 0.2)
self.update_interval = cfg.get("update_interval", 5)
self.warmup_batches = cfg.get("warmup_batches", 10)
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),
)
# Initialize EMA at current steps - 1 (neutral starting point)
self.ema_accept_len = float(self.current_steps - 1)
self._batch_count = 0
logger.info(
f"AdaptiveSpeculativeParams initialized: "
f"steps={self.current_steps}, candidate_steps={self.candidate_steps}"
)
def update(self, accept_lengths: list[int]) -> bool:
"""Update EMA with observed accept lengths. Returns True if params changed.
Args:
accept_lengths: Per-request accepted draft token counts from last verify.
"""
if not accept_lengths:
return False
batch_avg = sum(accept_lengths) / len(accept_lengths)
self.ema_accept_len = (
1 - self.ema_alpha
) * self.ema_accept_len + self.ema_alpha * batch_avg
self._batch_count += 1
if self._batch_count <= self.warmup_batches:
return False
if (self._batch_count - self.warmup_batches) % self.update_interval != 0:
return False
return self._recompute_params()
def _recompute_params(self) -> bool:
"""Recompute steps from EMA. Returns True if params changed."""
old_steps = self.current_steps
current_idx = self.candidate_steps.index(old_steps)
# TODO: Consider limiting step changes to avoid overshooting.
while current_idx > 0:
prev_step = self.candidate_steps[current_idx - 1]
drop_threshold = prev_step - 0.5 + self.down_hysteresis
if self.ema_accept_len <= drop_threshold:
current_idx -= 1
else:
break
while current_idx < len(self.candidate_steps) - 1:
current_step = self.candidate_steps[current_idx]
rise_threshold = current_step - 0.5 + self.up_hysteresis
if self.ema_accept_len > rise_threshold:
current_idx += 1
else:
break
target = self.candidate_steps[current_idx]
if target != old_steps:
self.current_steps = target
logger.info(
f"Adaptive spec params updated: steps {old_steps} -> {target} "
f"(ema_accept_len={self.ema_accept_len:.2f})"
)
return True
return False
@@ -54,7 +54,13 @@ class EagleDraftInputBuffers(ForwardInputBuffers):
class EAGLEDraftCudaGraphRunner:
def __init__(self, eagle_worker: EAGLEWorker):
def __init__(
self,
eagle_worker: EAGLEWorker,
*,
draft_attn_backend=None,
speculative_num_steps: Optional[int] = None,
):
# Parse args
self.eagle_worker = eagle_worker
if not hasattr(eagle_worker, "model_runner"):
@@ -72,8 +78,13 @@ class EAGLEDraftCudaGraphRunner:
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.tp_size
self.dp_size = self.model_runner.dp_size
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.speculative_num_steps = (
model_runner.server_args.speculative_num_steps
if speculative_num_steps is None
else speculative_num_steps
)
self.topk = model_runner.server_args.speculative_eagle_topk
self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
@@ -88,10 +99,8 @@ class EAGLEDraftCudaGraphRunner:
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
self.model_runner.draft_attn_backend.init_cuda_graph_state(
self.max_bs, self.max_num_token
)
self.seq_len_fill_value = self.model_runner.draft_attn_backend.attn_backends[
self.draft_attn_backend.init_cuda_graph_state(self.max_bs, self.max_num_token)
self.seq_len_fill_value = self.draft_attn_backend.attn_backends[
0
].get_cuda_graph_seq_len_fill_value()
seq_lens_cpu = torch.full(
@@ -310,9 +319,7 @@ class EAGLEDraftCudaGraphRunner:
)
# Attention backend
self.model_runner.draft_attn_backend.init_forward_metadata_capture_cuda_graph(
forward_batch
)
self.draft_attn_backend.init_forward_metadata_capture_cuda_graph(forward_batch)
# Run and capture
def run_once():
@@ -409,7 +416,7 @@ class EAGLEDraftCudaGraphRunner:
buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:bs]
self.model_runner.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
self.draft_attn_backend.init_forward_metadata_replay_cuda_graph(
forward_batch, bs
)
self.raw_bs = raw_bs
@@ -56,7 +56,13 @@ class EagleDraftExtendInputBuffers(ForwardInputBuffers):
class EAGLEDraftExtendCudaGraphRunner:
def __init__(self, eagle_worker: EAGLEWorker):
def __init__(
self,
eagle_worker: EAGLEWorker,
*,
draft_extend_attn_backend=None,
speculative_num_steps: Optional[int] = None,
):
# Parse args
self.eagle_worker = eagle_worker
if not hasattr(eagle_worker, "model_runner"):
@@ -77,8 +83,15 @@ class EAGLEDraftExtendCudaGraphRunner:
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.tp_size
self.dp_size = self.model_runner.dp_size
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.speculative_num_steps = (
model_runner.server_args.speculative_num_steps
if speculative_num_steps is None
else speculative_num_steps
)
self.topk = model_runner.server_args.speculative_eagle_topk
self.draft_extend_attn_backend = (
draft_extend_attn_backend or eagle_worker.draft_extend_attn_backend
)
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
@@ -93,11 +106,11 @@ class EAGLEDraftExtendCudaGraphRunner:
self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_bs
self.eagle_worker.draft_extend_attn_backend.init_cuda_graph_state(
self.draft_extend_attn_backend.init_cuda_graph_state(
self.max_bs, self.max_num_token
)
self.seq_len_fill_value = (
self.eagle_worker.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value()
self.draft_extend_attn_backend.get_cuda_graph_seq_len_fill_value()
)
seq_lens_cpu = torch.full(
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
@@ -362,11 +375,11 @@ class EAGLEDraftExtendCudaGraphRunner:
spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info,
capture_hidden_mode=CaptureHiddenMode.LAST,
attn_backend=self.eagle_worker.draft_extend_attn_backend,
attn_backend=self.draft_extend_attn_backend,
padded_static_len=self.padded_static_len,
)
self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph(
self.draft_extend_attn_backend.init_forward_metadata_capture_cuda_graph(
bs=bs,
num_tokens=num_tokens,
req_pool_indices=req_pool_indices,
@@ -493,7 +506,7 @@ class EAGLEDraftExtendCudaGraphRunner:
forward_batch.spec_info.positions = buffers.positions[:num_tokens]
forward_batch.spec_info.accept_length = buffers.accept_length[:bs]
self.eagle_worker.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
self.draft_extend_attn_backend.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=buffers.req_pool_indices,
seq_lens=buffers.seq_lens,
+162 -4
View File
@@ -1,5 +1,6 @@
import logging
import time
from contextlib import contextmanager
from typing import List, Optional, Tuple
import torch
@@ -24,6 +25,7 @@ from sglang.srt.mem_cache.common import (
alloc_token_slots,
get_last_loc,
)
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
@@ -32,6 +34,10 @@ from sglang.srt.model_executor.forward_batch_info import (
from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.observability.trace import get_global_tracing_enabled
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
SpecRuntimeState,
)
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
@@ -105,6 +111,13 @@ class EAGLEWorker(TpModelWorker):
server_args.speculative_algorithm
)
# 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
)
# 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
@@ -206,6 +219,20 @@ class EAGLEWorker(TpModelWorker):
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
self.init_attention_backend()
self.init_cuda_graphs()
if self.adaptive_controller is not None:
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_attn_backend,
cuda_graph_runner=self.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_extend_attn_backend,
cuda_graph_runner_for_draft_extend=self.cuda_graph_runner_for_draft_extend,
)
)
self.adaptive_controller.init_states()
# Some dummy tensors
self.num_new_pages_per_topk = torch.empty(
@@ -274,6 +301,130 @@ class EAGLEWorker(TpModelWorker):
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):
"""Apply a pre-built runtime state to this worker."""
if self.speculative_num_steps == state.speculative_num_steps:
return
logger.info(
"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}"
)
self.speculative_num_steps = state.speculative_num_steps
self.speculative_num_draft_tokens = state.speculative_num_draft_tokens
# Draft stage
self.draft_attn_backend = state.draft_attn_backend
self.draft_model_runner.draft_attn_backend = state.draft_attn_backend
self.cuda_graph_runner = state.cuda_graph_runner
# Verify stage
self.target_worker.model_runner.attn_backend = state.target_attn_backend
self.target_worker.model_runner.graph_runner = state.target_graph_runner
# Extend stage
self.draft_extend_attn_backend = state.draft_extend_attn_backend
self.cuda_graph_runner_for_draft_extend = (
state.cuda_graph_runner_for_draft_extend
)
# 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
)
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
):
# Reuse existing init methods for draft attention backend and cuda graphs
self.init_attention_backend()
self.init_cuda_graphs()
# Capture target attention backend and CUDA graph
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 stage
draft_attn_backend=self.draft_attn_backend,
cuda_graph_runner=self.cuda_graph_runner,
# Verify stage
target_attn_backend=target_attn_backend,
target_graph_runner=target_graph_runner,
# Extend stage
draft_extend_attn_backend=self.draft_extend_attn_backend,
cuda_graph_runner_for_draft_extend=self.cuda_graph_runner_for_draft_extend,
)
after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info(
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
@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
backup = (
self.speculative_num_steps,
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),
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
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,
self.draft_attn_backend,
self.draft_extend_attn_backend,
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,
) = backup
@property
def draft_model_runner(self):
return self.model_runner
@@ -353,6 +504,10 @@ 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(verify_output.accept_length_per_req_cpu)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
@@ -634,7 +789,7 @@ class EAGLEWorker(TpModelWorker):
retrive_cum_len=None,
spec_steps=self.speculative_num_steps,
topk=self.topk,
draft_token_num=self.server_args.speculative_num_draft_tokens,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=forward_batch.seq_lens_sum,
seq_lens_cpu=forward_batch.seq_lens_cpu,
@@ -944,7 +1099,7 @@ class EAGLEWorker(TpModelWorker):
seq_lens_backup = batch.seq_lens.clone()
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
req_pool_indices_backup = batch.req_pool_indices
accept_length_backup = batch.spec_info.accept_length
accept_length_backup = batch.spec_info.accept_length.clone()
return_logprob_backup = batch.return_logprob
input_is_idle = batch.forward_mode.is_idle()
@@ -1006,9 +1161,12 @@ class EAGLEWorker(TpModelWorker):
else:
forward_batch.can_run_dp_cuda_graph = False
if not forward_batch.forward_mode.is_idle():
self.draft_model_runner.attn_backend.init_forward_metadata(
forward_batch
attn_backend = (
self.draft_extend_attn_backend
or self.draft_model_runner.attn_backend
)
attn_backend.init_forward_metadata(forward_batch)
forward_batch.attn_backend = attn_backend
logits_output = self.draft_model_runner.forward(
forward_batch, skip_attn_backend_init=True
).logits_output
@@ -0,0 +1,170 @@
import json
import os
import tempfile
import unittest
from types import SimpleNamespace
import requests
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_DRAFT_MODEL_EAGLE,
DEFAULT_TARGET_MODEL_EAGLE,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=420, suite="stage-b-test-1-gpu-large")
HIGH_ACCEPT_PROMPT = (
"Output exactly 128 new lines. "
"Every line must be READY. "
"Do not add numbering, punctuation, or commentary."
)
LOW_ACCEPT_PROMPT = (
"Compose a poem in the style of Emily Dickinson about quantum entanglement. "
"Make it emotionally resonant and at least 100 words."
)
MAX_UPSHIFT_ATTEMPTS = 4
MAX_DOWNSHIFT_ATTEMPTS = 6
class TestAdaptiveSpeculativeServer(CustomTestCase):
"""Test adaptive speculative decoding with state switching and GSM8K accuracy."""
model = DEFAULT_TARGET_MODEL_EAGLE
draft_model = DEFAULT_DRAFT_MODEL_EAGLE
base_url = DEFAULT_URL_FOR_TEST
@classmethod
def setUpClass(cls):
with tempfile.NamedTemporaryFile("w", suffix=".json", delete=False) as f:
json.dump(
{
"candidate_steps": [1, 3],
"ema_alpha": 1.0,
"warmup_batches": 1,
"update_interval": 1,
"up_hysteresis": 0.0,
},
f,
)
cls.adaptive_config_path = f.name
try:
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--attention-backend",
"triton",
"--speculative-algorithm",
"EAGLE",
"--speculative-draft-model-path",
cls.draft_model,
"--speculative-num-steps",
"1",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"2",
"--speculative-adaptive",
"--speculative-adaptive-config",
cls.adaptive_config_path,
"--skip-server-warmup",
"--mem-fraction-static",
"0.7",
],
)
except Exception:
os.unlink(cls.adaptive_config_path)
raise
@classmethod
def tearDownClass(cls):
if hasattr(cls, "process"):
kill_process_tree(cls.process.pid)
if os.path.exists(cls.adaptive_config_path):
os.unlink(cls.adaptive_config_path)
def _get_internal_state(self) -> dict:
response = requests.get(self.base_url + "/server_info", timeout=30)
self.assertEqual(response.status_code, 200, response.text)
return response.json()["internal_states"][0]
def _generate(self, prompt: str, max_new_tokens: int = 64) -> dict:
response = requests.post(
self.base_url + "/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": max_new_tokens,
"ignore_eos": True,
},
},
timeout=180,
)
self.assertEqual(response.status_code, 200, response.text)
return response.json()
def _drive_upshift(self) -> dict:
"""Send high-acceptance prompts until steps upshift to 3."""
state = self._get_internal_state()
for _ in range(MAX_UPSHIFT_ATTEMPTS):
self._generate(HIGH_ACCEPT_PROMPT)
state = self._get_internal_state()
if state["speculative_num_steps"] == 3:
return state
return state
def _drive_downshift(self) -> dict:
"""Send low-acceptance prompts until steps downshift to 1."""
state = self._get_internal_state()
for _ in range(MAX_DOWNSHIFT_ATTEMPTS):
self._generate(LOW_ACCEPT_PROMPT)
state = self._get_internal_state()
if state["speculative_num_steps"] == 1:
return state
return state
def test_gsm8k_after_adaptive_switches(self):
"""Exercise up/down/up adaptive switches, then verify GSM8K accuracy."""
state = self._drive_upshift()
self.assertEqual(state["speculative_num_steps"], 3, f"Never upshifted: {state}")
state = self._drive_downshift()
self.assertEqual(
state["speculative_num_steps"], 1, f"Never downshifted: {state}"
)
self._drive_upshift()
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=100,
num_threads=64,
)
metrics = run_eval(args)
print(f"GSM8K after adaptive switches: {metrics}")
self.assertGreater(metrics["score"], 0.20)
server_info = requests.get(self.base_url + "/server_info").json()
avg_accept_len = server_info["internal_states"][0]["avg_spec_accept_length"]
print(f"avg_spec_accept_length={avg_accept_len:.4f}")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,195 @@
import unittest
from sglang.srt.speculative.adaptive_spec_params import AdaptiveSpeculativeParams
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=2, suite="stage-a-test-cpu")
class TestAdaptiveSpeculativeParams(unittest.TestCase):
def test_initial_steps_snap_to_nearest_candidate_preferring_larger_step(self):
params = AdaptiveSpeculativeParams(
initial_steps=2,
config={"candidate_steps": [1, 3, 7]},
)
self.assertEqual(params.current_steps, 3)
self.assertEqual(params.ema_accept_len, 2.0)
def test_update_respects_warmup_and_interval(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 1,
"update_interval": 2,
},
)
self.assertFalse(params.update([0, 0]))
self.assertEqual(params.current_steps, 3)
self.assertFalse(params.update([0, 0]))
self.assertEqual(params.current_steps, 3)
self.assertTrue(params.update([0, 0]))
self.assertEqual(params.current_steps, 1)
def test_empty_batches_do_not_consume_warmup_or_shift_steps(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 1,
"update_interval": 1,
},
)
self.assertFalse(params.update([]))
self.assertEqual(params.current_steps, 3)
self.assertEqual(params.ema_accept_len, 2.0)
self.assertFalse(params.update([0, 0]))
self.assertEqual(params.current_steps, 3)
self.assertTrue(params.update([0, 0]))
self.assertEqual(params.current_steps, 1)
def test_update_scales_up_across_candidates(self):
params = AdaptiveSpeculativeParams(
initial_steps=1,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
"up_hysteresis": 0.0,
},
)
self.assertTrue(params.update([1, 1]))
self.assertEqual(params.current_steps, 3)
self.assertTrue(params.update([3, 3]))
self.assertEqual(params.current_steps, 7)
def test_update_can_scale_down_across_candidates_in_one_recompute(self):
params = AdaptiveSpeculativeParams(
initial_steps=7,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
},
)
self.assertTrue(params.update([0, 0]))
self.assertEqual(params.current_steps, 1)
def test_exact_rise_threshold_does_not_upshift(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
"up_hysteresis": 0.0,
},
)
self.assertFalse(params.update([2, 3]))
self.assertEqual(params.current_steps, 3)
self.assertEqual(params.ema_accept_len, 2.5)
self.assertTrue(params.update([3, 3]))
self.assertEqual(params.current_steps, 7)
def test_exact_drop_threshold_does_downshift(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
"down_hysteresis": 0.0,
"up_hysteresis": 0.5,
},
)
self.assertTrue(params.update([0, 1]))
self.assertEqual(params.current_steps, 1)
self.assertEqual(params.ema_accept_len, 0.5)
def test_hysteresis_can_prevent_premature_upshift(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
"up_hysteresis": 0.75,
},
)
self.assertFalse(params.update([3, 3]))
self.assertEqual(params.current_steps, 3)
self.assertTrue(params.update([4, 4]))
self.assertEqual(params.current_steps, 7)
def test_down_hysteresis_can_prevent_premature_downshift(self):
params = AdaptiveSpeculativeParams(
initial_steps=7,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 1.0,
"warmup_batches": 0,
"update_interval": 1,
"down_hysteresis": -0.75,
},
)
self.assertFalse(params.update([2, 2]))
self.assertEqual(params.current_steps, 7)
self.assertTrue(params.update([1, 1]))
self.assertEqual(params.current_steps, 3)
def test_multi_batch_sequence_can_ramp_up_then_back_down(self):
params = AdaptiveSpeculativeParams(
initial_steps=3,
config={
"candidate_steps": [1, 3, 7],
"ema_alpha": 0.5,
"warmup_batches": 0,
"update_interval": 1,
"up_hysteresis": 0.0,
"down_hysteresis": 0.0,
},
)
self.assertTrue(params.update([4, 4]))
self.assertEqual(params.current_steps, 7)
self.assertEqual(params.ema_accept_len, 3.0)
self.assertTrue(params.update([0, 0]))
self.assertEqual(params.current_steps, 3)
self.assertEqual(params.ema_accept_len, 1.5)
self.assertFalse(params.update([0, 0]))
self.assertEqual(params.current_steps, 3)
self.assertEqual(params.ema_accept_len, 0.75)
self.assertTrue(params.update([0, 0]))
self.assertEqual(params.current_steps, 1)
self.assertEqual(params.ema_accept_len, 0.375)
if __name__ == "__main__":
unittest.main()