[refactor] Resolve config declarations onto server_args at the end of __post_init__ (#30297)
This commit is contained in:
@@ -17,6 +17,10 @@ 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)
|
||||
|
||||
|
||||
|
||||
@@ -102,6 +102,10 @@ 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
|
||||
|
||||
@@ -42,7 +42,7 @@ class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
|
||||
|
||||
self.assertEqual(model.num_fused_shared_experts, 0)
|
||||
self.assertTrue(get_flags().disable_shared_experts_fusion)
|
||||
# dual-apply transition: the published config carries the value too
|
||||
# post-init declaration writes through to the published config
|
||||
self.assertTrue(server_args.disable_shared_experts_fusion)
|
||||
|
||||
def test_enables_shared_fusion_when_enforced(self):
|
||||
|
||||
@@ -354,7 +354,10 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
|
||||
args._handle_attention_backend_compatibility()
|
||||
|
||||
self.assertEqual(args.page_size, 128)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
||||
self.assertEqual(resolved_view(args).page_size, 128)
|
||||
|
||||
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
|
||||
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
|
||||
@@ -364,7 +367,10 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
|
||||
|
||||
args._handle_attention_backend_compatibility()
|
||||
|
||||
self.assertEqual(args.page_size, 128)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
self.assertEqual(args.page_size, 1) # dual-apply retired: pristine
|
||||
self.assertEqual(resolved_view(args).page_size, 128)
|
||||
|
||||
|
||||
class TestContextParallelServerArgs(CustomTestCase):
|
||||
@@ -1119,7 +1125,11 @@ class TestDeepEPWaterfillArgs(CustomTestCase):
|
||||
# dummy-model path short-circuits __post_init__; invoke the handler directly.
|
||||
server_args._handle_a2a_moe()
|
||||
|
||||
self.assertFalse(server_args.disable_shared_experts_fusion)
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
# dual-apply retired: the fields stay pristine, the declarations win
|
||||
self.assertTrue(server_args.disable_shared_experts_fusion)
|
||||
self.assertFalse(resolved_view(server_args).disable_shared_experts_fusion)
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
|
||||
def test_waterfill_overrides_moe_a2a_backend_to_deepep(self):
|
||||
@@ -1131,7 +1141,10 @@ class TestDeepEPWaterfillArgs(CustomTestCase):
|
||||
# dummy-model path short-circuits __post_init__; invoke the handler directly.
|
||||
server_args._handle_a2a_moe()
|
||||
|
||||
self.assertEqual(server_args.moe_a2a_backend, "deepep")
|
||||
from sglang.srt.arg_groups.overrides import resolved_view
|
||||
|
||||
self.assertEqual(server_args.moe_a2a_backend, "none") # pristine
|
||||
self.assertEqual(resolved_view(server_args).moe_a2a_backend, "deepep")
|
||||
self.assertTrue(server_args.enforce_shared_experts_fusion)
|
||||
|
||||
def test_waterfill_supports_deepep_low_latency_mode(self):
|
||||
|
||||
@@ -19,11 +19,10 @@ 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_declarations_to_server_args,
|
||||
apply_model_overrides,
|
||||
assert_flag_parity,
|
||||
collect_model_override_declarations,
|
||||
register_model_override,
|
||||
validate_declarations,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
_StaticFlags,
|
||||
@@ -207,7 +206,7 @@ class TestResolvedViewAndPasses(CustomTestCase):
|
||||
self.assertEqual(view.a, 10)
|
||||
self.assertEqual(view.b, 2)
|
||||
|
||||
def test_run_pass_appends_stash_and_dual_applies(self):
|
||||
def test_run_pass_appends_stash_and_stays_pristine(self):
|
||||
from sglang.srt.arg_groups.overrides import run_post_process_pass
|
||||
|
||||
live = SimpleNamespace(x=None, _resolved_overrides=[])
|
||||
@@ -216,11 +215,12 @@ class TestResolvedViewAndPasses(CustomTestCase):
|
||||
return {"x": "filled"} if view.x is None else {}
|
||||
|
||||
run_post_process_pass(live, _fill_x)
|
||||
self.assertEqual(live.x, "filled") # dual-applied in place
|
||||
self.assertIsNone(live.x) # never applied in place
|
||||
self.assertEqual(
|
||||
live._resolved_overrides, [(_fill_x.__qualname__, {"x": "filled"})]
|
||||
)
|
||||
run_post_process_pass(live, _fill_x) # now a no-op
|
||||
# the next invocation sees the declared value through the overlay
|
||||
run_post_process_pass(live, _fill_x)
|
||||
self.assertEqual(len(live._resolved_overrides), 1)
|
||||
|
||||
def test_run_pass_rejects_non_dict(self):
|
||||
@@ -415,7 +415,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_mistral_large3_forces_bfloat16(self):
|
||||
sa = self._construct("MistralLarge3ForCausalLM", "mistral")
|
||||
self.assertEqual(sa.dtype, "bfloat16") # dual-apply == legacy write
|
||||
self.assertEqual(sa.dtype, "bfloat16") # materialized at end of resolution
|
||||
self.assertIn(
|
||||
("MODEL_OVERRIDES['MistralLarge3ForCausalLM']", {"dtype": "bfloat16"}),
|
||||
sa._resolved_overrides,
|
||||
@@ -424,15 +424,15 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_pixtral_forces_bfloat16(self):
|
||||
sa = self._construct("PixtralForConditionalGeneration", "pixtral")
|
||||
self.assertEqual(sa.dtype, "bfloat16")
|
||||
self.assertEqual(sa.dtype, "bfloat16") # materialized
|
||||
self.assertEqual(self._publish(sa).dtype, "bfloat16")
|
||||
|
||||
def test_user_requested_dtype_is_still_overridden(self):
|
||||
# Legacy fidelity: the arch branch overwrote dtype unconditionally,
|
||||
# so the declaration must too. The pristine request survives only on
|
||||
# provenance (and, post-V3, as the un-overridden server_args field).
|
||||
# so the declaration must too. The pristine request survives on
|
||||
# provenance; the materialized field carries the override.
|
||||
sa = self._construct("MistralLarge3ForCausalLM", "mistral", dtype="float16")
|
||||
self.assertEqual(sa.dtype, "bfloat16")
|
||||
self.assertEqual(sa.dtype, "bfloat16") # materialized
|
||||
self.assertEqual(self._publish(sa).dtype, "bfloat16")
|
||||
|
||||
def test_control_arch_keeps_pristine_dtype(self):
|
||||
@@ -446,7 +446,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_minimax_m2_enables_tf32_matmul(self):
|
||||
sa = self._construct("MiniMaxM2ForCausalLM", "llama")
|
||||
self.assertTrue(sa.enable_tf32_matmul) # dual-apply == legacy write
|
||||
self.assertTrue(sa.enable_tf32_matmul) # materialized
|
||||
self.assertIn(
|
||||
("_minimax_m2_overrides", {"enable_tf32_matmul": True}),
|
||||
sa._resolved_overrides,
|
||||
@@ -491,7 +491,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
config_extra=config_extra,
|
||||
enable_hierarchical_cache=True,
|
||||
)
|
||||
# dual-apply == legacy writes
|
||||
# materialized at the end of resolution
|
||||
self.assertEqual(sa.swa_full_tokens_ratio, 1.0)
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory)
|
||||
flags = self._publish(sa)
|
||||
@@ -500,7 +500,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_gemma2_disables_hybrid_swa_memory(self):
|
||||
sa = self._construct("Gemma2ForCausalLM", "llama")
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory) # dual-apply == legacy
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized
|
||||
self.assertIn(
|
||||
("_gemma2_gemma3_overrides", {"disable_hybrid_swa_memory": True}),
|
||||
sa._resolved_overrides,
|
||||
@@ -509,7 +509,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_olmo2_disables_hybrid_swa_memory(self):
|
||||
sa = self._construct("Olmo2ForCausalLM", "llama")
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory)
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized
|
||||
self.assertTrue(self._publish(sa).disable_hybrid_swa_memory)
|
||||
|
||||
def test_exaone_conditional_on_sliding_window_pattern(self):
|
||||
@@ -520,7 +520,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
config_extra={"sliding_window_pattern": "LLLG"},
|
||||
attention_backend="fa3",
|
||||
)
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory)
|
||||
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized
|
||||
self.assertTrue(self._publish(sa).disable_hybrid_swa_memory)
|
||||
|
||||
def test_exaone_without_pattern_declares_nothing(self):
|
||||
@@ -543,7 +543,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
"llama",
|
||||
config_extra={"quantization_config": {"quant_method": "mxfp4"}},
|
||||
)
|
||||
self.assertEqual(sa.dtype, "bfloat16") # dual-apply == legacy
|
||||
self.assertEqual(sa.dtype, "bfloat16") # materialized
|
||||
self.assertEqual(self._publish(sa).dtype, "bfloat16")
|
||||
|
||||
def test_gpt_oss_without_mxfp4_keeps_pristine_dtype(self):
|
||||
@@ -569,7 +569,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
sa = self._construct("LlamaForCausalLM", "llama")
|
||||
expected = "flashinfer" if is_flashinfer_available() else "pytorch"
|
||||
self.assertEqual(sa.sampling_backend, expected)
|
||||
self.assertEqual(sa.sampling_backend, expected) # materialized
|
||||
self.assertIn(
|
||||
("_sampling_backend_default", {"sampling_backend": expected}),
|
||||
sa._resolved_overrides,
|
||||
@@ -587,20 +587,20 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
"LlamaForCausalLM", "llama", enable_deterministic_inference=True
|
||||
)
|
||||
# two pass writers chain: default fill, then the deterministic force —
|
||||
# last writer wins on the flags leaf and parity holds end-to-end.
|
||||
# last writer wins; materialization lands the end state on the fields.
|
||||
self.assertEqual(sa.sampling_backend, "pytorch")
|
||||
flags = self._publish(sa)
|
||||
self.assertEqual(flags.sampling_backend, "pytorch")
|
||||
# the deterministic attention fill declared a compatible backend and
|
||||
# the compatibility default-fill then had nothing to do
|
||||
self.assertIn(
|
||||
(
|
||||
"_deterministic_attention_backend",
|
||||
{"attention_backend": sa.attention_backend},
|
||||
),
|
||||
sa._resolved_overrides,
|
||||
)
|
||||
self.assertEqual(flags.attn.backend, sa.attention_backend)
|
||||
deterministic_fills = [
|
||||
decl["attention_backend"]
|
||||
for source, decl in sa._resolved_overrides
|
||||
if source == "_deterministic_attention_backend"
|
||||
]
|
||||
self.assertEqual(len(deterministic_fills), 1)
|
||||
self.assertEqual(sa.attention_backend, deterministic_fills[0])
|
||||
self.assertEqual(flags.attn.backend, deterministic_fills[0])
|
||||
|
||||
def test_deterministic_incompatible_backend_raises(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
@@ -631,13 +631,15 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_dllm_forces_flashinfer_with_cuda_graph(self):
|
||||
# CUDA path: cuda graph enabled by default -> dllm forces flashinfer.
|
||||
# A real dllm arch: the page pass now runs regardless of the radix
|
||||
# switch and builds DllmConfig for it.
|
||||
sa = self._construct(
|
||||
"LlamaForCausalLM",
|
||||
"SDARForCausalLM",
|
||||
"llama",
|
||||
dllm_algorithm="LowConfidence",
|
||||
disable_radix_cache=True,
|
||||
)
|
||||
self.assertEqual(sa.attention_backend, "flashinfer")
|
||||
self.assertEqual(sa.attention_backend, "flashinfer") # materialized
|
||||
self.assertIn(
|
||||
("_dllm_attention_backend", {"attention_backend": "flashinfer"}),
|
||||
sa._resolved_overrides,
|
||||
@@ -647,24 +649,35 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_attention_backend_leaf_materializes_end_state(self):
|
||||
# The default-fill pass declares the platform-selected backend; the
|
||||
# leaf must equal the final server_args value (publish parity).
|
||||
# leaf must equal the last declared value while the server_args field
|
||||
# stays pristine (dual-apply retired).
|
||||
sa = self._construct("LlamaForCausalLM", "llama")
|
||||
declared = {f for _s, d in sa._resolved_overrides for f in d}
|
||||
self.assertIn("attention_backend", declared) # default fill declared
|
||||
self.assertEqual(self._publish(sa).attn.backend, sa.attention_backend)
|
||||
declared_values = [
|
||||
d["attention_backend"]
|
||||
for _s, d in sa._resolved_overrides
|
||||
if "attention_backend" in d
|
||||
]
|
||||
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])
|
||||
|
||||
def test_runner_side_adjustment_can_refresh_declaration(self):
|
||||
from sglang.srt.arg_groups.overrides import refresh_declared_fields
|
||||
def test_post_materialize_pass_writes_through(self):
|
||||
from sglang.srt.arg_groups.overrides import run_post_process_pass
|
||||
|
||||
# A pass invoked after materialization (a post-init slot, like the
|
||||
# legacy runner-side adjustments) declares AND writes through, so
|
||||
# field readers and the publish see the same end state.
|
||||
sa = self._construct("LlamaForCausalLM", "llama")
|
||||
declared = {f for _s, d in sa._resolved_overrides for f in d}
|
||||
self.assertIn("attention_backend", declared)
|
||||
# Simulate a legacy runner-side overwrite between collection and publish
|
||||
# (model_specific_adjustment forces attention_backend for HRM-Text).
|
||||
sa.attention_backend = "fa3" if sa.attention_backend != "fa3" else "triton"
|
||||
with self.assertRaises(AssertionError):
|
||||
self._publish(sa) # stale declaration breaks parity
|
||||
refresh_declared_fields(sa, ("attention_backend",))
|
||||
resolved_before = sa.attention_backend
|
||||
|
||||
def _force_triton(view):
|
||||
if view.attention_backend != "triton":
|
||||
return {"attention_backend": "triton"}
|
||||
return {}
|
||||
|
||||
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)
|
||||
|
||||
def test_attention_backend_user_choice_declares_nothing_extra(self):
|
||||
@@ -798,9 +811,41 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
return_value=SimpleNamespace(block_size=32),
|
||||
):
|
||||
self.assertEqual(_view() and _dllm_page_size(_view()), {"page_size": 32})
|
||||
self.assertEqual(_dllm_page_size(_view(page_size=64)), {}) # aligned
|
||||
# aligned but larger than the block: the scheduler-init fallback
|
||||
# (folded into this pass) still caps the page at the block size
|
||||
self.assertEqual(_dllm_page_size(_view(page_size=64)), {"page_size": 32})
|
||||
self.assertEqual(_dllm_page_size(_view(page_size=32)), {}) # equal
|
||||
# radix disabled skips the alignment fill but keeps the cap
|
||||
self.assertEqual(_dllm_page_size(_view(disable_radix_cache=True)), {})
|
||||
self.assertEqual(
|
||||
_dllm_page_size(_view(disable_radix_cache=True, page_size=64)),
|
||||
{"page_size": 32},
|
||||
)
|
||||
self.assertEqual(_dllm_page_size(_view(dllm_algorithm=None)), {})
|
||||
self.assertEqual(_dllm_page_size(_view(disable_radix_cache=True)), {})
|
||||
|
||||
def test_declaration_overlay_mechanics(self):
|
||||
from sglang.srt.arg_groups.overrides import run_post_process_pass
|
||||
|
||||
live = SimpleNamespace(x="user", y=None, _resolved_overrides=[])
|
||||
|
||||
def _resolve_x(view):
|
||||
return {"x": "resolved"} if view.x == "user" else {}
|
||||
|
||||
def _read_x(view):
|
||||
return {"y": view.x}
|
||||
|
||||
run_post_process_pass(live, _resolve_x)
|
||||
# declaration recorded, but server_args stays pristine
|
||||
self.assertEqual(
|
||||
live._resolved_overrides, [(_resolve_x.__qualname__, {"x": "resolved"})]
|
||||
)
|
||||
self.assertEqual(live.x, "user")
|
||||
# a later pass sees the resolved value through the view overlay
|
||||
run_post_process_pass(live, _read_x)
|
||||
self.assertEqual(
|
||||
live._resolved_overrides[-1], (_read_x.__qualname__, {"y": "resolved"})
|
||||
)
|
||||
self.assertIsNone(live.y) # never applied in place
|
||||
|
||||
def test_overlap_disable_passes(self):
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
@@ -1531,9 +1576,12 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
|
||||
def test_page_size_leaf_materializes_end_state(self):
|
||||
sa = self._construct("LlamaForCausalLM", "llama")
|
||||
declared = {f for _s, d in sa._resolved_overrides for f in d}
|
||||
self.assertIn("page_size", declared) # default fill declared
|
||||
self.assertEqual(self._publish(sa).page_size, sa.page_size)
|
||||
declared_values = [
|
||||
d["page_size"] for _s, d in sa._resolved_overrides if "page_size" in d
|
||||
]
|
||||
self.assertTrue(declared_values) # default fill declared
|
||||
self.assertEqual(sa.page_size, declared_values[-1]) # materialized
|
||||
self.assertEqual(self._publish(sa).page_size, declared_values[-1])
|
||||
|
||||
def test_qwen3_5_hybrid_coupled_declaration(self):
|
||||
from sglang.srt.arg_groups.overrides import _qwen3_5_hybrid_overrides
|
||||
@@ -1544,7 +1592,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
_get_default_attn_backend=lambda **_: default_backend,
|
||||
use_mla_backend=lambda: False,
|
||||
get_model_config=lambda: None,
|
||||
enable_mamba_extra_buffer=lambda: False,
|
||||
mamba_radix_cache_strategy="auto",
|
||||
disable_radix_cache=False,
|
||||
speculative_algorithm=None,
|
||||
)
|
||||
@@ -1571,6 +1619,24 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
),
|
||||
{},
|
||||
)
|
||||
# the mamba pass ran before this dispatch and stashed the
|
||||
# extra-buffer strategy: the callable must see it through the
|
||||
# view (SM100 hybrid keeps trtllm_mha + page 64)
|
||||
self.assertEqual(
|
||||
_qwen3_5_hybrid_overrides(
|
||||
_args(
|
||||
"trtllm_mha",
|
||||
_resolved_overrides=[
|
||||
(
|
||||
"_mamba_radix_cache_declarations",
|
||||
{"mamba_radix_cache_strategy": "extra_buffer"},
|
||||
)
|
||||
],
|
||||
),
|
||||
None,
|
||||
),
|
||||
{"attention_backend": "trtllm_mha", "page_size": 64},
|
||||
)
|
||||
with patch.object(overrides_module, "is_sm100_supported", return_value=False):
|
||||
self.assertEqual(_qwen3_5_hybrid_overrides(_args("fa3"), None), {})
|
||||
|
||||
@@ -1738,7 +1804,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
self.assertEqual(
|
||||
_intel_xpu_page_constraint(
|
||||
_view(
|
||||
get_attention_backends=lambda: (None, "intel_xpu"),
|
||||
decode_attention_backend="intel_xpu",
|
||||
use_mla_backend=lambda: False,
|
||||
)
|
||||
),
|
||||
@@ -1747,7 +1813,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
self.assertEqual(
|
||||
_intel_xpu_page_constraint(
|
||||
_view(
|
||||
get_attention_backends=lambda: (None, "intel_xpu"),
|
||||
decode_attention_backend="intel_xpu",
|
||||
use_mla_backend=lambda: True,
|
||||
page_size=16, # MLA decode accepts 16
|
||||
)
|
||||
@@ -2089,22 +2155,22 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
self.assertEqual(_step3p_overrides(_args(), None), {})
|
||||
|
||||
|
||||
class TestDualApplyParity(CustomTestCase):
|
||||
def test_dual_apply_replays_and_parity_holds(self):
|
||||
class TestDeclarationValidation(CustomTestCase):
|
||||
def test_declarations_never_mutate_server_args(self):
|
||||
flags, args = _FakeFlags(), _FakeArgs()
|
||||
declarations = [("src", {"resolved_by_model": "dsv4", "also_resolved": 7})]
|
||||
apply_model_overrides(flags, args, declarations)
|
||||
apply_declarations_to_server_args(args, declarations)
|
||||
self.assertEqual(args.resolved_by_model, "dsv4")
|
||||
self.assertEqual(args.also_resolved, 7)
|
||||
assert_flag_parity(flags, args, ["resolved_by_model", "also_resolved"])
|
||||
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)
|
||||
self.assertEqual(args.resolved_by_model, _FakeArgs.resolved_by_model)
|
||||
self.assertEqual(args.also_resolved, _FakeArgs.also_resolved)
|
||||
|
||||
def test_parity_detects_drift(self):
|
||||
flags, args = _FakeFlags(), _FakeArgs()
|
||||
apply_model_overrides(flags, args, [("src", {"resolved_by_model": "x"})])
|
||||
# dual-apply skipped -> server_args still pristine -> drift is caught
|
||||
with self.assertRaises(AssertionError):
|
||||
assert_flag_parity(flags, args, ["resolved_by_model"])
|
||||
def test_validation_rejects_unknown_fields(self):
|
||||
args = _FakeArgs()
|
||||
with self.assertRaises(ValueError):
|
||||
validate_declarations(args, [("src", {"nope": 1})])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -339,15 +339,6 @@ class TestRuntimeResolutionStages(_IsolatedServerArgs):
|
||||
self.assertEqual(get_flags().sampling_backend, "pytorch")
|
||||
self.assertEqual(get_flags().page_size, 64) # earlier stage survives
|
||||
|
||||
def test_record_parity_failure_rolls_back(self):
|
||||
self._publish(page_size=1)
|
||||
flags_before = get_flags()
|
||||
with self.assertRaises(AssertionError):
|
||||
# declared value diverges from the live server_args (no dual-apply)
|
||||
get_context().record_runtime_overrides([("bad", {"page_size": 64})])
|
||||
self.assertIs(get_flags(), flags_before) # previous flags intact
|
||||
self.assertEqual(get_context()._runtime_overrides, []) # rolled back
|
||||
|
||||
def test_record_whitelist_violation_rolls_back(self):
|
||||
self._publish()
|
||||
with self.assertRaises(ValueError):
|
||||
@@ -366,13 +357,15 @@ class TestRuntimeResolutionStages(_IsolatedServerArgs):
|
||||
finally:
|
||||
reset_context()
|
||||
|
||||
def test_declare_load_time_override_dual_applies_and_records(self):
|
||||
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})
|
||||
self.assertEqual(args.page_size, 64) # dual-applied onto server_args
|
||||
self.assertEqual(get_flags().page_size, 64) # resolved into the leaf
|
||||
# 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})],
|
||||
@@ -415,6 +408,21 @@ class TestRuntimeResolutionStages(_IsolatedServerArgs):
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user