split test_dsa_models_mtp into 4 files (#25318)

This commit is contained in:
Liangsheng Yin
2026-05-15 14:16:57 -07:00
committed by GitHub
parent f9caf43095
commit 494ee71189
8 changed files with 216 additions and 371 deletions
@@ -36,6 +36,7 @@ class GSM8KMixin:
gsm8k_accept_length_thres: Optional[float] = None
gsm8k_num_questions: int = 200
gsm8k_num_threads: int = 128
gsm8k_num_shots: int = 5
def test_gsm8k(self):
assert (
@@ -52,6 +53,7 @@ class GSM8KMixin:
max_tokens=512,
num_examples=self.gsm8k_num_questions,
num_threads=self.gsm8k_num_threads,
num_shots=self.gsm8k_num_shots,
)
metrics = run_eval(args)
print(f"{metrics=}")
+11 -3
View File
@@ -1,3 +1,5 @@
import requests
from sglang.test.send_one import BenchArgs, send_one_prompt
from sglang.test.test_utils import is_in_ci, write_github_step_summary
@@ -5,12 +7,18 @@ from sglang.test.test_utils import is_in_ci, write_github_step_summary
class SpecDecodingMixin:
bs_1_speed_thres: float
accept_length_thres: float
bs_1_speed_attempts: int = 3
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
acc_length, speed = send_one_prompt(args)
print(f"{acc_length=:.2f} {speed=:.2f}")
acc_length, speed = 0.0, 0.0
for attempt in range(1, self.bs_1_speed_attempts + 1):
requests.get(self.base_url + "/flush_cache")
acc_length, speed = send_one_prompt(args)
print(f"attempt {attempt}: {acc_length=:.2f} {speed=:.2f}")
if acc_length > self.accept_length_thres and speed > self.bs_1_speed_thres:
break
requests.get(self.base_url + "/flush_cache")
if is_in_ci():
write_github_step_summary(
@@ -0,0 +1,93 @@
"""DSA model + MTP (EAGLE) speculative-decoding server fixture.
Variants combine `DsaMtpServerBase` (server lifecycle) with
`DsaMtpEvalConfigDefaults` (shared eval thresholds/params),
`GSM8KMixin` and `SpecDecodingMixin`, then set `model` and per-variant
overrides (`enable_dp_attention`, `mem_fraction_static`, `bs_1_speed_thres`).
Example:
class TestDsv32DP(
DsaMtpServerBase,
DsaMtpEvalConfigDefaults,
GSM8KMixin,
SpecDecodingMixin,
):
model = "deepseek-ai/DeepSeek-V3.2"
enable_dp_attention = True
bs_1_speed_thres = 90
The base itself is NOT a runnable test (no `test_*` methods until a subclass
mixes in the kits), so unittest discovery picks it up as empty.
"""
from sglang.srt.utils import kill_process_tree
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
class DsaMtpEvalConfigDefaults:
"""Eval thresholds & params shared across DSA-MTP regression variants."""
# GSM8KMixin defaults.
gsm8k_accuracy_thres = 0.94
gsm8k_accept_length_thres = 2.7
gsm8k_num_questions = 500
gsm8k_num_threads = 500
gsm8k_num_shots = 20
# SpecDecodingMixin default; per-variant subclasses set `bs_1_speed_thres`.
accept_length_thres = 2.7
class DsaMtpServerBase(CustomTestCase):
base_url = DEFAULT_URL_FOR_TEST
# Subclasses must set `model`; the others have sensible defaults.
model: str = ""
mem_fraction_static: float = 0.7
enable_dp_attention: bool = False
# EAGLE MTP config (fixed across DSA-MTP variants).
speculative_algorithm: str = "EAGLE"
speculative_num_steps: int = 3
speculative_eagle_topk: int = 1
speculative_num_draft_tokens: int = 4
@classmethod
def get_server_args(cls):
assert cls.model, f"{cls.__name__} must set `model`"
args = ["--trust-remote-code", "--tp", "8"]
if cls.enable_dp_attention:
args += ["--dp", "8", "--enable-dp-attention"]
args += [
"--speculative-algorithm",
cls.speculative_algorithm,
"--speculative-num-steps",
str(cls.speculative_num_steps),
"--speculative-eagle-topk",
str(cls.speculative_eagle_topk),
"--speculative-num-draft-tokens",
str(cls.speculative_num_draft_tokens),
"--mem-frac",
str(cls.mem_fraction_static),
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
]
return args
@classmethod
def setUpClass(cls):
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=cls.get_server_args(),
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
@@ -1,368 +0,0 @@
import unittest
from types import SimpleNamespace
import requests
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.send_one import BenchArgs, send_one_prompt
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
is_in_ci,
popen_launch_server,
write_github_step_summary,
)
register_cuda_ci(
est_time=1030,
stage="stage-c",
runner_config="8-gpu-h200",
)
FULL_DEEPSEEK_V32_MODEL_PATH = "deepseek-ai/DeepSeek-V3.2"
GLM5_MODEL_PATH = "zai-org/GLM-5-FP8"
class TestDeepseekV32DPMTP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = FULL_DEEPSEEK_V32_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--tp",
"8",
"--dp",
"8",
"--enable-dp-attention",
"--speculative-algorithm",
"EAGLE",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-frac",
"0.7",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_a_gsm8k(
self,
): # Append an "a" to make this test run first (alphabetically) to warm up the server
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=500,
num_threads=500,
num_shots=20,
)
metrics = run_eval(args)
print(f"{metrics=}")
server_info = requests.get(self.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
if is_in_ci():
write_github_step_summary(
f"### test_gsm8k (deepseek-v32 mtp)\n"
f'{metrics["score"]=:.3f}\n'
f"{avg_spec_accept_length=:.2f}\n"
)
self.assertGreater(metrics["score"], 0.94)
self.assertGreater(avg_spec_accept_length, 2.7)
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
acc_length, speed = send_one_prompt(args)
print(f"{acc_length=:.2f} {speed=:.2f}")
if is_in_ci():
write_github_step_summary(
f"### test_bs_1_speed (deepseek-v32 mtp)\n"
f"{acc_length=:.2f}\n"
f"{speed=:.2f} token/s\n"
)
self.assertGreater(acc_length, 2.7)
self.assertGreater(speed, 90)
class TestDeepseekV32TPMTP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = FULL_DEEPSEEK_V32_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--tp",
"8",
"--speculative-algorithm",
"EAGLE",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-frac",
"0.7",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_a_gsm8k(
self,
): # Append an "a" to make this test run first (alphabetically) to warm up the server
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=500,
num_threads=500,
num_shots=20,
)
metrics = run_eval(args)
print(f"{metrics=}")
server_info = requests.get(self.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
if is_in_ci():
write_github_step_summary(
f"### test_gsm8k (deepseek-v32 mtp)\n"
f'{metrics["score"]=:.3f}\n'
f"{avg_spec_accept_length=:.2f}\n"
)
self.assertGreater(metrics["score"], 0.94)
self.assertGreater(avg_spec_accept_length, 2.7)
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
acc_length, speed = send_one_prompt(args)
print(f"{acc_length=:.2f} {speed=:.2f}")
if is_in_ci():
write_github_step_summary(
f"### test_bs_1_speed (deepseek-v32 mtp)\n"
f"{acc_length=:.2f}\n"
f"{speed=:.2f} token/s\n"
)
self.assertGreater(acc_length, 2.7)
self.assertGreater(speed, 180)
class TestGLM5DPMTP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GLM5_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--tp",
"8",
"--dp",
"8",
"--enable-dp-attention",
"--speculative-algorithm",
"EAGLE",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-frac",
"0.8",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_a_gsm8k(
self,
): # Append an "a" to make this test run first (alphabetically) to warm up the server
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=500,
num_threads=500,
num_shots=20,
)
metrics = run_eval(args)
print(f"{metrics=}")
server_info = requests.get(self.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
if is_in_ci():
write_github_step_summary(
f"### test_gsm8k (glm-5 mtp)\n"
f'{metrics["score"]=:.3f}\n'
f"{avg_spec_accept_length=:.2f}\n"
)
self.assertGreater(metrics["score"], 0.94)
self.assertGreater(avg_spec_accept_length, 2.7)
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
acc_length, speed = send_one_prompt(args)
print(f"{acc_length=:.2f} {speed=:.2f}")
if is_in_ci():
write_github_step_summary(
f"### test_bs_1_speed (glm-5 mtp)\n"
f"{acc_length=:.2f}\n"
f"{speed=:.2f} token/s\n"
)
self.assertGreater(acc_length, 2.7)
self.assertGreater(speed, 70)
class TestGLM5TPMTP(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = GLM5_MODEL_PATH
cls.base_url = DEFAULT_URL_FOR_TEST
other_args = [
"--trust-remote-code",
"--tp",
"8",
"--speculative-algorithm",
"EAGLE",
"--speculative-num-steps",
"3",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"4",
"--mem-frac",
"0.8",
"--model-loader-extra-config",
'{"enable_multithread_load": true, "num_threads": 64}',
]
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=other_args,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_a_gsm8k(
self,
): # Append an "a" to make this test run first (alphabetically) to warm up the server
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=500,
num_threads=500,
num_shots=20,
)
metrics = run_eval(args)
print(f"{metrics=}")
server_info = requests.get(self.base_url + "/server_info")
avg_spec_accept_length = server_info.json()["internal_states"][0][
"avg_spec_accept_length"
]
print(f"{avg_spec_accept_length=}")
if is_in_ci():
write_github_step_summary(
f"### test_gsm8k (glm-5 mtp)\n"
f'{metrics["score"]=:.3f}\n'
f"{avg_spec_accept_length=:.2f}\n"
)
self.assertGreater(metrics["score"], 0.94)
self.assertGreater(avg_spec_accept_length, 2.7)
def test_bs_1_speed(self):
args = BenchArgs(port=int(self.base_url.split(":")[-1]), max_new_tokens=2048)
acc_length, speed = send_one_prompt(args)
print(f"{acc_length=:.2f} {speed=:.2f}")
if is_in_ci():
write_github_step_summary(
f"### test_bs_1_speed (glm-5 mtp)\n"
f"{acc_length=:.2f}\n"
f"{speed=:.2f} token/s\n"
)
self.assertGreater(acc_length, 2.7)
self.assertGreater(speed, 150)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,28 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin
from sglang.test.server_fixtures.dsa_mtp_fixture import (
DsaMtpEvalConfigDefaults,
DsaMtpServerBase,
)
register_cuda_ci(
est_time=600,
stage="extra-b",
runner_config="8-gpu-h200",
)
class TestDeepseekV32DPMTP(
DsaMtpServerBase, DsaMtpEvalConfigDefaults, GSM8KMixin, SpecDecodingMixin
):
model = "deepseek-ai/DeepSeek-V3.2"
mem_fraction_static = 0.7
enable_dp_attention = True
bs_1_speed_thres = 90
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,27 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin
from sglang.test.server_fixtures.dsa_mtp_fixture import (
DsaMtpEvalConfigDefaults,
DsaMtpServerBase,
)
register_cuda_ci(
est_time=400,
stage="extra-b",
runner_config="8-gpu-h200",
)
class TestDeepseekV32TPMTP(
DsaMtpServerBase, DsaMtpEvalConfigDefaults, GSM8KMixin, SpecDecodingMixin
):
model = "deepseek-ai/DeepSeek-V3.2"
mem_fraction_static = 0.7
bs_1_speed_thres = 180
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,28 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin
from sglang.test.server_fixtures.dsa_mtp_fixture import (
DsaMtpEvalConfigDefaults,
DsaMtpServerBase,
)
register_cuda_ci(
est_time=400,
stage="stage-c",
runner_config="8-gpu-h200",
)
class TestGLM5DPMTP(
DsaMtpServerBase, DsaMtpEvalConfigDefaults, GSM8KMixin, SpecDecodingMixin
):
model = "zai-org/GLM-5-FP8"
mem_fraction_static = 0.8
enable_dp_attention = True
bs_1_speed_thres = 70
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,27 @@
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin
from sglang.test.server_fixtures.dsa_mtp_fixture import (
DsaMtpEvalConfigDefaults,
DsaMtpServerBase,
)
register_cuda_ci(
est_time=400,
stage="stage-c",
runner_config="8-gpu-h200",
)
class TestGLM5TPMTP(
DsaMtpServerBase, DsaMtpEvalConfigDefaults, GSM8KMixin, SpecDecodingMixin
):
model = "zai-org/GLM-5-FP8"
mem_fraction_static = 0.8
bs_1_speed_thres = 150
if __name__ == "__main__":
unittest.main()