Drop the mmlu case from the unified radix cache kit (#34667)
This commit is contained in:
@@ -16,7 +16,7 @@ def _random_suffixes(n, length, seed):
|
||||
|
||||
|
||||
class UnifiedRadixTreeTestMixin:
|
||||
"""Mixin: gsm8k, mmlu and multi-turn KL tests with multi-branch interleaving."""
|
||||
"""Mixin: gsm8k and multi-turn KL tests with multi-branch interleaving."""
|
||||
|
||||
kl_threshold: float = 0.003
|
||||
max_new_tokens: int = 512
|
||||
@@ -30,7 +30,6 @@ class UnifiedRadixTreeTestMixin:
|
||||
decode_hit_inter_batch_delay_s: float = 0
|
||||
|
||||
gsm8k_threshold: float = 0.93
|
||||
mmlu_threshold: float = 0.8
|
||||
num_gsm8k_questions: int = 200
|
||||
|
||||
def test_gsm8k(self):
|
||||
@@ -54,24 +53,6 @@ class UnifiedRadixTreeTestMixin:
|
||||
)
|
||||
self.assertGreaterEqual(metrics["accuracy"], self.gsm8k_threshold)
|
||||
|
||||
def test_mmlu(self):
|
||||
"""Simple-evals MMLU multi-task accuracy."""
|
||||
from sglang.test.run_eval import run_eval as run_simple_eval
|
||||
|
||||
args = SimpleNamespace(
|
||||
base_url=self.base_url,
|
||||
model=self.model,
|
||||
eval_name="mmlu",
|
||||
num_examples=64,
|
||||
num_threads=32,
|
||||
)
|
||||
metrics = run_simple_eval(args)
|
||||
print(
|
||||
f"[{self.__class__.__name__}] MMLU score: {metrics['score']:.3f} "
|
||||
f"(threshold: {self.mmlu_threshold})"
|
||||
)
|
||||
self.assertGreaterEqual(metrics["score"], self.mmlu_threshold)
|
||||
|
||||
def test_multiturn_logprobs_match(self):
|
||||
"""Helper 1: 3-turn, no explicit cache seeding."""
|
||||
ids = self.input_ids[:4]
|
||||
|
||||
-1
@@ -36,7 +36,6 @@ class TestUnifiedQwen3HiCachePP(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||
kl_threshold = 0.005
|
||||
gsm8k_threshold = 0.7
|
||||
num_gsm8k_questions = 50
|
||||
mmlu_threshold = 0.7
|
||||
decode_cache_assert = staticmethod(_assert_pp_decode_cached_tokens)
|
||||
|
||||
def test_gsm8k(self):
|
||||
|
||||
@@ -24,7 +24,6 @@ class TestUnifiedQwen3HiCacheCP(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||
max_running_requests = 32
|
||||
kl_threshold = 0.005
|
||||
gsm8k_threshold = 0.7
|
||||
mmlu_threshold = 0.7
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
|
||||
@@ -42,7 +42,6 @@ class TestUnifiedKimiLinearDcpHiCache(UnifiedRadixTreeTestMixin, CustomTestCase)
|
||||
|
||||
kl_threshold = 0.01
|
||||
gsm8k_threshold = 0.85
|
||||
mmlu_threshold = 0.4
|
||||
prefill_cache_assert = staticmethod(
|
||||
make_mamba_prefill_assert(chunk_size=WIDENED_PAGE)
|
||||
)
|
||||
|
||||
+3
-3
@@ -25,9 +25,9 @@ Measured avg_kl_div:
|
||||
extra coverage rather than as the guard for one fix. CI runs `1-gpu-large`, which
|
||||
is SM90.
|
||||
|
||||
These classes do not use UnifiedRadixTreeTestMixin: it bundles gsm8k and mmlu,
|
||||
which an undertrained checkpoint cannot gate on, and each class here runs the
|
||||
harness its regression was actually reproduced with.
|
||||
These classes do not use UnifiedRadixTreeTestMixin: it bundles a gsm8k case an
|
||||
undertrained checkpoint cannot gate on, and each class here runs the harness its
|
||||
regression was actually reproduced with.
|
||||
|
||||
The imported `test_`-prefixed helpers are aliased so pytest does not collect them
|
||||
as tests.
|
||||
|
||||
@@ -8,7 +8,6 @@ from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
is_in_ci,
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
@@ -22,11 +21,6 @@ class TestUnifiedSWARadixCache(UnifiedRadixTreeTestMixin, CustomTestCase):
|
||||
|
||||
kl_threshold = 0.03
|
||||
gsm8k_threshold = 0.7
|
||||
mmlu_threshold = 0.7
|
||||
|
||||
@unittest.skipIf(is_in_ci(), "SWA model mmlu eval not stable enough")
|
||||
def test_mmlu(self):
|
||||
super().test_mmlu()
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
|
||||
Reference in New Issue
Block a user