Refactor EAGLE infer tests: shared fixture + kits + overlap matrix (#26871)

This commit is contained in:
Liangsheng Yin
2026-06-01 03:55:01 -07:00
committed by GitHub
parent 89410b380b
commit 1bff7a290f
12 changed files with 1147 additions and 872 deletions
+566
View File
@@ -0,0 +1,566 @@
"""Reusable test-method mixins (kits) for EAGLE/EAGLE3 spec-decoding servers.
Pair these with ``SpecEagleServerBase`` (sglang.test.server_fixtures.spec_eagle_fixture).
Each kit is a cohesive group of ``test_*`` methods with no launch logic; concrete
test classes mix in the fixture (which owns launch knobs) + whichever kits apply.
Thresholds are read off ``self`` so a config can tune them as class attributes.
"""
import concurrent.futures
import json
import random
import threading
from concurrent.futures import ThreadPoolExecutor
from functools import partial
from types import SimpleNamespace
import numpy as np
import requests
from sglang.srt.utils.common import kill_process_tree
from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
popen_launch_server,
run_logprob_check,
)
class SpecCorrectnessKit:
"""Acceptance-quality + EOS checks (single server, cheap)."""
# Tunable thresholds (override per config class).
acc_length_thres = 3.1
batch_accept_len_thres = 1.75
def test_acc_length(self):
prompt = [
"Human: Give me a fully functional FastAPI server. Show the python code.\n\nAssistant:",
] * 5
sampling_params = {"temperature": 0, "max_new_tokens": 512}
output = requests.post(
self.base_url + "/generate",
json={"text": prompt, "sampling_params": sampling_params},
).json()[0]
meta = output["meta_info"]
if "spec_verify_ct" in meta and meta["spec_verify_ct"] > 0:
acc_length = meta["completion_tokens"] / meta["spec_verify_ct"]
else:
acc_length = 1.0
print(f"{acc_length=:.4f}")
self.assertGreater(acc_length, self.acc_length_thres)
def test_batch_generation(self):
prompts = [
"Hello, my name is",
"The president of the United States is",
"The capital of France is",
"The future of AI is",
]
results = requests.post(
self.base_url + "/generate",
json={
"text": prompts,
"sampling_params": {"temperature": 0, "max_new_tokens": 50},
},
).json()
# Accept length from per-request meta_info (self-contained). The
# internal_states `avg_spec_accept_length` isn't populated on the v1 /
# disable-overlap path after a small batch, so don't read server_info.
total_completion, total_verify = 0, 0
for r in results:
self.assertIn("text", r, f"Server error: {r}")
meta = r["meta_info"]
total_completion += meta["completion_tokens"]
total_verify += meta.get("spec_verify_ct", 0)
if total_verify > 0:
acc_length = total_completion / total_verify
print(f"batch {acc_length=:.4f}")
self.assertGreater(acc_length, self.batch_accept_len_thres)
def test_eos_token(self):
prompt = "[INST] <<SYS>>\nYou are a helpful assistant.\n<</SYS>>\nToday is a sunny day and I like [/INST]"
res = requests.post(
self.base_url + "/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0.1,
"max_new_tokens": 1024,
"skip_special_tokens": False,
},
},
).json()
output = res["text"]
tokens = self.tokenizer.encode(output, truncation=False)
self.assertNotIn(self.tokenizer.eos_token_id, tokens)
def test_first_token_finish(self):
# Very short max_new_tokens (1-3): exercise the immediate-finish path,
# where a request stops within the first draft window. Just must not crash.
prompts = [
f"There are {i} apples on the table. How to divide them equally?"
for i in range(8)
]
sampling_params = [
{"temperature": 0, "max_new_tokens": random.randint(1, 3)} for _ in range(8)
]
results = requests.post(
self.base_url + "/generate",
json={"text": prompts, "sampling_params": sampling_params},
).json()
for r in results:
self.assertIn("text", r, f"Server error: {r}")
def _greedy(url, text, max_new_tokens=48):
return requests.post(
url + "/generate",
json={
"text": text,
"sampling_params": {"temperature": 0, "max_new_tokens": max_new_tokens},
},
).json()["text"]
class SpecParityKit:
"""Lossless output parity vs a non-spec reference.
Sequential (NOT concurrent): launch a non-spec reference server on the
standard port, capture greedy outputs, tear it down, THEN let the fixture
launch the spec server. Only one model is resident at a time -- two 8B
servers don't fit on one GPU. Mix this kit FIRST in the bases so its
setUpClass runs before the fixture's: ``class T(SpecParityKit, Eagle3Base)``.
"""
parity_prompts = [
"The capital of France is",
"Once upon a time, there was a",
"The three primary colors are",
"def fibonacci(n):",
]
@classmethod
def setUpClass(cls):
ref_url = DEFAULT_URL_FOR_TEST
ref_proc = popen_launch_server(
cls.model,
ref_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--mem-fraction-static",
"0.8", # ref alone -> full GPU available
"--attention-backend",
cls.attention_backend,
"--page-size",
"1",
"--dtype",
cls.dtype,
*(["--trust-remote-code"] if cls.trust_remote_code else []),
],
)
try:
cls.parity_ref_outputs = {
p: _greedy(ref_url, p) for p in cls.parity_prompts
}
finally:
kill_process_tree(ref_proc.pid, wait_timeout=60)
# Now the spec server (same port; ref is gone).
super().setUpClass()
def test_parity_vs_reference(self):
"""Spec decode greedy output must equal the non-spec reference."""
for prompt in self.parity_prompts:
spec_out = _greedy(self.base_url, prompt)
self.assertEqual(
spec_out,
self.parity_ref_outputs[prompt],
f"spec != ref for prompt {prompt!r}",
)
class SpecAccuracyKit:
"""gsm8k accuracy + acceptance length, and throughput at max_tokens=1."""
gsm8k_num_examples = 200
gsm8k_score_thres = 0.20
gsm8k_check_accept_len = True
# If set, use this; else fall back to topk-based default (2.5 / 3.47).
gsm8k_accept_len_thres = None
def test_gsm8k(self):
requests.get(self.base_url + "/flush_cache")
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=self.gsm8k_num_examples,
num_threads=128,
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreater(metrics["score"], self.gsm8k_score_thres)
if self.gsm8k_check_accept_len:
server_info = requests.get(self.base_url + "/server_info").json()
avg_spec_accept_length = server_info["internal_states"][0].get(
"avg_spec_accept_length"
)
print(f"{avg_spec_accept_length=}")
# The metric isn't always populated (e.g. v1 / disable-overlap).
# Only enforce the threshold when it's reported.
if avg_spec_accept_length is not None:
topk = server_info["speculative_eagle_topk"]
thres = self.gsm8k_accept_len_thres
if thres is None:
thres = 2.5 if topk == 1 else 3.47
self.assertGreater(avg_spec_accept_length, thres)
class SpecPerfKit:
"""Throughput perf check (GPU-specific -> run on the reference/Hopper runner)."""
perf_output_throughput_thres = 50
def test_max_token_one(self):
requests.get(self.base_url + "/flush_cache")
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=1,
num_examples=200,
num_threads=128,
)
metrics = run_eval(args)
self.assertGreater(
metrics["output_throughput"], self.perf_output_throughput_thres
)
class SpecLogprobKit:
"""Logprob correctness: start_len, prefill-rescore match, mixed sweep,
spec-v2 decode-vs-prefill match, and ragged token_ids_logprob."""
def test_logprob_start_len(self):
logprob_start_len = 4
new_tokens = 4
prompts = [
"I have a very good idea on",
"Today is a sunndy day and",
]
response = requests.post(
self.base_url + "/generate",
json={
"text": prompts,
"sampling_params": {
"temperature": 0,
"max_new_tokens": new_tokens,
},
"return_logprob": True,
"top_logprobs_num": 5,
"logprob_start_len": logprob_start_len,
},
)
response_json = response.json()
for res in response_json:
self.assertEqual(
res["meta_info"]["prompt_tokens"],
logprob_start_len + len(res["meta_info"]["input_token_logprobs"]),
)
self.assertEqual(res["meta_info"]["completion_tokens"], new_tokens)
self.assertEqual(len(res["meta_info"]["output_token_logprobs"]), new_tokens)
def test_logprob_match(self):
"""Output logprobs should match a fresh prefill of the same sequence."""
def run_generate(
prompt,
return_logprob=False,
max_new_tokens=512,
logprob_start_len=-1,
temperature=1.0,
):
if isinstance(prompt, str):
prompt_kwargs = {"text": prompt}
else:
prompt_kwargs = {"input_ids": prompt}
response = requests.post(
self.base_url + "/generate",
json={
**prompt_kwargs,
"sampling_params": {
"temperature": temperature,
"max_new_tokens": max_new_tokens,
"ignore_eos": True,
},
"return_logprob": return_logprob,
"return_text_in_logprobs": True,
"logprob_start_len": logprob_start_len,
},
)
return response.json()
prompt = "I have a very good idea on how to"
for temperature in [1.0]:
gen = run_generate(
prompt,
return_logprob=True,
logprob_start_len=0,
temperature=temperature,
)
output_logprobs = np.array(
[x[0] for x in gen["meta_info"]["output_token_logprobs"]]
)
num_prompts_tokens = gen["meta_info"]["prompt_tokens"]
input_tokens = [x[1] for x in gen["meta_info"]["input_token_logprobs"]]
output_tokens = [x[1] for x in gen["meta_info"]["output_token_logprobs"]]
new_prompt = input_tokens + output_tokens
score = run_generate(
new_prompt,
return_logprob=True,
logprob_start_len=0,
max_new_tokens=0,
temperature=temperature,
)
output_logprobs_score = np.array(
[
x[0]
for x in score["meta_info"]["input_token_logprobs"][
num_prompts_tokens:
]
]
)
diff = np.abs(output_logprobs - output_logprobs_score)
max_diff = np.max(diff)
self.assertLess(max_diff, 0.255)
def test_logprob_mixed(self):
args = []
temperature = 0
# input_len, output_len, temperature, logprob_start_len, return_logprob, top_logprobs_num
for input_len in [200, 500, 1000, 2000]:
for output_len in [4, 8]:
for logprob_start_len in [0, 100, 300, 800, 1998]:
for return_logprob in [True, False]:
for top_logprobs_num in [0, 5]:
if logprob_start_len >= input_len:
continue
args.append(
(
input_len,
output_len,
temperature,
logprob_start_len,
return_logprob,
top_logprobs_num,
)
)
random.shuffle(args)
func = partial(run_logprob_check, self)
with ThreadPoolExecutor(8) as executor:
list(executor.map(func, args))
def test_logprob_spec_v2_match(self):
"""Verify spec v2 decode logprobs match prefill scoring logprobs."""
top_k = 5
probe_token_ids = [1, 2, 10, 100, 1000]
prompts = [
"The capital of France is",
"Explain quantum computing in simple terms:",
]
for round_idx, prompt in enumerate(prompts):
with self.subTest(round=round_idx, prompt=prompt):
gen_res = requests.post(
self.base_url + "/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 32,
"ignore_eos": True,
},
"return_logprob": True,
"top_logprobs_num": top_k,
"token_ids_logprob": probe_token_ids,
"logprob_start_len": 0,
},
).json()
decode_logprobs = gen_res["meta_info"]["output_token_logprobs"]
decode_top_logprobs = gen_res["meta_info"]["output_top_logprobs"]
decode_tid_logprobs = gen_res["meta_info"]["output_token_ids_logprobs"]
input_token_ids = [
t[1] for t in gen_res["meta_info"]["input_token_logprobs"]
]
output_token_ids = [t[1] for t in decode_logprobs]
num_prompt_tokens = gen_res["meta_info"]["prompt_tokens"]
score_res = requests.post(
self.base_url + "/generate",
json={
"input_ids": input_token_ids + output_token_ids,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 0,
},
"return_logprob": True,
"top_logprobs_num": top_k,
"token_ids_logprob": probe_token_ids,
"logprob_start_len": 0,
},
).json()
score_logprobs = score_res["meta_info"]["input_token_logprobs"][
num_prompt_tokens:
]
score_top_logprobs = score_res["meta_info"]["input_top_logprobs"][
num_prompt_tokens:
]
score_tid_logprobs = score_res["meta_info"]["input_token_ids_logprobs"][
num_prompt_tokens:
]
self.assertEqual(len(decode_logprobs), len(score_logprobs))
decode_vals = np.array([t[0] for t in decode_logprobs])
score_vals = np.array([t[0] for t in score_logprobs])
max_diff = np.max(np.abs(decode_vals - score_vals))
print(f"[round {round_idx}] logprob max_diff={max_diff:.6f}")
self.assertLess(max_diff, 0.255)
for pos in range(len(decode_logprobs)):
dec_top = {t[1]: t[0] for t in decode_top_logprobs[pos]}
scr_top = {t[1]: t[0] for t in score_top_logprobs[pos]}
common_ids = set(dec_top.keys()) & set(scr_top.keys())
self.assertGreater(len(common_ids), 0)
for tid in common_ids:
self.assertAlmostEqual(dec_top[tid], scr_top[tid], delta=0.255)
self.assertEqual(len(decode_tid_logprobs), len(score_tid_logprobs))
for pos in range(len(decode_tid_logprobs)):
dec_tid = {t[1]: t[0] for t in decode_tid_logprobs[pos]}
scr_tid = {t[1]: t[0] for t in score_tid_logprobs[pos]}
self.assertEqual(set(dec_tid.keys()), set(scr_tid.keys()))
for tid in dec_tid:
self.assertAlmostEqual(dec_tid[tid], scr_tid[tid], delta=0.255)
def test_token_ids_logprob_ragged(self):
"""Regression: ragged token_ids_logprob lists in one batch must not crash."""
def send(probe_ids):
return requests.post(
self.base_url + "/generate",
json={
"text": "Hello world",
"sampling_params": {"temperature": 0, "max_new_tokens": 8},
"return_logprob": True,
"top_logprobs_num": 3,
"token_ids_logprob": probe_ids,
},
).json()
ragged_probes = [
[1, 2],
[3, 4, 5],
[6],
[10, 20, 30, 40],
[1, 2],
[3, 4, 5],
[6],
[10, 20, 30, 40],
]
with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool:
futs = [pool.submit(send, ids) for ids in ragged_probes]
for f in concurrent.futures.as_completed(futs):
res = f.result()
self.assertIn("text", res, f"Server error: {res}")
class SpecPenaltyKit:
"""Penalty parameters under concurrency must not crash / corrupt output."""
def test_penalty_mixed(self):
args = [
{},
{},
{},
{"frequency_penalty": 2},
{"presence_penalty": 1},
{"min_new_tokens": 16},
{"frequency_penalty": 0.2},
{"presence_penalty": 0.4},
{"min_new_tokens": 8},
{"frequency_penalty": 0.4, "presence_penalty": 0.8},
{"frequency_penalty": 0.4, "min_new_tokens": 12},
{"presence_penalty": 0.8, "min_new_tokens": 12},
{"presence_penalty": -0.3, "frequency_penalty": 1.3, "min_new_tokens": 32},
{"presence_penalty": 0.3, "frequency_penalty": -1.3, "min_new_tokens": 32},
]
random.shuffle(args * 5)
with ThreadPoolExecutor(8) as executor:
list(executor.map(self.run_decode, args))
class SpecFeatureKit:
"""Radix attention, constrained decoding, concurrent abort."""
def test_radix_attention(self):
run_radix_attention_test(self.base_url)
self.assertIsNone(self.process.poll())
def test_request_abort(self):
concurrency = 4
threads = [
threading.Thread(target=self.send_request) for _ in range(concurrency)
] + [
threading.Thread(target=self.send_requests_abort)
for _ in range(concurrency)
]
for worker in threads:
worker.start()
for p in threads:
p.join()
def test_constrained_decoding(self):
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Give me a json"},
]
response = requests.post(
self.base_url + "/v1/chat/completions",
json={
"model": self.model,
"messages": messages,
"temperature": 0,
"response_format": {"type": "json_object"},
},
)
self.assertEqual(response.status_code, 200)
res = response.json()
self.assertIn("choices", res)
self.assertEqual(len(res["choices"]), 1)
self.assertIn("message", res["choices"][0])
self.assertIn("content", res["choices"][0]["message"])
content_json = res["choices"][0]["message"]["content"]
try:
content = json.loads(content_json)
self.assertIsInstance(content, dict)
except Exception:
self.fail(f"parse JSON failed: {content_json}")
@@ -0,0 +1,227 @@
"""Unified EAGLE/EAGLE3 speculative-decoding server fixture.
A single popen-server base whose launch is fully described by class attributes,
so concrete test classes only flip knobs (overlap on/off, model, topk, page size,
backend, env overrides). Pair it with the kits in
``sglang.test.kits.spec_server_kits`` to assemble test classes.
The primary axis is ``disable_overlap``:
- ``False`` -> spec v2 (overlap scheduler)
- ``True`` -> spec v1 (overlap disabled)
"""
import contextlib
import random
import time
import requests
from sglang.srt.environ import envs
from sglang.srt.utils.common import kill_process_tree
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
from sglang.test.test_utils import (
DEFAULT_DRAFT_MODEL_EAGLE,
DEFAULT_DRAFT_MODEL_EAGLE3,
DEFAULT_TARGET_MODEL_EAGLE,
DEFAULT_TARGET_MODEL_EAGLE3,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
# Chat-style prompts shared by send_request / send_requests_abort.
PROMPTS = [
"[INST] <<SYS>>\\nYou are a helpful assistant.\\n<</SYS>>\\nToday is a sunny day and I like[/INST]"
'[INST] <<SYS>>\\nYou are a helpful assistant.\\n<</SYS>>\\nWhat are the mental triggers in Jeff Walker\'s Product Launch Formula and "Launch" book?[/INST]',
"[INST] <<SYS>>\\nYou are a helpful assistant.\\n<</SYS>>\\nSummarize Russell Brunson's Perfect Webinar Script...[/INST]",
"[INST] <<SYS>>\\nYou are a helpful assistant.\\n<</SYS>>\\nwho are you?[/INST]",
"[INST] <<SYS>>\\nYou are a helpful assistant.\\n<</SYS>>\\nwhere are you from?[/INST]",
]
class SpecEagleServerBase(CustomTestCase):
"""Launch a single EAGLE/EAGLE3 spec server from class-attribute knobs."""
# -- model --
model = DEFAULT_TARGET_MODEL_EAGLE3
draft_model = DEFAULT_DRAFT_MODEL_EAGLE3
spec_algo = "EAGLE3"
# -- speculative config --
spec_steps = 5
spec_topk = 1
spec_tokens = 6
# -- runtime config --
page_size = 1
attention_backend = "flashinfer"
# Primary axis: False -> spec v2 (overlap); True -> spec v1 (overlap off).
disable_overlap = False
mem_fraction_static = 0.75
max_running_requests = 8
chunked_prefill_size = 128
dtype = "float16"
cuda_graph_max_bs = None
trust_remote_code = True
# -- extras --
# env_overrides: iterable of (env_var_obj, value) applied only around launch.
env_overrides = ()
extra_args = ()
@classmethod
def _launch_args(cls):
args = [
"--speculative-algorithm",
cls.spec_algo,
"--speculative-draft-model-path",
cls.draft_model,
"--speculative-num-steps",
str(cls.spec_steps),
"--speculative-eagle-topk",
str(cls.spec_topk),
"--speculative-num-draft-tokens",
str(cls.spec_tokens),
"--page-size",
str(cls.page_size),
"--attention-backend",
cls.attention_backend,
"--mem-fraction-static",
str(cls.mem_fraction_static),
"--max-running-requests",
str(cls.max_running_requests),
"--chunked-prefill-size",
str(cls.chunked_prefill_size),
"--dtype",
cls.dtype,
]
if cls.disable_overlap:
args.append("--disable-overlap-schedule")
if cls.trust_remote_code:
args.append("--trust-remote-code")
if cls.cuda_graph_max_bs is not None:
args += ["--cuda-graph-max-bs", str(cls.cuda_graph_max_bs)]
args += [str(a) for a in cls.extra_args]
return args
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
# Alias so kit methods can use either name.
cls.target_model = cls.model
cls._tokenizer = None
with contextlib.ExitStack() as stack:
stack.enter_context(envs.SGLANG_ENABLE_ASYNC_ASSERT.override(True))
stack.enter_context(
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True)
)
for env_var, value in cls.env_overrides:
stack.enter_context(env_var.override(value))
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=cls._launch_args(),
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid, wait_timeout=60)
@property
def tokenizer(self):
if type(self)._tokenizer is None:
type(self)._tokenizer = get_tokenizer(self.model)
return type(self)._tokenizer
# -- helpers used by kits --
def send_request(self):
time.sleep(random.uniform(0, 2))
for prompt in PROMPTS:
url = self.base_url + "/generate"
data = {
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 1024,
},
}
response = requests.post(url, json=data)
assert response.status_code == 200
def send_requests_abort(self):
for prompt in PROMPTS:
try:
time.sleep(random.uniform(0, 2))
url = self.base_url + "/generate"
data = {
"model": "base",
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 1024,
},
}
# set timeout = 1s, mock disconnected
requests.post(url, json=data, timeout=1)
except Exception as e:
print(e)
pass
def run_decode(self, sampling_params):
response = requests.post(
self.base_url + "/generate",
json={
"text": "Human: Write a travel blog post to Hawaii.\n\nAssistant:",
"sampling_params": {
"max_new_tokens": 48,
"n": 1,
"temperature": 0.7,
**sampling_params,
},
"return_logprob": True,
"top_logprobs_num": 5,
"return_text_in_logprobs": True,
"logprob_start_len": 0,
},
)
self.assertEqual(response.status_code, 200)
class Eagle3Base(SpecEagleServerBase):
"""EAGLE3 (Llama-3.1) config preset, topk=1 / page_size=1 by default."""
model = DEFAULT_TARGET_MODEL_EAGLE3
draft_model = DEFAULT_DRAFT_MODEL_EAGLE3
spec_algo = "EAGLE3"
spec_steps = 5
spec_topk = 1
spec_tokens = 6
attention_backend = "flashinfer"
chunked_prefill_size = 1024
# EAGLE3 topk=1 accepts modestly; tune against CI if needed.
acc_length_thres = 1.6
batch_accept_len_thres = 1.3
gsm8k_score_thres = 0.7
gsm8k_accept_len_thres = 1.3
class EagleLlama2Base(SpecEagleServerBase):
"""EAGLE (Llama-2) config preset. topk=8 tree -> spec v1; gsm8k is low."""
model = DEFAULT_TARGET_MODEL_EAGLE
draft_model = DEFAULT_DRAFT_MODEL_EAGLE
spec_algo = "EAGLE"
spec_steps = 5
spec_topk = 8
spec_tokens = 64
attention_backend = "flashinfer"
chunked_prefill_size = 128
mem_fraction_static = 0.7
gsm8k_score_thres = 0.20
acc_length_thres = 3.0
batch_accept_len_thres = 1.8
# EAGLE topk>1 already routes to v1; force it explicitly to preserve intent.
env_overrides = ((envs.SGLANG_ENABLE_SPEC_V2, False),)
@@ -1,203 +0,0 @@
import random
import unittest
import sglang as sgl
from sglang.srt.environ import envs
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_DRAFT_MODEL_EAGLE,
DEFAULT_DRAFT_MODEL_EAGLE3,
DEFAULT_TARGET_MODEL_EAGLE,
DEFAULT_TARGET_MODEL_EAGLE3,
CustomTestCase,
)
register_cuda_ci(est_time=357, stage="extra-a", runner_config="1-gpu-large")
class TestEAGLEEngine(CustomTestCase):
BASE_CONFIG = {
"model_path": DEFAULT_TARGET_MODEL_EAGLE,
"speculative_draft_model_path": DEFAULT_DRAFT_MODEL_EAGLE,
"speculative_algorithm": "EAGLE",
"speculative_num_steps": 5,
"speculative_eagle_topk": 4,
"speculative_num_draft_tokens": 8,
"mem_fraction_static": 0.7,
"cuda_graph_max_bs": 5,
"trust_remote_code": True,
}
NUM_CONFIGS = 2
THRESHOLDS = {
"batch_avg_accept_len": 1.9,
"accept_len": 3.6,
}
@classmethod
def setUpClass(cls):
envs.SGLANG_ENABLE_SPEC_V2.set(False)
@classmethod
def tearDownClass(cls):
envs.SGLANG_ENABLE_SPEC_V2.clear()
def setUp(self):
self.prompt = "Today is a sunny day and I like"
self.sampling_params = {"temperature": 0, "max_new_tokens": 8}
ref_engine = sgl.Engine(
model_path=self.BASE_CONFIG["model_path"], cuda_graph_max_bs=1
)
self.ref_output = ref_engine.generate(self.prompt, self.sampling_params)["text"]
ref_engine.shutdown()
def test_correctness(self):
configs = [
# Basic config
self.BASE_CONFIG,
# Chunked prefill
{**self.BASE_CONFIG, "chunked_prefill_size": 4},
]
for i, config in enumerate(configs[: self.NUM_CONFIGS]):
with self.subTest(i=i):
print(f"{config=}")
engine = sgl.Engine(**config, log_level="info", decode_log_interval=10)
try:
self._test_single_generation(engine)
self._test_first_token_finish(engine)
self._test_batch_generation(engine)
self._test_eos_token(engine)
self._test_acc_length(engine)
finally:
engine.flush_cache() # check engine alive
engine.shutdown()
print("=" * 100)
def _test_single_generation(self, engine):
output = engine.generate(self.prompt, self.sampling_params)["text"]
print(f"{output=}, {self.ref_output=}")
self.assertEqual(output, self.ref_output)
def _test_batch_generation(self, engine):
prompts = [
"Hello, my name is",
"The president of the United States is",
"The capital of France is",
"The future of AI is",
]
params = {"temperature": 0, "max_new_tokens": 50}
outputs = engine.generate(prompts, params)
for prompt, output in zip(prompts, outputs):
print(f"Prompt: {prompt}")
print(f"Generated: {output['text']}")
print("-" * 40)
print(f"{engine.get_server_info()=}")
avg_spec_accept_length = engine.get_server_info()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
self.assertGreater(
avg_spec_accept_length, self.THRESHOLDS["batch_avg_accept_len"]
)
def _test_first_token_finish(self, engine):
prompt = [
f"There are {i} apples on the table. How to divide them equally?"
for i in range(8)
]
params = [
{"temperature": 0, "max_new_tokens": random.randint(1, 3)} for _ in range(8)
]
outputs = engine.generate(prompt, params)
for i, output in enumerate(outputs):
print(f"Prompt: {prompt[i]}")
print(f"Generated: {output['text']}")
print("-" * 40)
def _test_eos_token(self, engine):
prompt = "[INST] <<SYS>>\nYou are a helpful assistant.\n<</SYS>>\nToday is a sunny day and I like [/INST]"
params = {
"temperature": 0.1,
"max_new_tokens": 1024,
"skip_special_tokens": False,
}
tokenizer = get_tokenizer(DEFAULT_TARGET_MODEL_EAGLE)
output = engine.generate(prompt, params)["text"]
print(f"{output=}")
tokens = tokenizer.encode(output, truncation=False)
self.assertNotIn(tokenizer.eos_token_id, tokens)
def _test_acc_length(self, engine):
prompt = [
"Human: Give me a fully functional FastAPI server. Show the python code.\n\nAssistant:",
] * 5 # test batched generation
sampling_params = {"temperature": 0, "max_new_tokens": 512}
output = engine.generate(prompt, sampling_params)
output = output[0]
if "spec_verify_ct" in output["meta_info"]:
acc_length = (
output["meta_info"]["completion_tokens"]
/ output["meta_info"]["spec_verify_ct"]
)
else:
acc_length = 1.0
speed = (
output["meta_info"]["completion_tokens"]
/ output["meta_info"]["e2e_latency"]
)
print(f"{acc_length=:.4f}, {speed=}")
self.assertGreater(acc_length, self.THRESHOLDS["accept_len"])
class TestEAGLEEngineTokenMap(TestEAGLEEngine):
BASE_CONFIG = {
"model_path": "meta-llama/Meta-Llama-3-8B-Instruct",
"speculative_draft_model_path": "lmsys/sglang-EAGLE-LLaMA3-Instruct-8B",
"speculative_algorithm": "EAGLE",
"speculative_num_steps": 5,
"speculative_eagle_topk": 4,
"speculative_num_draft_tokens": 8,
"speculative_token_map": "thunlp/LLaMA3-Instruct-8B-FR-Spec/freq_32768.pt",
"mem_fraction_static": 0.7,
"cuda_graph_max_bs": 5,
"dtype": "float16",
}
NUM_CONFIGS = 1
THRESHOLDS = {
"batch_avg_accept_len": 1.9,
"accept_len": 2.5,
}
class TestEAGLE3Engine(TestEAGLEEngine):
BASE_CONFIG = {
"model_path": DEFAULT_TARGET_MODEL_EAGLE3,
"speculative_draft_model_path": DEFAULT_DRAFT_MODEL_EAGLE3,
"speculative_algorithm": "EAGLE3",
"speculative_num_steps": 5,
"speculative_eagle_topk": 16,
"speculative_num_draft_tokens": 64,
"mem_fraction_static": 0.7,
"cuda_graph_max_bs": 5,
"dtype": "float16",
}
NUM_CONFIGS = 1
THRESHOLDS = {
"batch_avg_accept_len": 1.75,
"accept_len": 3.1,
}
if __name__ == "__main__":
unittest.main()
@@ -1,388 +0,0 @@
import json
import random
import threading
import time
import unittest
from concurrent.futures import ThreadPoolExecutor
from functools import partial
from types import SimpleNamespace
import numpy as np
import requests
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.abort_timeout_kit import (
AbortAllMixin,
RunningTimeoutTwoWaveMixin,
WaitingTimeoutMixin,
)
from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test
from sglang.test.run_eval import run_eval
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
from sglang.test.test_utils import DEFAULT_TARGET_MODEL_EAGLE, run_logprob_check
register_cuda_ci(est_time=847, stage="base-b", runner_config="1-gpu-large")
class TestEAGLEServerBasic(EagleServerBase):
"""Core tests that run on every server config variant."""
extra_args = ["--chunked-prefill-size", 128, "--max-running-requests", 8]
@classmethod
def setUpClass(cls):
with envs.SGLANG_ENABLE_SPEC_V2.override(False):
super().setUpClass()
# FIXME(lsyin): move the test methods to kits
def test_request_abort(self):
concurrency = 4
threads = [
threading.Thread(target=self.send_request) for _ in range(concurrency)
] + [
threading.Thread(target=self.send_requests_abort)
for _ in range(concurrency)
]
for worker in threads:
worker.start()
for p in threads:
p.join()
def test_gsm8k(self):
requests.get(self.base_url + "/flush_cache")
args = SimpleNamespace(
base_url=self.base_url,
model=self.target_model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=200,
num_threads=128,
)
metrics = run_eval(args)
print(f"{metrics=}")
self.assertGreater(metrics["score"], 0.20)
server_info = requests.get(self.base_url + "/server_info").json()
avg_spec_accept_length = server_info["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
speculative_eagle_topk = server_info["speculative_eagle_topk"]
if speculative_eagle_topk == 1:
self.assertGreater(avg_spec_accept_length, 2.5)
else:
self.assertGreater(avg_spec_accept_length, 3.47)
# Wait a little bit so that the memory check happens.
time.sleep(4)
class TestEAGLEServerAdditional(TestEAGLEServerBasic):
spec_topk = 5
spec_steps = 8
spec_tokens = 64
extra_args = [
"--max-running-requests",
8,
"--cuda-graph-max-bs",
5,
"--attention-backend",
"fa3",
"--page-size",
256,
"--dtype",
"float16",
]
def test_radix_attention(self):
run_radix_attention_test(self.base_url)
self.assertIsNone(self.process.poll())
def test_max_token_one(self):
requests.get(self.base_url + "/flush_cache")
args = SimpleNamespace(
base_url=self.base_url,
model=self.target_model,
eval_name="gsm8k",
api="completion",
max_tokens=1,
num_examples=200,
num_threads=128,
)
metrics = run_eval(args)
self.assertGreater(metrics["output_throughput"], 50)
def test_logprob_start_len(self):
logprob_start_len = 4
new_tokens = 4
prompts = [
"I have a very good idea on",
"Today is a sunndy day and",
]
response = requests.post(
self.base_url + "/generate",
json={
"text": prompts,
"sampling_params": {
"temperature": 0,
"max_new_tokens": new_tokens,
},
"return_logprob": True,
"top_logprobs_num": 5,
"logprob_start_len": logprob_start_len,
},
)
response_json = response.json()
print(json.dumps(response_json, indent=2))
for res in response_json:
self.assertEqual(
res["meta_info"]["prompt_tokens"],
logprob_start_len + len(res["meta_info"]["input_token_logprobs"]),
)
self.assertEqual(res["meta_info"]["completion_tokens"], new_tokens)
self.assertEqual(len(res["meta_info"]["output_token_logprobs"]), new_tokens)
def test_logprob_match(self):
"""Test the output logprobs are close to the input logprobs if we run a prefill again."""
def run_generate(
prompt,
return_logprob=False,
max_new_tokens=512,
logprob_start_len=-1,
temperature=1.0,
):
if isinstance(prompt, str):
prompt_kwargs = {"text": prompt}
else:
prompt_kwargs = {"input_ids": prompt}
response = requests.post(
self.base_url + "/generate",
json={
**prompt_kwargs,
"sampling_params": {
"temperature": temperature,
"max_new_tokens": max_new_tokens,
"ignore_eos": True,
},
"return_logprob": return_logprob,
"return_text_in_logprobs": True,
"logprob_start_len": logprob_start_len,
},
)
return response.json()
prompt = "I have a very good idea on how to"
for temperature in [1.0]:
gen = run_generate(
prompt,
return_logprob=True,
logprob_start_len=0,
temperature=temperature,
)
output_logprobs = np.array(
[x[0] for x in gen["meta_info"]["output_token_logprobs"]]
)
num_prompts_tokens = gen["meta_info"]["prompt_tokens"]
input_tokens = [x[1] for x in gen["meta_info"]["input_token_logprobs"]]
output_tokens = [x[1] for x in gen["meta_info"]["output_token_logprobs"]]
new_prompt = input_tokens + output_tokens
score = run_generate(
new_prompt,
return_logprob=True,
logprob_start_len=0,
max_new_tokens=0,
temperature=temperature,
)
output_logprobs_score = np.array(
[
x[0]
for x in score["meta_info"]["input_token_logprobs"][
num_prompts_tokens:
]
]
)
print(f"{output_logprobs[-10:]=}")
print(f"{output_logprobs_score[-10:]=}")
diff = np.abs(output_logprobs - output_logprobs_score)
max_diff = np.max(diff)
self.assertLess(max_diff, 0.255)
def test_logprob_mixed(self):
args = []
temperature = 0
# input_len, output_len, temperature, logprob_start_len, return_logprob, top_logprobs_num
# Llama 2 context length seems to be only 2k, so we can only test small length.
for input_len in [200, 500, 1000, 2000]:
for output_len in [4, 8]:
for logprob_start_len in [0, 100, 300, 800, 1998]:
for return_logprob in [True, False]:
for top_logprobs_num in [0, 5]:
if logprob_start_len >= input_len:
continue
args.append(
(
input_len,
output_len,
temperature,
logprob_start_len,
return_logprob,
top_logprobs_num,
)
)
random.shuffle(args)
func = partial(run_logprob_check, self)
with ThreadPoolExecutor(8) as executor:
list(executor.map(func, args))
def test_penalty_mixed(self):
args = [
{},
{},
{},
{"frequency_penalty": 2},
{"presence_penalty": 1},
{"min_new_tokens": 16},
{"frequency_penalty": 0.2},
{"presence_penalty": 0.4},
{"min_new_tokens": 8},
{"frequency_penalty": 0.4, "presence_penalty": 0.8},
{"frequency_penalty": 0.4, "min_new_tokens": 12},
{"presence_penalty": 0.8, "min_new_tokens": 12},
{"presence_penalty": -0.3, "frequency_penalty": 1.3, "min_new_tokens": 32},
{"presence_penalty": 0.3, "frequency_penalty": -1.3, "min_new_tokens": 32},
]
random.shuffle(args * 5)
with ThreadPoolExecutor(8) as executor:
list(executor.map(self.run_decode, args))
def test_constrained_decoding(self):
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Give me a json"},
]
response = requests.post(
self.base_url + "/v1/chat/completions",
json={
"model": DEFAULT_TARGET_MODEL_EAGLE,
"messages": messages,
"temperature": 0,
"response_format": {"type": "json_object"},
},
)
self.assertEqual(response.status_code, 200)
res = response.json()
# Validate response structure
self.assertIn("choices", res)
self.assertEqual(len(res["choices"]), 1)
self.assertIn("message", res["choices"][0])
self.assertIn("content", res["choices"][0]["message"])
# Validate JSON content
content_json = res["choices"][0]["message"]["content"]
is_valid_json = True
try:
content = json.loads(content_json)
self.assertIsInstance(content, dict)
except Exception:
print(f"parse JSON failed: {content_json}")
is_valid_json = False
self.assertTrue(is_valid_json)
class TestEAGLERetract(TestEAGLEServerBasic):
extra_args = [
"--chunked-prefill-size=128",
"--max-running-requests=64",
"--max-total-tokens=4500", # Set a smaller KV cache to trigger retract more easily
]
@classmethod
def setUpClass(cls):
# These config helps find a leak.
with envs.SGLANG_TEST_RETRACT.override(True):
super().setUpClass()
class TestEAGLEServerTriton(TestEAGLEServerBasic):
extra_args = ["--attention-backend=triton", "--max-running-requests=8"]
class TestEAGLEServerPageSize(TestEAGLEServerBasic):
spec_steps = 5
spec_topk = 1
spec_tokens = 6
extra_args = [
"--chunked-prefill-size=128",
"--max-running-requests=8",
"--page-size=4",
"--attention-backend=flashinfer",
]
@classmethod
def setUpClass(cls):
# Runtime check only supported for topk=1, and can help to find a leak.
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1):
super().setUpClass()
class TestEAGLEServerPageSizeTopk(TestEAGLEServerBasic):
# default topk=8 and tokens=64
extra_args = [
"--chunked-prefill-size=128",
"--max-running-requests=8",
"--page-size=4",
"--attention-backend=flashinfer",
]
class TestEAGLEAbortAll(AbortAllMixin, EagleServerBase):
abort_all_max_new_tokens = 4000
extra_args = ["--max-running-requests=8"]
class TestEAGLEWaitingTimeout(WaitingTimeoutMixin, EagleServerBase):
extra_args = ["--max-running-requests=1"]
@classmethod
def setUpClass(cls):
with envs.SGLANG_REQ_WAITING_TIMEOUT.override(0.001):
super().setUpClass()
class TestEAGLERunningTimeout(RunningTimeoutTwoWaveMixin, EagleServerBase):
# Regression test for https://github.com/sgl-project/sglang/pull/18760
extra_args = ["--max-running-requests=16"]
@classmethod
def setUpClass(cls):
with envs.SGLANG_REQ_RUNNING_TIMEOUT.override(3):
super().setUpClass()
if __name__ == "__main__":
unittest.main()
@@ -1,281 +0,0 @@
import unittest
from types import SimpleNamespace
import numpy as np
import requests
from sglang.srt.environ import envs
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.matched_stop_kit import MatchedStopMixin
from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_DRAFT_MODEL_EAGLE3,
DEFAULT_TARGET_MODEL_EAGLE3,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=369, stage="base-b", runner_config="1-gpu-small")
class TestEagle3ServerBase(CustomTestCase, MatchedStopMixin):
max_running_requests = 64
attention_backend = "triton"
spec_steps = 5
spec_topk = 1
spec_draft_tokens = 6
page_size = 1
other_launch_args = []
model = DEFAULT_TARGET_MODEL_EAGLE3
draft_model = DEFAULT_DRAFT_MODEL_EAGLE3
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
launch_args = [
"--trust-remote-code",
"--dtype=float16",
"--chunked-prefill-size",
"1024",
"--attention-backend",
cls.attention_backend,
"--speculative-algorithm",
"EAGLE3",
"--speculative-draft-model",
cls.draft_model,
"--speculative-num-steps",
cls.spec_steps,
"--speculative-eagle-topk",
cls.spec_topk,
"--speculative-num-draft-tokens",
cls.spec_draft_tokens,
"--page-size",
str(cls.page_size),
"--mem-fraction-static",
"0.75",
"--max-running-requests",
str(cls.max_running_requests),
"--cuda-graph-bs",
*[str(i) for i in range(1, cls.max_running_requests + 1)],
]
launch_args.extend(cls.other_launch_args)
with (
envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1),
envs.SGLANG_ENABLE_ASYNC_ASSERT.override(True),
envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True),
):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=launch_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_radix_attention(self):
run_radix_attention_test(self.base_url)
assert self.process.poll() is None
def test_gsm8k(self):
args = SimpleNamespace(
base_url=self.base_url,
model=self.model,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=1000,
num_threads=128,
)
metrics = run_eval(args)
print(f"TestEagle3LargeBS -- {metrics=}")
self.assertGreater(metrics["score"], 0.7)
assert self.process.poll() is None
def test_logprob_spec_v2_match(self):
"""Verify spec v2 decode logprobs match prefill scoring logprobs.
Generate tokens with spec v2, then score the same sequence via
prefill-only (no speculation). The two sets of logprobs should be
close, validating that spec v2 computes logprobs correctly.
Runs two rounds with different prompts to catch state-dependent bugs.
"""
top_k = 5
probe_token_ids = [1, 2, 10, 100, 1000]
prompts = [
"The capital of France is",
"Explain quantum computing in simple terms:",
]
for round_idx, prompt in enumerate(prompts):
with self.subTest(round=round_idx, prompt=prompt):
gen_res = requests.post(
self.base_url + "/generate",
json={
"text": prompt,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 32,
"ignore_eos": True,
},
"return_logprob": True,
"top_logprobs_num": top_k,
"token_ids_logprob": probe_token_ids,
"logprob_start_len": 0,
},
).json()
decode_logprobs = gen_res["meta_info"]["output_token_logprobs"]
decode_top_logprobs = gen_res["meta_info"]["output_top_logprobs"]
decode_tid_logprobs = gen_res["meta_info"]["output_token_ids_logprobs"]
input_token_ids = [
t[1] for t in gen_res["meta_info"]["input_token_logprobs"]
]
output_token_ids = [t[1] for t in decode_logprobs]
num_prompt_tokens = gen_res["meta_info"]["prompt_tokens"]
score_res = requests.post(
self.base_url + "/generate",
json={
"input_ids": input_token_ids + output_token_ids,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 0,
},
"return_logprob": True,
"top_logprobs_num": top_k,
"token_ids_logprob": probe_token_ids,
"logprob_start_len": 0,
},
).json()
score_logprobs = score_res["meta_info"]["input_token_logprobs"][
num_prompt_tokens:
]
score_top_logprobs = score_res["meta_info"]["input_top_logprobs"][
num_prompt_tokens:
]
score_tid_logprobs = score_res["meta_info"]["input_token_ids_logprobs"][
num_prompt_tokens:
]
self.assertEqual(len(decode_logprobs), len(score_logprobs))
# Check per-token logprobs
decode_vals = np.array([t[0] for t in decode_logprobs])
score_vals = np.array([t[0] for t in score_logprobs])
max_diff = np.max(np.abs(decode_vals - score_vals))
print(
f"[round {round_idx}] prompt={prompt!r} "
f"logprob max_diff={max_diff:.6f}"
)
print(f"[round {round_idx}] decode_vals[-5:]={decode_vals[-5:]}")
print(f"[round {round_idx}] score_vals[-5:]={score_vals[-5:]}")
self.assertLess(max_diff, 0.255)
# Check top-k logprobs
for pos in range(len(decode_logprobs)):
dec_top = {t[1]: t[0] for t in decode_top_logprobs[pos]}
scr_top = {t[1]: t[0] for t in score_top_logprobs[pos]}
common_ids = set(dec_top.keys()) & set(scr_top.keys())
self.assertGreater(len(common_ids), 0)
for tid in common_ids:
self.assertAlmostEqual(dec_top[tid], scr_top[tid], delta=0.255)
# Check token_ids_logprob
self.assertEqual(len(decode_tid_logprobs), len(score_tid_logprobs))
for pos in range(len(decode_tid_logprobs)):
dec_tid = {t[1]: t[0] for t in decode_tid_logprobs[pos]}
scr_tid = {t[1]: t[0] for t in score_tid_logprobs[pos]}
self.assertEqual(set(dec_tid.keys()), set(scr_tid.keys()))
for tid in dec_tid:
self.assertAlmostEqual(dec_tid[tid], scr_tid[tid], delta=0.255)
def test_token_ids_logprob_ragged(self):
"""Regression: get_token_ids_logprobs_raw crashes on ragged token_ids_logprob lists.
Sends concurrent requests with different-length token_ids_logprob lists
so they land in the same batch. torch.tensor() on ragged input will crash.
"""
import concurrent.futures
def send(probe_ids):
return requests.post(
self.base_url + "/generate",
json={
"text": "Hello world",
"sampling_params": {
"temperature": 0,
"max_new_tokens": 8,
},
"return_logprob": True,
"top_logprobs_num": 3,
"token_ids_logprob": probe_ids,
},
).json()
ragged_probes = [
[1, 2],
[3, 4, 5],
[6],
[10, 20, 30, 40],
[1, 2],
[3, 4, 5],
[6],
[10, 20, 30, 40],
]
with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool:
futs = [pool.submit(send, ids) for ids in ragged_probes]
for f in concurrent.futures.as_completed(futs):
res = f.result()
self.assertIn("text", res, f"Server error: {res}")
def test_penalty(self):
"""Verify spec v2 handles penalty parameters without crashing."""
import concurrent.futures
args = [
{"max_new_tokens": 32},
{"max_new_tokens": 16, "frequency_penalty": 2},
{"max_new_tokens": 48, "presence_penalty": 1},
{"max_new_tokens": 8, "frequency_penalty": 0.4, "presence_penalty": 0.8},
{"max_new_tokens": 64, "frequency_penalty": -0.5, "presence_penalty": 0.3},
{"max_new_tokens": 24, "min_new_tokens": 8, "frequency_penalty": 0.4},
{"max_new_tokens": 32, "repetition_penalty": 1.5},
]
def run_decode(sampling_params):
response = requests.post(
self.base_url + "/generate",
json={
"text": "The capital of France is",
"sampling_params": sampling_params,
},
)
self.assertEqual(response.status_code, 200)
res = response.json()
self.assertIn("text", res, f"Server error: {res}")
self.assertIsInstance(
res["text"],
str,
f"Expected 'text' to be str, got {type(res['text']).__name__}: {res}",
)
with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool:
list(pool.map(run_decode, args * 3))
assert self.process.poll() is None
class TestEagle3ServerPage(TestEagle3ServerBase):
other_launch_args = ["--page-size", "64"]
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,51 @@
"""EAGLE3 spec-decoding core: overlap (spec v2) x no-overlap (spec v1) matrix,
same standard config (topk=1, page_size=1), only ``disable_overlap`` differs.
flashinfer is pinned (the 5090 default) so a default-selection change can't
silently alter what this exercises.
"""
import unittest
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.matched_stop_kit import MatchedStopMixin
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecCorrectnessKit,
SpecFeatureKit,
SpecLogprobKit,
SpecPenaltyKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
register_cuda_ci(est_time=480, stage="base-b", runner_config="1-gpu-small")
_KITS = (
SpecCorrectnessKit,
SpecAccuracyKit,
SpecLogprobKit,
SpecPenaltyKit,
SpecFeatureKit,
MatchedStopMixin,
)
class _Core(Eagle3Base):
# Busy-time pool accounting check (topk=1 only).
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
class TestEagle3Overlap(_Core, *_KITS):
"""Spec v2 (overlap scheduler on)."""
disable_overlap = False
class TestEagle3NoOverlap(_Core, *_KITS):
"""Spec v1 (overlap scheduler off)."""
disable_overlap = True
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,49 @@
"""fa3 attention backend -- Hopper-only (FlashAttention-3 is sm_90).
fa3 is the real H200 default for MHA spec at topk=1, so this also covers the
"what an H200 user actually runs" path. Requires the large (Hopper) runner.
"""
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecCorrectnessKit,
SpecFeatureKit,
SpecLogprobKit,
SpecPenaltyKit,
SpecPerfKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-large")
class TestEagle3Fa3(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprobKit):
"""EAGLE3 spec v2 topk=1 on fa3 (the H200 default backend)."""
attention_backend = "fa3"
disable_overlap = False
class TestEagleLlama2Fa3Page256(
EagleLlama2Base,
SpecAccuracyKit,
SpecLogprobKit,
SpecPenaltyKit,
SpecPerfKit,
SpecFeatureKit,
):
"""EAGLE/Llama-2 topk=5 tree on fa3 + page_size=256 (spec v1)."""
spec_topk = 5
spec_steps = 8
attention_backend = "fa3"
page_size = 256
chunked_prefill_size = 4096 # must be divisible by page_size (256)
cuda_graph_max_bs = 5
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,47 @@
"""page_size > 1 variants (flashinfer).
EAGLE3 page64 (spec v2) + EAGLE/Llama-2 page4 (topk1 and topk8, spec v1).
Runs on the cheap (5090) runner.
"""
import unittest
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecFeatureKit,
SpecLogprobKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
register_cuda_ci(est_time=540, stage="base-b", runner_config="1-gpu-small")
class TestEagle3Page64(Eagle3Base, SpecAccuracyKit, SpecLogprobKit, SpecFeatureKit):
"""EAGLE3 spec v2, page_size=64 (flashinfer): + logprob losslessness."""
page_size = 64
disable_overlap = False
class TestEagleLlama2Page4Topk1(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
"""Llama-2 topk=1 + page_size=4; busy-time pool check (topk=1 only)."""
spec_topk = 1
spec_tokens = 6
page_size = 4
env_overrides = (
(envs.SGLANG_ENABLE_SPEC_V2, False),
(envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),
)
class TestEagleLlama2Page4Topk8(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
"""Llama-2 topk>1 tree + page_size=4 (spec v1)."""
page_size = 4
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,28 @@
"""Lossless output parity: spec-decode greedy output == a non-spec reference.
The reference is a separate non-spec server, launched and torn down BEFORE the
spec server (sequential -- one model resident at a time; see SpecParityKit).
"""
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.spec_server_kits import SpecParityKit
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
register_cuda_ci(est_time=360, stage="base-b", runner_config="1-gpu-large")
class TestEagle3Parity(SpecParityKit, Eagle3Base):
"""EAGLE3 spec v2 (flashinfer) greedy output == non-spec reference.
SpecParityKit is first so its setUpClass runs the reference server (and tears
it down) before the fixture launches the spec server -- sequential, one model
at a time.
"""
disable_overlap = False
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,65 @@
"""Perf + stress: throughput, retract-under-pressure, abort storms, timeouts.
These need memory headroom / measure load behavior, so they run on the large
(Hopper) runner.
"""
import unittest
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.abort_timeout_kit import (
AbortAllMixin,
RunningTimeoutTwoWaveMixin,
WaitingTimeoutMixin,
)
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecFeatureKit,
SpecPerfKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-large")
class TestEagle3Perf(Eagle3Base, SpecPerfKit):
"""Decode throughput (max_new_tokens=1) on EAGLE3 spec v2."""
disable_overlap = False
class TestEagleLlama2Retract(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
"""Retract under a small KV budget; must not leak."""
max_running_requests = 64
extra_args = ("--max-total-tokens", 4500) # small KV to trigger retract
env_overrides = (
(envs.SGLANG_ENABLE_SPEC_V2, False),
(envs.SGLANG_TEST_RETRACT, True),
)
class TestEagleLlama2AbortAll(EagleLlama2Base, AbortAllMixin):
abort_all_max_new_tokens = 4000
class TestEagleLlama2WaitingTimeout(EagleLlama2Base, WaitingTimeoutMixin):
max_running_requests = 1
env_overrides = (
(envs.SGLANG_ENABLE_SPEC_V2, False),
(envs.SGLANG_REQ_WAITING_TIMEOUT, 0.001),
)
class TestEagleLlama2RunningTimeout(EagleLlama2Base, RunningTimeoutTwoWaveMixin):
# Regression: https://github.com/sgl-project/sglang/pull/18760
max_running_requests = 16
env_overrides = (
(envs.SGLANG_ENABLE_SPEC_V2, False),
(envs.SGLANG_REQ_RUNNING_TIMEOUT, 3),
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,67 @@
"""topk > 1 tree drafting (EAGLE3 topk16 + EAGLE/Llama-2 topk8).
topk > 1 always routes to spec v1; flashinfer is pinned (topk > 1 can't use fa3).
Runs on the cheap (5090) runner -- functional sanity only, no perf/stress.
"""
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecCorrectnessKit,
SpecFeatureKit,
SpecLogprobKit,
SpecPenaltyKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
register_cuda_ci(est_time=840, stage="base-b", runner_config="1-gpu-small")
class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprobKit):
"""EAGLE3 topk=16 tree (spec v1): correctness + gsm8k + logprob losslessness."""
spec_topk = 16
spec_tokens = 64
disable_overlap = True # topk>1 -> spec v1
cuda_graph_max_bs = 5
acc_length_thres = 3.1
batch_accept_len_thres = 1.75
gsm8k_accept_len_thres = 2.4 # EAGLE3 topk16 gsm8k accept ~2.48
class TestEagleLlama2Suite(
EagleLlama2Base,
SpecCorrectnessKit,
SpecAccuracyKit,
SpecLogprobKit,
SpecPenaltyKit,
SpecFeatureKit,
):
"""EAGLE/Llama-2 topk=8 full coverage (kits listed in bases)."""
class TestEagleLlama2Chunked4(EagleLlama2Base, SpecCorrectnessKit):
"""Correctness under tiny chunked prefill."""
chunked_prefill_size = 4
class TestEagleLlama3TokenMap(EagleLlama2Base, SpecAccuracyKit):
"""EAGLE on Llama-3-8B with a FR-Spec token map (topk=4)."""
model = "meta-llama/Meta-Llama-3-8B-Instruct"
draft_model = "lmsys/sglang-EAGLE-LLaMA3-Instruct-8B"
spec_topk = 4
spec_tokens = 8
cuda_graph_max_bs = 5
gsm8k_accept_len_thres = 2.5 # FR-Spec token map lowers accept (~2.57)
extra_args = (
"--speculative-token-map",
"thunlp/LLaMA3-Instruct-8B-FR-Spec/freq_32768.pt",
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,47 @@
"""triton attention backend (EAGLE3 spec v2 + EAGLE/Llama-2 spec v1).
triton runs everywhere, so this stays on the cheap (5090) runner.
"""
import unittest
from sglang.srt.environ import envs
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.matched_stop_kit import MatchedStopMixin
from sglang.test.kits.spec_server_kits import (
SpecAccuracyKit,
SpecFeatureKit,
SpecLogprobKit,
SpecPenaltyKit,
)
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
register_cuda_ci(est_time=480, stage="base-b", runner_config="1-gpu-small")
class TestEagle3Triton(
Eagle3Base,
MatchedStopMixin,
SpecAccuracyKit,
SpecLogprobKit,
SpecPenaltyKit,
SpecFeatureKit,
):
"""EAGLE3 spec v2 on triton (kits listed in bases)."""
attention_backend = "triton"
max_running_requests = 64
cuda_graph_max_bs = 64
gsm8k_num_examples = 1000
gsm8k_check_accept_len = False
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
class TestEagleLlama2Triton(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit):
"""EAGLE/Llama-2 topk=8 on triton (spec v1)."""
attention_backend = "triton"
if __name__ == "__main__":
unittest.main()