Refactor EAGLE infer tests: shared fixture + kits + overlap matrix (#26871)
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user