[refactor] Wire the config resolution pipeline (dispatch, stash, dual-apply, publish) (stack 6/15) (#30068)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
c3d751b231
commit
df6491d80c
@@ -85,7 +85,12 @@ class Arg:
|
|||||||
@functools.lru_cache(maxsize=None)
|
@functools.lru_cache(maxsize=None)
|
||||||
def model_overridable_fields(cls) -> frozenset:
|
def model_overridable_fields(cls) -> frozenset:
|
||||||
"""Names of ``cls`` dataclass fields whose ``Arg`` metadata declares
|
"""Names of ``cls`` dataclass fields whose ``Arg`` metadata declares
|
||||||
``model_overridable=True`` — the whitelist for model-override resolution."""
|
``model_overridable=True`` — the whitelist for model-override resolution.
|
||||||
|
|
||||||
|
Non-dataclass types (e.g. mock config objects in tests) have no Arg
|
||||||
|
metadata and yield an empty whitelist."""
|
||||||
|
if not dataclasses.is_dataclass(cls):
|
||||||
|
return frozenset()
|
||||||
hints = get_type_hints(cls, include_extras=True)
|
hints = get_type_hints(cls, include_extras=True)
|
||||||
names = set()
|
names = set()
|
||||||
for field in dataclasses.fields(cls):
|
for field in dataclasses.fields(cls):
|
||||||
|
|||||||
@@ -172,7 +172,23 @@ def apply_declarations_to_server_args(
|
|||||||
|
|
||||||
Retired per field once that field's readers have all flipped to the flags
|
Retired per field once that field's readers have all flipped to the flags
|
||||||
tier (at which point the server_args field returns to pristine).
|
tier (at which point the server_args field returns to pristine).
|
||||||
|
|
||||||
|
Validates against the same whitelist as the publish gate BEFORE any write:
|
||||||
|
a registry typo or a not-yet-resolvable field must fail fast here, not
|
||||||
|
mutate ``server_args`` and only be rejected at publish time.
|
||||||
"""
|
"""
|
||||||
|
# Non-dataclass fixtures carry no Arg metadata (mirrors the
|
||||||
|
# model_overridable_fields escape); only real ServerArgs is validated.
|
||||||
|
if dataclasses.is_dataclass(type(server_args)):
|
||||||
|
whitelist = model_overridable_fields(type(server_args))
|
||||||
|
for source, decl in list(declarations) + list(terminal):
|
||||||
|
unknown = set(decl) - whitelist
|
||||||
|
if unknown:
|
||||||
|
raise ValueError(
|
||||||
|
f"{source}: {sorted(unknown)} not model-overridable; the "
|
||||||
|
"transition dual-apply refuses fields the publish gate "
|
||||||
|
"would reject."
|
||||||
|
)
|
||||||
for _source, decl in list(declarations) + list(terminal):
|
for _source, decl in list(declarations) + list(terminal):
|
||||||
for field, value in decl.items():
|
for field, value in decl.items():
|
||||||
setattr(server_args, field, value)
|
setattr(server_args, field, value)
|
||||||
|
|||||||
@@ -359,9 +359,40 @@ class RuntimeContext:
|
|||||||
Overwrite-allowed: a re-publish replaces the slot (test kits re-publish
|
Overwrite-allowed: a re-publish replaces the slot (test kits re-publish
|
||||||
per test; production ordering discipline lives at the call-sites, e.g.
|
per test; production ordering discipline lives at the call-sites, e.g.
|
||||||
the draft-worker guard in ``ModelRunner.__init__``).
|
the draft-worker guard in ``ModelRunner.__init__``).
|
||||||
|
|
||||||
|
Publishing also resolves the stashed model-override declarations
|
||||||
|
into the flags tier (skipped for objects without the stash — dummy /
|
||||||
|
"none" fixture ServerArgs and test-kit mocks never compute it).
|
||||||
|
Resolution runs first: if it fails, the previous publish stays intact.
|
||||||
"""
|
"""
|
||||||
|
self._resolve_flags(server_args)
|
||||||
self._server_args = server_args
|
self._server_args = server_args
|
||||||
|
|
||||||
|
def _resolve_flags(self, server_args: ServerArgs) -> None:
|
||||||
|
declarations = getattr(server_args, "_resolved_overrides", None)
|
||||||
|
if declarations is None:
|
||||||
|
return
|
||||||
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
apply_model_overrides,
|
||||||
|
assert_flag_parity,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Resolve into a fresh container and only install it once everything
|
||||||
|
# passed: a failed resolution (gate validation or the parity assert)
|
||||||
|
# must not leave the process-global flags half-written for callers
|
||||||
|
# that catch the error or republish (same install-fresh semantics as
|
||||||
|
# reset_context()).
|
||||||
|
flags = Flags()
|
||||||
|
apply_model_overrides(flags, server_args, declarations)
|
||||||
|
# Transition-period drift guard: dual-apply keeps the declared fields
|
||||||
|
# on server_args byte-identical to the resolved flag leaves.
|
||||||
|
assert_flag_parity(
|
||||||
|
flags,
|
||||||
|
server_args,
|
||||||
|
{field for _source, decl in declarations for field in decl},
|
||||||
|
)
|
||||||
|
self.flags = flags
|
||||||
|
|
||||||
|
|
||||||
_PARALLEL = ParallelContext()
|
_PARALLEL = ParallelContext()
|
||||||
_CONTEXT = RuntimeContext(parallel=_PARALLEL)
|
_CONTEXT = RuntimeContext(parallel=_PARALLEL)
|
||||||
|
|||||||
@@ -2569,6 +2569,12 @@ class ServerArgs:
|
|||||||
Orchestrates the handling of various server arguments, ensuring proper configuration and validation.
|
Orchestrates the handling of various server arguments, ensuring proper configuration and validation.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# Declaration stash for the override/post-process passes. Set before any
|
||||||
|
# short-circuit (none/dummy model paths) so run_post_process_pass and
|
||||||
|
# direct handler invocations can rely on it even when
|
||||||
|
# _handle_model_specific_adjustments never runs.
|
||||||
|
self._resolved_overrides = []
|
||||||
|
|
||||||
self._maybe_download_model_for_runai()
|
self._maybe_download_model_for_runai()
|
||||||
|
|
||||||
# Normalize load balancing defaults early (before dummy-model short-circuit).
|
# Normalize load balancing defaults early (before dummy-model short-circuit).
|
||||||
@@ -3709,6 +3715,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
self.uses_mamba_radix_cache = False
|
self.uses_mamba_radix_cache = False
|
||||||
if parse_connector_type(self.model_path) == ConnectorType.INSTANCE:
|
if parse_connector_type(self.model_path) == ConnectorType.INSTANCE:
|
||||||
|
self._resolved_overrides = []
|
||||||
return
|
return
|
||||||
|
|
||||||
hf_config = self.get_model_config().hf_config
|
hf_config = self.get_model_config().hf_config
|
||||||
@@ -3718,6 +3725,22 @@ class ServerArgs:
|
|||||||
if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
|
if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
|
||||||
self._handle_mamba_radix_cache(model_arch=model_arch)
|
self._handle_mamba_radix_cache(model_arch=model_arch)
|
||||||
|
|
||||||
|
# Collect the declarative model overrides (registry) on the
|
||||||
|
# pristine config and stash them for publish-time flags resolution.
|
||||||
|
# Transition dual-apply: the same declarations are applied to
|
||||||
|
# server_args right here, byte-identical to the imperative arch
|
||||||
|
# branches this dispatch gradually replaces (dual-apply is retired
|
||||||
|
# per field once that field's readers migrate to the flags tier).
|
||||||
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
apply_declarations_to_server_args,
|
||||||
|
collect_model_override_declarations,
|
||||||
|
)
|
||||||
|
|
||||||
|
self._resolved_overrides = collect_model_override_declarations(
|
||||||
|
model_arch, self, hf_config
|
||||||
|
)
|
||||||
|
apply_declarations_to_server_args(self, self._resolved_overrides)
|
||||||
|
|
||||||
if model_arch in [
|
if model_arch in [
|
||||||
"MistralLarge3ForCausalLM",
|
"MistralLarge3ForCausalLM",
|
||||||
"PixtralForConditionalGeneration",
|
"PixtralForConditionalGeneration",
|
||||||
|
|||||||
@@ -21,7 +21,12 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
collect_model_override_declarations,
|
collect_model_override_declarations,
|
||||||
register_model_override,
|
register_model_override,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import _StaticFlags
|
from sglang.srt.runtime_context import (
|
||||||
|
_StaticFlags,
|
||||||
|
get_context,
|
||||||
|
get_server_args,
|
||||||
|
reset_context,
|
||||||
|
)
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
|
||||||
@@ -50,6 +55,9 @@ class TestModelOverridableWhitelist(CustomTestCase):
|
|||||||
|
|
||||||
self.assertEqual(model_overridable_fields(ServerArgs), frozenset())
|
self.assertEqual(model_overridable_fields(ServerArgs), frozenset())
|
||||||
|
|
||||||
|
def test_non_dataclass_yields_empty_whitelist(self):
|
||||||
|
self.assertEqual(model_overridable_fields(SimpleNamespace), frozenset())
|
||||||
|
|
||||||
|
|
||||||
class _IsolatedRegistry(CustomTestCase):
|
class _IsolatedRegistry(CustomTestCase):
|
||||||
"""Run each test against empty registries (they are process-global)."""
|
"""Run each test against empty registries (they are process-global)."""
|
||||||
@@ -199,6 +207,59 @@ class TestApplyModelOverridesGate(CustomTestCase):
|
|||||||
self.assertEqual(flags.resolved_by_model, "unset") # flat leaf untouched
|
self.assertEqual(flags.resolved_by_model, "unset") # flat leaf untouched
|
||||||
|
|
||||||
|
|
||||||
|
class _IsolatedPublish(CustomTestCase):
|
||||||
|
"""Publishing writes the process-global context; save/restore around it."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
super().setUp()
|
||||||
|
self._saved_server_args = get_context()._server_args
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
reset_context()
|
||||||
|
if self._saved_server_args is not None:
|
||||||
|
get_context()._server_args = self._saved_server_args
|
||||||
|
super().tearDown()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
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."""
|
||||||
|
|
||||||
|
def test_dummy_fixture_has_empty_stash_and_publishes_cleanly(self):
|
||||||
|
from sglang.srt.server_args import (
|
||||||
|
ServerArgs,
|
||||||
|
set_global_server_args_for_scheduler,
|
||||||
|
)
|
||||||
|
|
||||||
|
sa = ServerArgs(model_path="dummy") # __post_init__ early-returns
|
||||||
|
# The stash is created before the dummy short-circuit and stays empty.
|
||||||
|
self.assertEqual(sa._resolved_overrides, [])
|
||||||
|
set_global_server_args_for_scheduler(sa)
|
||||||
|
self.assertIs(get_server_args(), sa)
|
||||||
|
|
||||||
|
def test_empty_stash_publish_runs_gate_as_noop(self):
|
||||||
|
sa = _NoOverridableArgs()
|
||||||
|
sa._resolved_overrides = []
|
||||||
|
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 TestDualApplyParity(CustomTestCase):
|
class TestDualApplyParity(CustomTestCase):
|
||||||
def test_dual_apply_replays_and_parity_holds(self):
|
def test_dual_apply_replays_and_parity_holds(self):
|
||||||
flags, args = _FakeFlags(), _FakeArgs()
|
flags, args = _FakeFlags(), _FakeArgs()
|
||||||
|
|||||||
Reference in New Issue
Block a user