config: read resolved config via namespace accessors (#33013)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user