config: read resolved config via namespace accessors (#33013)

This commit is contained in:
Cheng Wan
2026-07-31 15:06:59 -07:00
committed by GitHub
parent 4862edc85f
commit 55b6769b0e
187 changed files with 1110 additions and 923 deletions
@@ -19,6 +19,7 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt import runtime_context as rc
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
@@ -136,6 +137,17 @@ class TestGetDcpLens(CustomTestCase):
)
allocators = {}
# The configurator's bag reads (disaggregation_mode / page_size /
# enable_hisparse) come from the published context; the per-iteration
# dcp_size stays on the injected instance stand-in.
self._sa_override = rc.get_context().override_server_args(
disaggregation_mode="null",
page_size=physical_page_size,
enable_hisparse=False,
)
self._sa_override.install()
self.addCleanup(self._sa_override.restore)
for dcp_size in (1, 4):
configurator = SimpleNamespace(
server_args=SimpleNamespace(
@@ -117,8 +117,17 @@ class TestMambaRatioEnvGate(unittest.TestCase):
enable_mamba_extra_buffer_lazy=lambda: lazy,
)
fake = SimpleNamespace(server_args=server_args)
# The bag reads (disable_radix_cache / disable_overlap_schedule) come
# from the published context; the derived-method calls stay on the
# injected stand-in.
from sglang.srt import runtime_context as rc
with envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.override(skip):
return KVCacheConfigurator._calculate_mamba_ratio(fake)
with rc.get_context().override_server_args(
disable_radix_cache=False,
disable_overlap_schedule=disable_overlap,
):
return KVCacheConfigurator._calculate_mamba_ratio(fake)
def test_flag_off_restores_original_ratios(self):
r = lambda **kw: self._ratio(skip=False, **kw)
@@ -340,6 +340,10 @@ class TestPrefetchDispatch(CustomTestCase):
drop_cache,
),
),
patch(
"sglang.srt.model_loader.loader.get_model",
return_value=self._server_args(prefetch, disable_mmap, drop_cache),
),
patch(
"sglang.srt.model_loader.loader."
"buffered_multi_thread_safetensors_weights_iterator",
@@ -356,12 +360,13 @@ class TestPrefetchDispatch(CustomTestCase):
"""Prefetch on + no explicit multithread config -> single-threaded,
and the opt-out warning fires once."""
loader = self._make_loader({})
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -375,12 +380,13 @@ class TestPrefetchDispatch(CustomTestCase):
"""Explicit enable_multithread_load=true is the escape hatch; the
override and its warning must not fire."""
loader = self._make_loader({"enable_multithread_load": True})
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -395,12 +401,13 @@ class TestPrefetchDispatch(CustomTestCase):
default) also signals multi-thread intent, so the override must not
fire and num_threads stays live."""
loader = self._make_loader({"num_threads": 64})
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -416,12 +423,13 @@ class TestPrefetchDispatch(CustomTestCase):
"""Prefetch off -> multi-threaded iterator is used (default), no
override warning."""
loader = self._make_loader({})
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -435,12 +443,13 @@ class TestPrefetchDispatch(CustomTestCase):
"""Prefetch is a no-op without mmap, so the override and its warning
must not fire."""
loader = self._make_loader({})
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True, disable_mmap=True
)
with (
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -454,7 +463,7 @@ class TestPrefetchDispatch(CustomTestCase):
"""FASTSAFETENSORS ignores both flags; override + warning must not
fire."""
loader = self._make_loader({}, load_format=LoadFormat.FASTSAFETENSORS)
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=True
)
with (
@@ -464,6 +473,7 @@ class TestPrefetchDispatch(CustomTestCase):
) as mock_fast,
p_prep,
p_args,
p_model,
p_buffered as mock_buffered,
p_single as mock_single,
p_warn as mock_warning,
@@ -482,7 +492,7 @@ class TestPrefetchDispatch(CustomTestCase):
loader = self._make_loader(
{"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS
)
p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch(
p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch(
prefetch=False,
drop_cache=True,
)
@@ -493,6 +503,7 @@ class TestPrefetchDispatch(CustomTestCase):
) as mock_fast,
p_prep,
p_args,
p_model,
p_buffered,
p_single,
p_warn,
@@ -819,6 +819,16 @@ class TestShardConfig(unittest.TestCase):
), mock.patch(
"sglang.srt.model_loader.loader.get_parallel",
return_value=parallel,
), mock.patch(
"sglang.srt.model_loader.loader.get_exec",
return_value=SimpleNamespace(
features=SimpleNamespace(enable_fp32_lm_head=True),
moe=SimpleNamespace(
ep_num_redundant_experts=4,
enable_eplb=True,
init_expert_location="trivial",
),
),
), mock.patch.object(
loader, "_compute_structural_signature", return_value="sig16"
):
@@ -6,6 +6,7 @@ register_cpu_ci(est_time=9, suite="base-a-test-cpu")
register_cpu_ci(est_time=8, suite="base-c-test-cpu")
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
@@ -455,6 +456,22 @@ class TestCopyForForward(CustomTestCase):
# from_schedule_batch
class TestFromScheduleBatch(CustomTestCase):
def setUp(self):
super().setUp()
# from_schedule_batch reads these two flags from the exec bag; give
# each test a mutable stand-in so it does not depend on a published
# (or leaked) process context.
self._exec_ns = SimpleNamespace(
deterministic=SimpleNamespace(enable_deterministic_inference=False),
features=SimpleNamespace(enable_custom_logit_processor=False),
)
exec_patch = patch(
"sglang.srt.sampling.sampling_batch_info.get_exec",
return_value=self._exec_ns,
)
exec_patch.start()
self.addCleanup(exec_patch.stop)
def _make_req(
self,
temp=1.0,
@@ -537,6 +554,7 @@ class TestFromScheduleBatch(CustomTestCase):
"""Test that explicit seed=123 is kept and missing seed defaults to 42."""
mock_server_args.return_value.enable_deterministic_inference = True
mock_server_args.return_value.enable_custom_logit_processor = False
self._exec_ns.deterministic.enable_deterministic_inference = True
reqs = [self._make_req(seed=123), self._make_req(seed=None)]
batch = MagicMock()
@@ -585,6 +603,7 @@ class TestFromScheduleBatch(CustomTestCase):
mock_server_args.return_value.enable_deterministic_inference = False
mock_server_args.return_value.enable_custom_logit_processor = True
self._exec_ns.features.enable_custom_logit_processor = True
proc_str = DisallowedTokensLogitsProcessor.to_str()
req1 = self._make_req()
@@ -200,6 +200,9 @@ class TestNgramMambaVerifyUpdate(CustomTestCase):
), patch(
"sglang.srt.speculative.spec_utils.get_server_args",
return_value=MagicMock(mamba_track_interval=256),
), patch(
"sglang.srt.speculative.spec_utils.get_exec",
return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)),
):
commit_mamba_states_after_verify(
target_worker,
+17 -13
View File
@@ -10,7 +10,7 @@ import unittest
from unittest.mock import patch
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups.arg_utils import A, Arg
from sglang.srt.arg_groups.arg_utils import NS, A, Arg
from sglang.srt.runtime_context import (
Flags,
ParallelContext,
@@ -19,6 +19,7 @@ from sglang.srt.runtime_context import (
get_context,
get_flags,
get_parallel,
get_schedule,
get_server_args,
reset_context,
)
@@ -375,8 +376,10 @@ class TestFlagsTier(_IsolatedServerArgs):
class _FakeResolvedArgs:
"""Publishable fixture with a resolvable whitelist (real flat leaves)."""
page_size: A[int | None, Arg(help="p", resolvable=True)] = None
sampling_backend: A[str | None, Arg(help="s", resolvable=True)] = None
page_size: A[int | None, Arg(help="p", resolvable=True), NS("schedule")] = None
sampling_backend: A[
str | None, Arg(help="s", resolvable=True), NS("exec.kernel")
] = None
_resolved_overrides: list = dataclasses.field(default_factory=list)
@@ -935,12 +938,15 @@ class TestPublishLifecycle(_IsolatedServerArgs):
get_context().set_server_args(object())
self.assertFalse(get_flags().capture.enable_torch_compile)
def test_declare_load_time_override_writes_through(self):
def test_declare_load_time_override_writes_the_bag(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
args = self._publish(page_size=1)
declare_load_time_override("model.load_time", {"page_size": 64})
self.assertEqual(args.page_size, 64)
# The declaration lands on the config bag; the pristine startup record
# (server_args) is untouched.
self.assertEqual(get_schedule().page_size, 64)
self.assertEqual(args.page_size, 1)
def test_declare_load_time_override_validates_whitelist(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
@@ -952,16 +958,14 @@ class TestPublishLifecycle(_IsolatedServerArgs):
def test_declare_load_time_override_records_provenance(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
from sglang.srt.server_args import ServerArgs
class _Args(_FakeResolvedArgs):
override = ServerArgs.override
args = _Args(page_size=1)
get_context().set_server_args(args)
self._publish(page_size=1)
declare_load_time_override("model.load_time", {"page_size": 64})
self.assertEqual(args.page_size, 64)
self.assertIn(("model.load_time", {"page_size": 64}), args._resolved_overrides)
self.assertEqual(get_schedule().page_size, 64)
self.assertIn(
("model.load_time", {"page_size": 64}),
get_context().overrides_log(),
)
if __name__ == "__main__":
@@ -49,7 +49,7 @@ _EXCLUDED = (
"multimodal_gen",
)
_BASELINE = 49
_BASELINE = 39
class TestServerArgsWriterRatchet(CustomTestCase):