[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
@@ -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()