[refactor] Read resolved config from server_args fields; retire the flags mirror tier (#30346)
This commit is contained in:
@@ -17,10 +17,6 @@ register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd")
|
||||
|
||||
|
||||
def _make_server_args(*, sampling_backend: str) -> SimpleNamespace:
|
||||
# The install gate reads the resolved backend from the flags tier.
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
|
||||
get_flags().sampling_backend = sampling_backend
|
||||
return SimpleNamespace(sampling_backend=sampling_backend)
|
||||
|
||||
|
||||
|
||||
@@ -23,15 +23,16 @@ register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd")
|
||||
|
||||
def _mock_global_server_args(backend="pytorch"):
|
||||
from sglang.srt.layers import sampler as sampler_mod
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
|
||||
sampler_mod.get_global_server_args = lambda: ServerArgs(
|
||||
model_path="dummy",
|
||||
sampling_backend=backend,
|
||||
# Publish for real: the sampler reads the context slot through
|
||||
# get_server_args(), which a module-attribute rebinding cannot intercept.
|
||||
set_global_server_args_for_scheduler(
|
||||
ServerArgs(model_path="dummy", sampling_backend=backend)
|
||||
)
|
||||
# The sampler reads the resolved backend from the flags tier.
|
||||
get_flags().sampling_backend = backend
|
||||
|
||||
class _DummyTPGroup:
|
||||
device_group = None
|
||||
|
||||
@@ -7,7 +7,6 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
get_global_server_args,
|
||||
@@ -45,7 +44,6 @@ class TestLMHeadFP32(unittest.TestCase):
|
||||
|
||||
def _make_logprocessor(self, vocab_size, enable_fp32):
|
||||
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
|
||||
get_flags().enable_dp_lm_head = False
|
||||
get_global_server_args().enable_fp32_lm_head = enable_fp32
|
||||
cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None)
|
||||
return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None)
|
||||
|
||||
@@ -38,11 +38,9 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch:
|
||||
|
||||
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
|
||||
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
|
||||
with get_parallel().override(attn_tp_size=1), patch.object(
|
||||
tbo, "get_global_server_args", lambda: fake_args
|
||||
), get_flags().attn.override(backend="fa3"):
|
||||
):
|
||||
return TboForwardBatchPreparer.filter_batch(
|
||||
batch,
|
||||
start_token_index=lo,
|
||||
|
||||
@@ -102,10 +102,6 @@ def _make_model_runner(
|
||||
|
||||
sa = SimpleNamespace()
|
||||
sa.swa_full_tokens_ratio = swa_full_tokens_ratio
|
||||
# The configurator reads the resolved ratio from the flags tier.
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
|
||||
get_flags().swa_full_tokens_ratio = swa_full_tokens_ratio
|
||||
sa.page_size = page_size
|
||||
sa.disable_radix_cache = disable_radix_cache
|
||||
sa.chunked_prefill_size = chunked_prefill_size
|
||||
|
||||
@@ -2,7 +2,7 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM
|
||||
from sglang.srt.runtime_context import get_context, get_flags, reset_context
|
||||
from sglang.srt.runtime_context import get_context, reset_context
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
@@ -10,9 +10,8 @@ register_cpu_ci(est_time=4, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
|
||||
"""The disable decision is a load-time resolution: it lands on the flags
|
||||
tier through declare_load_time_override (dual-applied onto the published
|
||||
config during the transition)."""
|
||||
"""The disable decision is a load-time resolution: it writes through to
|
||||
the published config via declare_load_time_override."""
|
||||
|
||||
def setUp(self):
|
||||
self._saved_server_args = get_context()._server_args
|
||||
@@ -41,7 +40,6 @@ class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
|
||||
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
|
||||
|
||||
self.assertEqual(model.num_fused_shared_experts, 0)
|
||||
self.assertTrue(get_flags().disable_shared_experts_fusion)
|
||||
# post-init declaration writes through to the published config
|
||||
self.assertTrue(server_args.disable_shared_experts_fusion)
|
||||
|
||||
@@ -52,7 +50,6 @@ class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
|
||||
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
|
||||
|
||||
self.assertEqual(model.num_fused_shared_experts, 1)
|
||||
self.assertFalse(get_flags().disable_shared_experts_fusion)
|
||||
self.assertFalse(server_args.disable_shared_experts_fusion)
|
||||
|
||||
|
||||
|
||||
@@ -18,14 +18,11 @@ from unittest.mock import patch
|
||||
from sglang.srt.arg_groups import overrides as overrides_module
|
||||
from sglang.srt.arg_groups.arg_utils import A, Arg, resolvable_fields
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
OverrideRecord,
|
||||
apply_model_overrides,
|
||||
collect_model_override_declarations,
|
||||
register_model_override,
|
||||
validate_declarations,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
_StaticFlags,
|
||||
get_context,
|
||||
get_server_args,
|
||||
reset_context,
|
||||
@@ -233,89 +230,6 @@ class TestResolvedViewAndPasses(CustomTestCase):
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _FakeAttnGroup(_StaticFlags):
|
||||
backend: str = "unset"
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _FakeFlags(_StaticFlags):
|
||||
attn: _FakeAttnGroup = dataclasses.field(default_factory=_FakeAttnGroup)
|
||||
resolved_by_model: str = "unset"
|
||||
also_resolved: Optional[int] = None
|
||||
|
||||
|
||||
class TestApplyModelOverridesGate(CustomTestCase):
|
||||
def _fresh(self):
|
||||
return _FakeFlags(), _FakeArgs()
|
||||
|
||||
def test_materializes_declared_and_pristine_leaves(self):
|
||||
flags, args = self._fresh()
|
||||
records = apply_model_overrides(
|
||||
flags, args, [("src", {"resolved_by_model": "dsv4"})]
|
||||
)
|
||||
self.assertEqual(flags.resolved_by_model, "dsv4") # declared
|
||||
self.assertIsNone(flags.also_resolved) # undeclared -> pristine value
|
||||
self.assertEqual(args.resolved_by_model, "auto") # server_args untouched
|
||||
self.assertEqual(
|
||||
records, [OverrideRecord("src", "resolved_by_model", "auto", "dsv4")]
|
||||
)
|
||||
|
||||
def test_last_writer_wins_then_terminal_wins_last(self):
|
||||
flags, args = self._fresh()
|
||||
records = apply_model_overrides(
|
||||
flags,
|
||||
args,
|
||||
[
|
||||
("first", {"resolved_by_model": "a"}),
|
||||
("second", {"resolved_by_model": "b"}),
|
||||
],
|
||||
terminal=[("enforce_disable", {"resolved_by_model": "off"})],
|
||||
)
|
||||
self.assertEqual(flags.resolved_by_model, "off")
|
||||
self.assertEqual([r.resolved for r in records], ["a", "b", "off"])
|
||||
self.assertEqual(records[1].base, "a") # provenance chains the writers
|
||||
|
||||
def test_non_whitelisted_field_rejected_before_any_write(self):
|
||||
flags, args = self._fresh()
|
||||
with self.assertRaises(ValueError):
|
||||
apply_model_overrides(
|
||||
flags,
|
||||
args,
|
||||
[("ok", {"resolved_by_model": "x"}), ("bad", {"plain": 1})],
|
||||
)
|
||||
self.assertEqual(flags.resolved_by_model, "unset") # transactional
|
||||
|
||||
def test_missing_leaf_rejected_before_any_write(self):
|
||||
flags, args = self._fresh()
|
||||
with self.assertRaises(ValueError):
|
||||
apply_model_overrides(
|
||||
flags,
|
||||
args,
|
||||
[("src", {"resolved_by_model": "x"})],
|
||||
whitelist={"resolved_by_model", "field_without_leaf"},
|
||||
)
|
||||
self.assertEqual(flags.resolved_by_model, "unset")
|
||||
|
||||
def test_frozen_flags_rejected(self):
|
||||
flags, args = self._fresh()
|
||||
flags.freeze()
|
||||
with self.assertRaises(RuntimeError):
|
||||
apply_model_overrides(flags, args, [("src", {"resolved_by_model": "x"})])
|
||||
|
||||
def test_leaf_map_routes_to_group_leaf(self):
|
||||
flags, args = self._fresh()
|
||||
apply_model_overrides(
|
||||
flags,
|
||||
args,
|
||||
[("src", {"resolved_by_model": "fa3"})],
|
||||
whitelist={"resolved_by_model"},
|
||||
leaf_map={"resolved_by_model": "attn.backend"},
|
||||
)
|
||||
self.assertEqual(flags.attn.backend, "fa3")
|
||||
self.assertEqual(flags.resolved_by_model, "unset") # flat leaf untouched
|
||||
|
||||
|
||||
class _IsolatedPublish(CustomTestCase):
|
||||
"""Publishing writes the process-global context; save/restore around it."""
|
||||
|
||||
@@ -335,9 +249,9 @@ class _NoOverridableArgs:
|
||||
x: int = 1
|
||||
|
||||
|
||||
class TestPublishResolvesFlags(_IsolatedPublish):
|
||||
"""Publish wiring: stash-carrying publishes resolve into flags via the
|
||||
gate; publishes without the stash skip resolution."""
|
||||
class TestPublishInstallsSlot(_IsolatedPublish):
|
||||
"""Publish wiring: set_server_args installs the already-resolved object
|
||||
into the context-owned slot (no transformation at publish time)."""
|
||||
|
||||
def test_dummy_fixture_has_empty_stash_and_publishes_cleanly(self):
|
||||
from sglang.srt.server_args import (
|
||||
@@ -357,23 +271,12 @@ class TestPublishResolvesFlags(_IsolatedPublish):
|
||||
get_context().set_server_args(sa)
|
||||
self.assertIs(get_server_args(), sa)
|
||||
|
||||
def test_non_whitelisted_declaration_fails_at_publish(self):
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
|
||||
flags_before = get_flags()
|
||||
sa = _NoOverridableArgs()
|
||||
sa._resolved_overrides = [("rogue", {"x": 2})]
|
||||
with self.assertRaises(ValueError):
|
||||
get_context().set_server_args(sa)
|
||||
# a failed publish must leave BOTH the slot and the flags untouched
|
||||
self.assertIs(get_flags(), flags_before)
|
||||
|
||||
|
||||
class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
"""Per-arch golden diff for migrated families: the declarative path must
|
||||
reproduce the legacy imperative writes byte-identically on server_args
|
||||
(dual-apply) and materialize the same values on the flags tier at
|
||||
publish."""
|
||||
reproduce the legacy imperative writes byte-identically on the
|
||||
materialized server_args fields; the publish round-trip returns the same
|
||||
object."""
|
||||
|
||||
_MINI_CONFIG = {
|
||||
"hidden_size": 64,
|
||||
@@ -408,11 +311,13 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
return ServerArgs(model_path=config_dir, **server_kwargs)
|
||||
|
||||
def _publish(self, server_args):
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.server_args import set_global_server_args_for_scheduler
|
||||
from sglang.srt.server_args import (
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
return get_flags()
|
||||
return get_global_server_args()
|
||||
|
||||
def test_mistral_large3_forces_bfloat16(self):
|
||||
sa = self._construct("MistralLarge3ForCausalLM", "mistral")
|
||||
@@ -601,7 +506,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
]
|
||||
self.assertEqual(len(deterministic_fills), 1)
|
||||
self.assertEqual(sa.attention_backend, deterministic_fills[0])
|
||||
self.assertEqual(flags.attn.backend, deterministic_fills[0])
|
||||
self.assertEqual(flags.attention_backend, deterministic_fills[0])
|
||||
|
||||
def test_deterministic_incompatible_backend_raises(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
@@ -645,8 +550,8 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
("_dllm_attention_backend", {"attention_backend": "flashinfer"}),
|
||||
sa._resolved_overrides,
|
||||
)
|
||||
# first MAPPED leaf: attention_backend routes to flags.attn.backend
|
||||
self.assertEqual(self._publish(sa).attn.backend, "flashinfer")
|
||||
# the deterministic fill lands on the attention_backend field
|
||||
self.assertEqual(self._publish(sa).attention_backend, "flashinfer")
|
||||
|
||||
def test_attention_backend_leaf_materializes_end_state(self):
|
||||
# The default-fill pass declares the platform-selected backend; the
|
||||
@@ -660,7 +565,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
]
|
||||
self.assertTrue(declared_values) # default fill declared
|
||||
self.assertEqual(sa.attention_backend, declared_values[-1]) # materialized
|
||||
self.assertEqual(self._publish(sa).attn.backend, declared_values[-1])
|
||||
self.assertEqual(self._publish(sa).attention_backend, declared_values[-1])
|
||||
|
||||
def test_post_materialize_pass_writes_through(self):
|
||||
from sglang.srt.arg_groups.overrides import run_post_process_pass
|
||||
@@ -679,12 +584,12 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
run_post_process_pass(sa, _force_triton)
|
||||
if resolved_before != "triton":
|
||||
self.assertEqual(sa.attention_backend, "triton")
|
||||
self.assertEqual(self._publish(sa).attn.backend, sa.attention_backend)
|
||||
self.assertEqual(self._publish(sa).attention_backend, sa.attention_backend)
|
||||
|
||||
def test_attention_backend_user_choice_declares_nothing_extra(self):
|
||||
sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton")
|
||||
self.assertEqual(sa.attention_backend, "triton")
|
||||
self.assertEqual(self._publish(sa).attn.backend, "triton")
|
||||
self.assertEqual(self._publish(sa).attention_backend, "triton")
|
||||
|
||||
def test_compatibility_passes_at_callable_level(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
@@ -2158,13 +2063,10 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
class TestDeclarationValidation(CustomTestCase):
|
||||
def test_declarations_never_mutate_server_args(self):
|
||||
flags, args = _FakeFlags(), _FakeArgs()
|
||||
args = _FakeArgs()
|
||||
declarations = [("src", {"resolved_by_model": "dsv4", "also_resolved": 7})]
|
||||
apply_model_overrides(flags, args, declarations)
|
||||
validate_declarations(args, declarations)
|
||||
# the leaves carry the declared values; the fields stay pristine
|
||||
self.assertEqual(flags.resolved_by_model, "dsv4")
|
||||
self.assertEqual(flags.also_resolved, 7)
|
||||
# validation is a pure whitelist check: the fields stay untouched
|
||||
self.assertEqual(args.resolved_by_model, _FakeArgs.resolved_by_model)
|
||||
self.assertEqual(args.also_resolved, _FakeArgs.also_resolved)
|
||||
|
||||
|
||||
@@ -15,13 +15,11 @@ from sglang.srt.runtime_context import (
|
||||
ParallelContext,
|
||||
RuntimeContext,
|
||||
_FlagGroupBase,
|
||||
_StaticFlags,
|
||||
get_context,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
reset_context,
|
||||
resolve_flag_leaf,
|
||||
)
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -215,91 +213,47 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
|
||||
self.assertFalse(hasattr(server_args_module, "_global_server_args"))
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _FakeStaticGroup(_StaticFlags):
|
||||
alpha: int = 1
|
||||
beta: str = "b"
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _FakeCaptureGroup(_FlagGroupBase):
|
||||
gamma: int = 0
|
||||
|
||||
|
||||
class TestFlagsTier(_IsolatedServerArgs):
|
||||
"""V3a skeleton: typed dataclass groups, freeze guard, override primitive."""
|
||||
"""Runtime-flags tier: typed groups, typo-safe writes, override primitive.
|
||||
|
||||
Resolved configuration lives on server_args fields (materialized at the
|
||||
end of __post_init__); the flags tier only carries runtime state
|
||||
(today: the capture lifecycle)."""
|
||||
|
||||
def test_wiring_and_groups(self):
|
||||
flags = get_flags()
|
||||
self.assertIs(flags, get_context().flags)
|
||||
self.assertIsInstance(flags, Flags)
|
||||
for group in ("attn", "moe", "capture"):
|
||||
self.assertTrue(hasattr(flags, group))
|
||||
self.assertFalse(flags.frozen)
|
||||
self.assertTrue(hasattr(flags, "capture"))
|
||||
|
||||
def test_typo_safety(self):
|
||||
group = _FakeStaticGroup()
|
||||
group = _FakeCaptureGroup()
|
||||
with self.assertRaises(AttributeError):
|
||||
group.alpha_misspelled = 2 # undeclared leaf
|
||||
group.gamma_misspelled = 2 # undeclared leaf
|
||||
with self.assertRaises(AttributeError):
|
||||
get_flags().not_a_flag = 1
|
||||
|
||||
def test_static_group_writable_until_freeze(self):
|
||||
group = _FakeStaticGroup()
|
||||
group.alpha = 5
|
||||
self.assertEqual(group.alpha, 5)
|
||||
group.freeze()
|
||||
with self.assertRaises(RuntimeError):
|
||||
group.alpha = 6
|
||||
self.assertEqual(group.alpha, 5)
|
||||
|
||||
def test_override_is_transactional_and_works_on_frozen(self):
|
||||
group = _FakeStaticGroup()
|
||||
group.freeze()
|
||||
with group.override(alpha=99, beta="x"):
|
||||
self.assertEqual(group.alpha, 99)
|
||||
self.assertEqual(group.beta, "x")
|
||||
self.assertEqual(group.alpha, 1)
|
||||
self.assertEqual(group.beta, "b")
|
||||
with self.assertRaises(ValueError):
|
||||
with group.override(alpha=2, gamma=3): # gamma undeclared
|
||||
pass
|
||||
self.assertEqual(group.alpha, 1) # validated before any write
|
||||
|
||||
def test_non_static_group_has_no_freeze(self):
|
||||
def test_override_is_transactional(self):
|
||||
group = _FakeCaptureGroup()
|
||||
group.gamma = 42
|
||||
self.assertEqual(group.gamma, 42)
|
||||
self.assertFalse(hasattr(group, "freeze"))
|
||||
with group.override(gamma=99):
|
||||
self.assertEqual(group.gamma, 99)
|
||||
self.assertEqual(group.gamma, 0)
|
||||
with self.assertRaises(ValueError):
|
||||
with group.override(gamma=2, delta=3): # delta undeclared
|
||||
pass
|
||||
self.assertEqual(group.gamma, 0) # validated before any write
|
||||
|
||||
def test_container_freeze_cascades_except_capture(self):
|
||||
flags = Flags() # fresh container, not the process singleton
|
||||
flags.freeze()
|
||||
self.assertTrue(flags.frozen)
|
||||
self.assertTrue(flags.attn.frozen)
|
||||
self.assertTrue(flags.moe.frozen)
|
||||
with self.assertRaises(RuntimeError):
|
||||
flags.attn = flags.attn # container leaves lock too
|
||||
self.assertFalse(getattr(flags.capture, "_frozen", False))
|
||||
|
||||
def test_resolve_flag_leaf_flat_default_and_mapped(self):
|
||||
flags = Flags()
|
||||
owner, leaf = resolve_flag_leaf(flags, "some_field")
|
||||
self.assertIs(owner, flags)
|
||||
self.assertEqual(leaf, "some_field")
|
||||
owner, leaf = resolve_flag_leaf(flags, "x", leaf_map={"x": "attn.x"})
|
||||
self.assertIs(owner, flags.attn)
|
||||
self.assertEqual(leaf, "x")
|
||||
|
||||
def test_reset_context_installs_fresh_unfrozen_flags(self):
|
||||
try:
|
||||
old = get_flags()
|
||||
old.freeze()
|
||||
reset_context()
|
||||
self.assertIsNot(get_flags(), old)
|
||||
self.assertFalse(get_flags().frozen)
|
||||
finally:
|
||||
reset_context() # never leave the singleton frozen for other tests
|
||||
def test_reset_context_installs_fresh_flags(self):
|
||||
old = get_flags()
|
||||
old.capture.enable_torch_compile = True
|
||||
reset_context()
|
||||
self.assertIsNot(get_flags(), old)
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -311,84 +265,15 @@ class _FakeResolvedArgs:
|
||||
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
||||
|
||||
|
||||
class TestRuntimeResolutionStages(_IsolatedServerArgs):
|
||||
"""Runtime stages: post-publish declarations re-resolve the flags tier
|
||||
atomically; freeze_flags() ends the resolution lifecycle."""
|
||||
class TestPublishLifecycle(_IsolatedServerArgs):
|
||||
"""Publish installs the resolved server_args and seeds the capture tier."""
|
||||
|
||||
def _publish(self, **kw):
|
||||
args = _FakeResolvedArgs(**kw)
|
||||
get_context().set_server_args(args)
|
||||
return args
|
||||
|
||||
def test_record_before_publish_raises(self):
|
||||
reset_context()
|
||||
with self.assertRaises(ValueError):
|
||||
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
|
||||
|
||||
def test_record_updates_leaves_and_accumulates_stages(self):
|
||||
args = self._publish(page_size=1, sampling_backend="flashinfer")
|
||||
self.assertEqual(get_flags().page_size, 1) # publish-time materialize
|
||||
# dual-apply transition: the call site keeps its imperative write
|
||||
args.page_size = 64
|
||||
get_context().record_runtime_overrides([("stage.runner", {"page_size": 64})])
|
||||
self.assertEqual(get_flags().page_size, 64)
|
||||
args.sampling_backend = "pytorch"
|
||||
get_context().record_runtime_overrides(
|
||||
[("stage.load", {"sampling_backend": "pytorch"})]
|
||||
)
|
||||
self.assertEqual(get_flags().sampling_backend, "pytorch")
|
||||
self.assertEqual(get_flags().page_size, 64) # earlier stage survives
|
||||
|
||||
def test_record_whitelist_violation_rolls_back(self):
|
||||
self._publish()
|
||||
with self.assertRaises(ValueError):
|
||||
get_context().record_runtime_overrides([("bad", {"nope": 1})])
|
||||
self.assertEqual(get_context()._runtime_overrides, [])
|
||||
|
||||
def test_freeze_ends_the_resolution_lifecycle(self):
|
||||
args = self._publish(page_size=1)
|
||||
try:
|
||||
get_context().freeze_flags()
|
||||
self.assertTrue(get_flags().frozen)
|
||||
with self.assertRaises(RuntimeError):
|
||||
get_context().record_runtime_overrides([("late", {"page_size": 64})])
|
||||
with self.assertRaises(RuntimeError):
|
||||
get_context().set_server_args(args)
|
||||
finally:
|
||||
reset_context()
|
||||
|
||||
def test_declare_load_time_override_applies_and_records(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})
|
||||
# post-init declaration: written through to the field and resolved
|
||||
# into the leaf
|
||||
self.assertEqual(args.page_size, 64)
|
||||
self.assertEqual(get_flags().page_size, 64)
|
||||
self.assertEqual(
|
||||
get_context()._runtime_overrides,
|
||||
[("model.load_time", {"page_size": 64})],
|
||||
)
|
||||
|
||||
def test_failed_republish_keeps_previous_lifecycle(self):
|
||||
args = self._publish(page_size=1)
|
||||
args.page_size = 64
|
||||
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
|
||||
flags_before = get_flags()
|
||||
bad = _FakeResolvedArgs(page_size=1)
|
||||
bad._resolved_overrides = [("bad", {"nope": 1})] # gate rejects
|
||||
with self.assertRaises(ValueError):
|
||||
get_context().set_server_args(bad)
|
||||
# previous publish fully intact: slot, flags, and the recorded stages
|
||||
self.assertIs(get_context()._server_args, args)
|
||||
self.assertIs(get_flags(), flags_before)
|
||||
self.assertEqual(
|
||||
get_context()._runtime_overrides, [("stage", {"page_size": 64})]
|
||||
)
|
||||
|
||||
def test_capture_tier_seeded_at_publish_and_survives_stages(self):
|
||||
# seeded from the published config
|
||||
def test_capture_tier_seeded_at_publish(self):
|
||||
args = self._publish(page_size=1)
|
||||
args.enable_torch_compile = True
|
||||
get_context().set_server_args(args) # re-publish picks up the value
|
||||
@@ -396,59 +281,38 @@ class TestRuntimeResolutionStages(_IsolatedServerArgs):
|
||||
# capture-time write (B4) targets the capture leaf
|
||||
get_flags().capture.enable_torch_compile = False
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
# a runtime-stage re-resolve must not clobber the capture write
|
||||
args.page_size = 64
|
||||
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
# capture stays writable after freeze
|
||||
try:
|
||||
get_context().freeze_flags()
|
||||
get_flags().capture.enable_torch_compile = True
|
||||
self.assertTrue(get_flags().capture.enable_torch_compile)
|
||||
finally:
|
||||
reset_context()
|
||||
|
||||
def test_declared_leaf_wins_over_stale_field(self):
|
||||
# A stash entry always drives the leaf at publish, even if the field
|
||||
# value diverged (e.g. a fixture that skipped materialization).
|
||||
@dataclasses.dataclass
|
||||
class _Args:
|
||||
enable_dp_lm_head: A[bool, Arg(help="d", resolvable=True)] = True
|
||||
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
||||
|
||||
args = _Args()
|
||||
args._resolved_overrides = [("dp", {"enable_dp_lm_head": False})]
|
||||
args._declarations_materialized = True
|
||||
args.enable_dp_lm_head = False
|
||||
get_context().set_server_args(args)
|
||||
self.assertFalse(get_flags().enable_dp_lm_head)
|
||||
|
||||
def test_bare_dataclass_publish_skips_materialization(self):
|
||||
# object.__new__(ServerArgs) fixtures (no __init__, no field values)
|
||||
# must publish without touching the flags tier — dataclass defaults
|
||||
# live on the class, so materializing from them would clobber
|
||||
# previously resolved flags with defaults.
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
self._publish(page_size=64)
|
||||
self.assertEqual(get_flags().page_size, 64)
|
||||
bare = object.__new__(ServerArgs)
|
||||
get_context().set_server_args(bare)
|
||||
self.assertIs(get_server_args(), bare)
|
||||
self.assertEqual(get_flags().page_size, 64) # not clobbered
|
||||
|
||||
def test_capture_tier_defaults_for_sentinel_publish(self):
|
||||
get_context().set_server_args(object())
|
||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||
|
||||
def test_republish_clears_runtime_overrides(self):
|
||||
def test_declare_load_time_override_writes_through(self):
|
||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||
|
||||
args = self._publish(page_size=1)
|
||||
args.page_size = 64
|
||||
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
|
||||
self.assertEqual(get_flags().page_size, 64)
|
||||
self._publish(page_size=1) # fresh lifecycle
|
||||
self.assertEqual(get_flags().page_size, 1)
|
||||
self.assertEqual(get_context()._runtime_overrides, [])
|
||||
declare_load_time_override("model.load_time", {"page_size": 64})
|
||||
self.assertEqual(args.page_size, 64)
|
||||
|
||||
def test_declare_load_time_override_validates_whitelist(self):
|
||||
from sglang.srt.arg_groups.overrides import declare_load_time_override
|
||||
|
||||
args = self._publish(page_size=1)
|
||||
with self.assertRaises(ValueError):
|
||||
declare_load_time_override("bad", {"nope": 1})
|
||||
self.assertEqual(args.page_size, 1)
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -68,8 +68,8 @@ class TestServerArgsMutationRatchet(CustomTestCase):
|
||||
f"server_args mutations outside the resolution pipeline grew: "
|
||||
f"{count} > baseline {_BASELINE}. Configuration is resolved in "
|
||||
"ServerArgs.__post_init__; declare through the pipeline "
|
||||
"(passes / declare_load_time_override / "
|
||||
"record_runtime_overrides) instead of assigning fields."
|
||||
"(passes / declare_load_time_override) or go through "
|
||||
"ServerArgs.override(source, ...) instead of assigning fields."
|
||||
)
|
||||
if count < _BASELINE:
|
||||
self.fail(
|
||||
|
||||
Reference in New Issue
Block a user