From 687cfe91982904ee789577c82ddae19c51e041d8 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Thu, 4 Jun 2026 16:46:20 -0400 Subject: [PATCH] Enable runtime busy memory check for speculation topk>1 (#27228) --- .../scheduler_components/invariant_checker.py | 21 ++++++++----------- .../server_fixtures/spec_eagle_fixture.py | 14 +++++++++++-- .../spec/eagle/test_eagle_dp_attention.py | 5 ++++- test/registered/spec/eagle/test_spec_eagle.py | 1 - .../spec/eagle/test_spec_eagle_fa3.py | 3 +++ .../spec/eagle/test_spec_eagle_page.py | 9 ++++---- .../spec/eagle/test_spec_eagle_parity.py | 2 ++ .../spec/eagle/test_spec_eagle_stress.py | 13 ++++++++---- .../spec/eagle/test_spec_eagle_topk.py | 6 ++++++ .../spec/eagle/test_spec_eagle_triton.py | 1 + 10 files changed, 50 insertions(+), 25 deletions(-) diff --git a/python/sglang/srt/managers/scheduler_components/invariant_checker.py b/python/sglang/srt/managers/scheduler_components/invariant_checker.py index ede8872f5..6cf379d8c 100644 --- a/python/sglang/srt/managers/scheduler_components/invariant_checker.py +++ b/python/sglang/srt/managers/scheduler_components/invariant_checker.py @@ -1,7 +1,6 @@ from __future__ import annotations import logging -import warnings from collections import deque from dataclasses import dataclass, field from typing import ( @@ -164,12 +163,17 @@ class SchedulerInvariantChecker: """ # After decode: running_batch IS last_batch (same object), count once. # After prefill: they differ, both hold uncached tokens. - batches = [self.get_last_batch()] + # Use identity (is / is not), not membership or ==: ScheduleBatch's + # dataclass __eq__ compares tensor fields and raises on ambiguous bools. + last_batch = self.get_last_batch() + running_batch = self.get_running_batch() + batches = [last_batch] if ( - self.get_running_batch() not in (None, self.get_last_batch()) - and not self.get_running_batch().is_empty() + running_batch is not None + and running_batch is not last_batch + and not running_batch.is_empty() ): - batches.append(self.get_running_batch()) + batches.append(running_batch) full_uncached = 0 swa_uncached = 0 @@ -196,13 +200,6 @@ class SchedulerInvariantChecker: if self.get_last_batch() is None: return - spec_topk = self.server_args.speculative_eagle_topk or 1 - if spec_topk > 1: - warnings.warn( - "Runtime memory check (busy) is not supported when speculation topk > 1." - ) - return - ps = self.pool_stats_observer.get_pool_stats() full_uncached, swa_uncached = self._get_total_uncached_sizes() diff --git a/python/sglang/test/server_fixtures/spec_eagle_fixture.py b/python/sglang/test/server_fixtures/spec_eagle_fixture.py index 32f5b86ba..0ef6ccf21 100644 --- a/python/sglang/test/server_fixtures/spec_eagle_fixture.py +++ b/python/sglang/test/server_fixtures/spec_eagle_fixture.py @@ -66,7 +66,9 @@ class SpecEagleServerBase(CustomTestCase): trust_remote_code = True # -- extras -- - # env_overrides: iterable of (env_var_obj, value) applied only around launch. + # env_overrides: (env_var_obj, value) pairs applied only around launch. + # Declare ONLY this class's own; _merged_env_overrides() unions them down the + # MRO (base first), so never restate a base's. Derived wins on a repeated env. env_overrides = () extra_args = () @@ -105,6 +107,14 @@ class SpecEagleServerBase(CustomTestCase): args += [str(a) for a in cls.extra_args] return args + @classmethod + def _merged_env_overrides(cls): + # Base first so a derived class wins for a repeated env var. + merged = [] + for klass in reversed(cls.__mro__): + merged.extend(klass.__dict__.get("env_overrides", ())) + return merged + @classmethod def setUpClass(cls): cls.base_url = DEFAULT_URL_FOR_TEST @@ -116,7 +126,7 @@ class SpecEagleServerBase(CustomTestCase): stack.enter_context( envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True) ) - for env_var, value in cls.env_overrides: + for env_var, value in cls._merged_env_overrides(): stack.enter_context(env_var.override(value)) cls.process = popen_launch_server( cls.model, diff --git a/test/registered/spec/eagle/test_eagle_dp_attention.py b/test/registered/spec/eagle/test_eagle_dp_attention.py index 8fc48cb1b..9fa6593c5 100644 --- a/test/registered/spec/eagle/test_eagle_dp_attention.py +++ b/test/registered/spec/eagle/test_eagle_dp_attention.py @@ -60,7 +60,10 @@ class TestEAGLE3EngineDPAttention(CustomTestCase): "--cuda-graph-max-bs", "64", ] - with envs.SGLANG_ENABLE_ASYNC_ASSERT.override(True): + with ( + envs.SGLANG_ENABLE_ASYNC_ASSERT.override(True), + envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(1), + ): cls.process = popen_launch_server( cls.model, cls.base_url, diff --git a/test/registered/spec/eagle/test_spec_eagle.py b/test/registered/spec/eagle/test_spec_eagle.py index 9d603b5ec..cee2355cc 100644 --- a/test/registered/spec/eagle/test_spec_eagle.py +++ b/test/registered/spec/eagle/test_spec_eagle.py @@ -31,7 +31,6 @@ _KITS = ( class _Core(Eagle3Base): - # Busy-time pool accounting check (topk=1 only). env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) diff --git a/test/registered/spec/eagle/test_spec_eagle_fa3.py b/test/registered/spec/eagle/test_spec_eagle_fa3.py index 7d00c2bfc..db7fff026 100644 --- a/test/registered/spec/eagle/test_spec_eagle_fa3.py +++ b/test/registered/spec/eagle/test_spec_eagle_fa3.py @@ -6,6 +6,7 @@ fa3 is the real H200 default for MHA spec at topk=1, so this also covers the 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, @@ -25,6 +26,7 @@ class TestEagle3Fa3(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprob attention_backend = "fa3" disable_overlap = False + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagleLlama2Fa3Page256( @@ -43,6 +45,7 @@ class TestEagleLlama2Fa3Page256( page_size = 256 chunked_prefill_size = 4096 # must be divisible by page_size (256) cuda_graph_max_bs = 5 + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) if __name__ == "__main__": diff --git a/test/registered/spec/eagle/test_spec_eagle_page.py b/test/registered/spec/eagle/test_spec_eagle_page.py index dc415db9c..ece61fbc9 100644 --- a/test/registered/spec/eagle/test_spec_eagle_page.py +++ b/test/registered/spec/eagle/test_spec_eagle_page.py @@ -23,24 +23,23 @@ class TestEagle3Page64(Eagle3Base, SpecAccuracyKit, SpecLogprobKit, SpecFeatureK page_size = 64 disable_overlap = False + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagleLlama2Page4Topk1(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit): - """Llama-2 topk=1 + page_size=4; busy-time pool check (topk=1 only).""" + """Llama-2 topk=1 + page_size=4.""" 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), - ) + env_overrides = ((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 + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) if __name__ == "__main__": diff --git a/test/registered/spec/eagle/test_spec_eagle_parity.py b/test/registered/spec/eagle/test_spec_eagle_parity.py index f20cb9c29..d8dd32f0c 100644 --- a/test/registered/spec/eagle/test_spec_eagle_parity.py +++ b/test/registered/spec/eagle/test_spec_eagle_parity.py @@ -6,6 +6,7 @@ spec server (sequential -- one model resident at a time; see SpecParityKit). 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 SpecParityKit from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base @@ -22,6 +23,7 @@ class TestEagle3Parity(SpecParityKit, Eagle3Base): """ disable_overlap = False + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) if __name__ == "__main__": diff --git a/test/registered/spec/eagle/test_spec_eagle_stress.py b/test/registered/spec/eagle/test_spec_eagle_stress.py index af4730377..40e54f585 100644 --- a/test/registered/spec/eagle/test_spec_eagle_stress.py +++ b/test/registered/spec/eagle/test_spec_eagle_stress.py @@ -27,6 +27,7 @@ class TestEagle3Perf(Eagle3Base, SpecPerfKit): """Decode throughput (max_new_tokens=1) on EAGLE3 spec v2.""" disable_overlap = False + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagleLlama2Retract(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit): @@ -35,8 +36,8 @@ class TestEagleLlama2Retract(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit): 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), + (envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1), ) @@ -51,18 +52,22 @@ class TestEagle3Topk16V2Retract(Eagle3Base, SpecAccuracyKit, SpecFeatureKit): max_running_requests = 64 gsm8k_accept_len_thres = 2.4 extra_args = ("--max-total-tokens", 4500) # small KV to trigger retract - env_overrides = ((envs.SGLANG_TEST_RETRACT, True),) + env_overrides = ( + (envs.SGLANG_TEST_RETRACT, True), + (envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1), + ) class TestEagleLlama2AbortAll(EagleLlama2Base, AbortAllMixin): abort_all_max_new_tokens = 4000 + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagleLlama2WaitingTimeout(EagleLlama2Base, WaitingTimeoutMixin): max_running_requests = 1 env_overrides = ( - (envs.SGLANG_ENABLE_SPEC_V2, False), (envs.SGLANG_REQ_WAITING_TIMEOUT, 0.001), + (envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1), ) @@ -70,8 +75,8 @@ 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), + (envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1), ) diff --git a/test/registered/spec/eagle/test_spec_eagle_topk.py b/test/registered/spec/eagle/test_spec_eagle_topk.py index 369a5bb7f..b4f05a26f 100644 --- a/test/registered/spec/eagle/test_spec_eagle_topk.py +++ b/test/registered/spec/eagle/test_spec_eagle_topk.py @@ -8,6 +8,7 @@ where fa3 (Hopper-only) isn't available -- functional sanity only, no perf/stres 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, @@ -31,6 +32,7 @@ class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogp acc_length_thres = 3.1 batch_accept_len_thres = 1.75 gsm8k_accept_len_thres = 2.4 # EAGLE3 topk16 gsm8k accept ~2.48 + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagle3Topk16SpecV2(TestEagle3Topk16, SpecFeatureKit): @@ -50,11 +52,14 @@ class TestEagleLlama2Suite( ): """EAGLE/Llama-2 topk=8 full coverage (kits listed in bases).""" + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) + class TestEagleLlama2Chunked4(EagleLlama2Base, SpecCorrectnessKit): """Correctness under tiny chunked prefill.""" chunked_prefill_size = 4 + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) class TestEagleLlama3TokenMap(EagleLlama2Base, SpecAccuracyKit): @@ -70,6 +75,7 @@ class TestEagleLlama3TokenMap(EagleLlama2Base, SpecAccuracyKit): "--speculative-token-map", "thunlp/LLaMA3-Instruct-8B-FR-Spec/freq_32768.pt", ) + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) if __name__ == "__main__": diff --git a/test/registered/spec/eagle/test_spec_eagle_triton.py b/test/registered/spec/eagle/test_spec_eagle_triton.py index d5d422c1f..63c5ce04b 100644 --- a/test/registered/spec/eagle/test_spec_eagle_triton.py +++ b/test/registered/spec/eagle/test_spec_eagle_triton.py @@ -41,6 +41,7 @@ class TestEagleLlama2Triton(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit): """EAGLE/Llama-2 topk=8 on triton (spec v1).""" attention_backend = "triton" + env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) if __name__ == "__main__":