[Config] msgspec.Struct for the config tier (#38753)

This commit is contained in:
Cheng Wan
2026-09-09 19:41:19 -07:00
committed by GitHub
parent 106cc561f9
commit f1a512c51c
44 changed files with 328 additions and 210 deletions
+25 -29
View File
@@ -47,7 +47,6 @@ test-only ``override(**kw)``.
from __future__ import annotations
import dataclasses
import functools
import logging
import math
@@ -56,6 +55,8 @@ import sys
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, Dict, Optional
import msgspec
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
@@ -461,28 +462,29 @@ def _install_derived_widths() -> None:
_install_derived_widths()
class _FlagGroupBase:
class _FlagGroupBase(msgspec.Struct):
"""Shared flag-group behavior: typo-safe writes + transactional ``override()``.
Groups are plain dataclasses; ``__dataclass_fields__`` is the single source
of truth for which leaves exist, so a mistyped name fails loudly instead of
creating a stray attribute.
``__struct_fields__`` is the single source of truth for which leaves exist,
so a mistyped name fails loudly instead of creating a stray attribute. The
write goes through ``super().__setattr__``: a ``Struct`` keeps its fields in
its own layout, so ``object.__setattr__`` does not reach them.
"""
def __setattr__(self, name: str, value: Any) -> None:
if name not in type(self).__dataclass_fields__:
if name not in type(self).__struct_fields__:
raise AttributeError(
f"{type(self).__name__} has no flag '{name}' (leaves are "
"declared as dataclass fields; check for typos)"
"declared as struct fields; check for typos)"
)
object.__setattr__(self, name, value)
super().__setattr__(name, value)
@contextmanager
def override(self, **kwargs):
"""Temporarily force flag values, restoring on exit. Transactional
(keys validated before any write) — the test-only injection
primitive."""
fields = type(self).__dataclass_fields__
fields = type(self).__struct_fields__
unknown = set(kwargs) - set(fields)
if unknown:
raise ValueError(
@@ -490,15 +492,14 @@ class _FlagGroupBase:
)
saved = {name: getattr(self, name) for name in kwargs}
for name, value in kwargs.items():
object.__setattr__(self, name, value)
setattr(self, name, value)
try:
yield self
finally:
for name, value in saved.items():
object.__setattr__(self, name, value)
setattr(self, name, value)
@dataclasses.dataclass
class CaptureFlags(_FlagGroupBase):
"""Capture-time flags; never frozen (written during cuda-graph capture)."""
@@ -512,7 +513,6 @@ class CaptureFlags(_FlagGroupBase):
disable_dispose_tensor: bool = False
@dataclasses.dataclass
class MoeFlags(_FlagGroupBase):
"""MoE runtime flags, materialized by ``initialize_moe_config`` (scheduler
init, after distributed setup). ``a2a_backend`` / ``runner_backend`` /
@@ -552,7 +552,6 @@ class MoeFlags(_FlagGroupBase):
speculative_context: bool = False
@dataclasses.dataclass
class DpFlags(_FlagGroupBase):
"""DP-attention runtime flags, materialized by ``initialize_dp_attention``
(after distributed setup; reads the model config). Topology values
@@ -576,7 +575,6 @@ class DpFlags(_FlagGroupBase):
buffer_device: Any = None
@dataclasses.dataclass
class SpFlags(_FlagGroupBase):
"""LayerNorm sequence-parallelism flags, materialized by
``initialize_layernorm_sp`` (after distributed setup; reads the model
@@ -585,7 +583,6 @@ class SpFlags(_FlagGroupBase):
enabled: bool = False
@dataclasses.dataclass
class Flags(_FlagGroupBase):
"""Root of the runtime-flags tier.
@@ -595,13 +592,12 @@ class Flags(_FlagGroupBase):
by lifecycle (``capture``) or subsystem (``moe`` / ``dp`` / ``sp``).
"""
capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags)
moe: MoeFlags = dataclasses.field(default_factory=MoeFlags)
dp: DpFlags = dataclasses.field(default_factory=DpFlags)
sp: SpFlags = dataclasses.field(default_factory=SpFlags)
capture: CaptureFlags = msgspec.field(default_factory=CaptureFlags)
moe: MoeFlags = msgspec.field(default_factory=MoeFlags)
dp: DpFlags = msgspec.field(default_factory=DpFlags)
sp: SpFlags = msgspec.field(default_factory=SpFlags)
@dataclasses.dataclass
class Resources(_FlagGroupBase):
"""Process-level resource handles: named slots with one reset lifecycle,
scoped test injection via ``override()``, and the creation/publish
@@ -616,16 +612,16 @@ class Resources(_FlagGroupBase):
expert_distribution_recorder: Any = None
expert_location_metadata: Any = None
# LPLB: layer_id -> solver.
lplb_solvers: dict = dataclasses.field(default_factory=dict)
lplb_solvers: dict = msgspec.field(default_factory=dict)
# Named side streams (see RuntimeContext.get_stream): name -> stream.
streams: dict = dataclasses.field(default_factory=dict)
streams: dict = msgspec.field(default_factory=dict)
# Named persistent buffers (see RuntimeContext.get_buffer): name -> tensor.
# Accessors with bespoke semantics (grow-only, per-device keys) manage
# their entries directly.
buffers: dict = dataclasses.field(default_factory=dict)
buffers: dict = msgspec.field(default_factory=dict)
# Persistent reusable CUDA events for non-EP DP TBO, keyed by
# (kind, subbatch) — see dp_attention._tbo_event for why reuse matters.
tbo_event_pool: dict = dataclasses.field(default_factory=dict)
tbo_event_pool: dict = msgspec.field(default_factory=dict)
# State capturers (installed by their subsystems when capture is on).
indexer_capturer: Any = None
experts_capturer: Any = None
@@ -1290,7 +1286,7 @@ class _ServerArgsOverride:
# Underscore names seed private property caches (the strict guard
# exempts them); everything else must be a real config field.
unknown = {name for name in self._fields if not name.startswith("_")} - set(
type(server_args).__dataclass_fields__
type(server_args).__struct_fields__
)
if unknown:
raise ValueError(
@@ -1303,7 +1299,7 @@ class _ServerArgsOverride:
# real field, and seeding it as a raw attribute would leave the earlier
# declaration authoritative, so `resolution_result` and the bag would
# both keep answering the pre-override value.
fields = set(type(server_args).__dataclass_fields__)
fields = set(type(server_args).__struct_fields__)
declared = {n: v for n, v in self._fields.items() if n in fields}
if declared:
declare_resolution(server_args, "override_server_args", **declared)
@@ -1311,7 +1307,7 @@ class _ServerArgsOverride:
# and friends), which are not configuration and never were.
seeds = {n: v for n, v in self._fields.items() if n not in fields}
for name, value in seeds.items():
object.__setattr__(server_args, name, value)
msgspec.Struct.__setattr__(server_args, name, value)
ctx.set_server_args(server_args)
self._installed = True
return server_args
@@ -1680,7 +1676,7 @@ def set_global_dwdp_manager(manager: Any) -> None:
def _group_leaves(group: _FlagGroupBase) -> dict[str, Any]:
"""The leaf values of a flag group, recursively."""
leaves: dict[str, Any] = {}
for name in type(group).__dataclass_fields__:
for name in type(group).__struct_fields__:
value = getattr(group, name)
if isinstance(value, _FlagGroupBase):
leaves[name] = _group_leaves(value)