Add EAGLE3 test with MMLU dataset. (#15945)
This commit is contained in:
@@ -115,8 +115,8 @@ class BaseFlashAttentionTest(CustomTestCase):
|
|||||||
self.assertGreater(metrics[metric_key], self.accuracy_threshold)
|
self.assertGreater(metrics[metric_key], self.accuracy_threshold)
|
||||||
|
|
||||||
if self.speculative_decode:
|
if self.speculative_decode:
|
||||||
server_info = requests.get(self.base_url + "/get_server_info")
|
server_info = requests.get(self.base_url + "/server_info").json()
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
@@ -125,9 +125,8 @@ class TestFlashMLAMTP(CustomTestCase):
|
|||||||
|
|
||||||
self.assertGreater(metrics["accuracy"], 0.60)
|
self.assertGreater(metrics["accuracy"], 0.60)
|
||||||
|
|
||||||
server_info = requests.get(self.base_url + "/get_server_info")
|
server_info = requests.get(self.base_url + "/server_info").json()
|
||||||
print(f"{server_info=}")
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
@@ -115,9 +115,8 @@ class TestFlashinferMLAMTP(CustomTestCase):
|
|||||||
|
|
||||||
self.assertGreater(metrics["accuracy"], 0.60)
|
self.assertGreater(metrics["accuracy"], 0.60)
|
||||||
|
|
||||||
server_info = requests.get(self.base_url + "/get_server_info")
|
server_info = requests.get(self.base_url + "/get_server_info").json()
|
||||||
print(f"{server_info=}")
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
@@ -79,8 +79,8 @@ class TestDeepseekV3FP4MTP(CustomTestCase):
|
|||||||
metrics = run_eval_few_shot_gsm8k(args)
|
metrics = run_eval_few_shot_gsm8k(args)
|
||||||
print(f"{metrics=}")
|
print(f"{metrics=}")
|
||||||
|
|
||||||
server_info = requests.get(self.base_url + "/get_server_info")
|
server_info = requests.get(self.base_url + "/server_info").json()
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.run_eval import run_eval
|
||||||
|
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
|
DEFAULT_TARGET_MODEL_EAGLE3,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=50, suite="stage-b-test-small-1-gpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestEagle3Basic(EagleServerBase):
|
||||||
|
target_model = DEFAULT_TARGET_MODEL_EAGLE3
|
||||||
|
draft_model = DEFAULT_DRAFT_MODEL_EAGLE3
|
||||||
|
|
||||||
|
spec_algo = "EAGLE3"
|
||||||
|
spec_steps = 2
|
||||||
|
spec_topk = 1
|
||||||
|
spec_tokens = 3
|
||||||
|
extra_args = ["--dtype=float16", "--chunked-prefill-size", 1024]
|
||||||
|
|
||||||
|
def test_mmlu(self):
|
||||||
|
"""Override to add EAGLE-specific assertions"""
|
||||||
|
args = SimpleNamespace(
|
||||||
|
base_url=self.base_url,
|
||||||
|
model=self.target_model,
|
||||||
|
eval_name="mmlu",
|
||||||
|
num_examples=64,
|
||||||
|
num_threads=32,
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = run_eval(args)
|
||||||
|
self.assertGreaterEqual(metrics["score"], 0.72)
|
||||||
|
|
||||||
|
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=}")
|
||||||
|
self.assertGreater(avg_spec_accept_length, 2.26)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
|
||||||
|
unittest.main()
|
||||||
@@ -75,7 +75,7 @@ class TestEAGLEServerBasic(EagleServerBase):
|
|||||||
print(f"{metrics=}")
|
print(f"{metrics=}")
|
||||||
self.assertGreater(metrics["accuracy"], 0.20)
|
self.assertGreater(metrics["accuracy"], 0.20)
|
||||||
|
|
||||||
server_info = requests.get(self.base_url + "/get_server_info").json()
|
server_info = requests.get(self.base_url + "/server_info").json()
|
||||||
avg_spec_accept_length = server_info["internal_states"][0][
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -34,10 +34,8 @@ def test_gsm8k(base_url: str):
|
|||||||
port=int(base_url.split(":")[-1]),
|
port=int(base_url.split(":")[-1]),
|
||||||
)
|
)
|
||||||
metrics = run_eval_few_shot_gsm8k(args)
|
metrics = run_eval_few_shot_gsm8k(args)
|
||||||
server_info = requests.get(base_url + "/get_server_info")
|
server_info = requests.get(base_url + "/server_info").json()
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
avg_spec_accept_length = server_info["internal_states"][0]["avg_spec_accept_length"]
|
||||||
"avg_spec_accept_length"
|
|
||||||
]
|
|
||||||
|
|
||||||
print(f"{metrics=}")
|
print(f"{metrics=}")
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
@@ -152,9 +152,8 @@ class TestHiCacheEagle(HiCacheBaseServer, HiCacheEvalMixin):
|
|||||||
self.assertGreaterEqual(metrics["score"], self.expected_mmlu_score)
|
self.assertGreaterEqual(metrics["score"], self.expected_mmlu_score)
|
||||||
|
|
||||||
# EAGLE-specific check
|
# EAGLE-specific check
|
||||||
server_info = requests.get(self.base_url + "/get_server_info")
|
server_info = requests.get(self.base_url + "/get_server_info").json()
|
||||||
print(f"{server_info=}")
|
avg_spec_accept_length = server_info["internal_states"][0][
|
||||||
avg_spec_accept_length = server_info.json()["internal_states"][0][
|
|
||||||
"avg_spec_accept_length"
|
"avg_spec_accept_length"
|
||||||
]
|
]
|
||||||
print(f"{avg_spec_accept_length=}")
|
print(f"{avg_spec_accept_length=}")
|
||||||
|
|||||||
Reference in New Issue
Block a user