config: ServerArgs holds the raw input (#36255)

This commit is contained in:
Cheng Wan
2026-08-26 05:14:05 -07:00
committed by GitHub
parent 27c36368b6
commit 413df1f8db
34 changed files with 1201 additions and 383 deletions
+54 -25
View File
@@ -22,12 +22,18 @@ flags/resources/forward tiers.
## Config: publish + namespace bags ## Config: publish + namespace bags
**`ServerArgs` is a pristine seed. Business code never reads it for decisions — **`ServerArgs` holds the raw input and nothing else. Resolution writes no field:
resolved configuration lives in the namespace bags.** it declares, and the declarations are what the namespace bags are projected from.
Business code never reads the record for a decision — and after this cut, a field
read there answers with what the operator typed, not with what resolution
decided.**
- Every publishing process entry calls `publish(server_args, role=...)` - Every publishing process entry calls `publish(server_args, role=...)`
(`run_scheduler_process`, the Ray `SchedulerActor`, the DP controller, tokenizer, (`run_scheduler_process`, the Ray `SchedulerActor`, the DP controller, tokenizer,
detokenizer, encoder, weight-cache daemon, ...); the roles are enumerated once, detokenizer, encoder, weight-cache daemon, the multi-tokenizer worker, the
spawned encoder TP/DP workers, the benchmark work functions, ...); constructors
do not publish — `ModelRunner`, `TokenizerManager` and `MMEncoder` call
`assert_published` and fail loudly if an entry forgot. The roles are enumerated once,
as the keys of `ROLE_NAMESPACE_SETS` — there is no `launcher` role, the launch as the keys of `ROLE_NAMESPACE_SETS` — there is no `launcher` role, the launch
path publishes as `tokenizer`. The remaining non-publisher is path publishes as `tokenizer`. The remaining non-publisher is
`run_multi_detokenizer_router_process`: it *is* handed a `ServerArgs`, and uses `run_multi_detokenizer_router_process`: it *is* handed a `ServerArgs`, and uses
@@ -82,13 +88,15 @@ resolved configuration lives in the namespace bags.**
the target runner. the target runner.
- **Late launcher-stage resolution (pre-publish)**: a few rules cannot run inside - **Late launcher-stage resolution (pre-publish)**: a few rules cannot run inside
`__post_init__` — LoRA normalization, and the auto-parser detection that needs a `__post_init__` — LoRA normalization, and the auto-parser detection that needs a
tokenizer/chat-template load. They are resolution, not mutation, and they write tokenizer/chat-template load. They are resolution, not mutation, and they
**in place** via `arg_groups.overrides.declare_late_resolution(server_args, **declare** via `arg_groups.overrides.declare_late_resolution(server_args,
source, **fields)`, which refuses the published instance. In place is the point: source, **fields)`, which refuses the published instance. The declaration lands
every holder of that object must see the resolved value — the HTTP server, the in the stash on that very object, so every holder of it carries the decision —
multi-tokenizer workers it is serialized for, the schedulers it forks. Returning a the HTTP server, the multi-tokenizer workers it is serialized for, the
variant here is a bug: the launcher rebinds its local and everyone else keeps the schedulers it forks — and each of them publishes bags projected from it. The
unresolved object. fields stay the operator's input; `resolution_result(sa, field)` and the bags
are what answer for the decision. Returning a variant here is a bug: the
launcher rebinds its local and everyone else keeps the unresolved object.
- **A value another runner / worker owns is a constructor argument, not a config - **A value another runner / worker owns is a constructor argument, not a config
copy.** The draft worker's `context_length`, load format and attention backend copy.** The draft worker's `context_length`, load format and attention backend
travel as arguments to `TpModelWorker` / `ModelRunner` and live on the runner travel as arguments to `TpModelWorker` / `ModelRunner` and live on the runner
@@ -99,8 +107,8 @@ resolved configuration lives in the namespace bags.**
**Why a bag override cannot stand in for late resolution or per-runner **Why a bag override cannot stand in for late resolution or per-runner
construction.** The bags are projected at construction.** The bags are projected at
publish *from the instance's fields*, so anything the runtime must read has to be on publish *from the declarations over the instance's raw fields*, so anything the
the instance before publish — an override afterwards puts instance and bags back out runtime must read has to be declared before publish — an override afterwards puts instance and bags back out
of agreement, and whole-object readers (`ModelConfig.from_server_args`, of agreement, and whole-object readers (`ModelConfig.from_server_args`,
`build_load_config`, `MMEncoder`'s own `self.server_args.X`) never see it. And bags do `build_load_config`, `MMEncoder`'s own `self.server_args.X`) never see it. And bags do
not cross a process boundary: a child publishes from the object it receives and not cross a process boundary: a child publishes from the object it receives and
@@ -114,7 +122,8 @@ bag to override at all.
- **Per-runner values** — there is no per-runner `ServerArgs` any more. The - **Per-runner values** — there is no per-runner `ServerArgs` any more. The
draft-worker config copy is gone: every worker (`TpModelWorker`, the draft draft-worker config copy is gone: every worker (`TpModelWorker`, the draft
workers in `speculative/`) is handed the *same* instance the process published, workers in `speculative/`) is handed the *same* instance the process published,
so `self.server_args.X` and the bag leaf agree **at publish** — a so a bag leaf is the decision and `self.server_args.X` is the operator's
input — a
post-publish `override` moves only the bag, which is exactly why a field that post-publish `override` moves only the bag, which is exactly why a field that
is process-wide config (`attention_backend`, `skip_tokenizer_init`, is process-wide config (`attention_backend`, `skip_tokenizer_init`,
`kv_cache_dtype`) reads from the bags like any other, and why a residual `kv_cache_dtype`) reads from the bags like any other, and why a residual
@@ -320,9 +329,9 @@ what sits beside it is residue, not a family — and not for one single reason:
without publishing has to keep patching the factory (or publish itself); without publishing has to keep patching the factory (or publish itself);
- `MMEncoder` publishes the very instance it is handed (`publish(server_args, - `MMEncoder` publishes the very instance it is handed (`publish(server_args,
role="encoder")`) and takes its per-worker device as a separate `gpu_id` role="encoder")`) and takes its per-worker device as a separate `gpu_id`
argument, so its `self.server_args` reads and the bag agree today. They are on argument. Its `self.server_args` reads are on this list as a construction-path
this list as a construction-path convention rather than a semantic exception — convention, and the residual is real: they answer with the raw input, so a leaf
and the residual is real: a post-publish `override` would not reach them. resolution decided and a post-publish `override` both pass them by.
Their tests are not one story: a `GrammarManager` built standalone turns the Their tests are not one story: a `GrammarManager` built standalone turns the
factory's bag read into "config namespace not published" unless the test patches factory's bag read into "config namespace not published" unless the test patches
@@ -357,14 +366,30 @@ if you do it, say so in the test.
### Mid-resolution reads (inside the pipeline only) ### Mid-resolution reads (inside the pipeline only)
Resolution itself still runs in `__post_init__`: handlers and hooks read the Resolution runs in `__post_init__` and **writes nothing onto the record**: a
in-flight state through `resolved_view(server_args)` / `self._resolved()`, fields are handler declares (`self._declare` / `declare_resolution`), the declaration goes
read-only during resolution, and declarations materialize once at the very end of into the stash, and the fields keep what the caller passed. So a mid-resolution
`__post_init__` (gate order, last writer wins) — *then* `publish` snapshots the read of a field answers with the *raw input* — every reader in the pipeline goes
resolved values into the bags. `resolved_view` is pipeline-internal through a view instead:
(`server_args.py` / `arg_groups/`, plus helpers the pipeline itself invokes
mid-resolution, e.g. `adaptive_spec_params`); do not introduce new - `resolving_view(server_args)` / `self._resolved()` — the live view (walks the
out-of-pipeline call sites. stash per read). This is what handlers and hooks bind, conventionally as
`cfg = resolving_view(self)` at the top of the handler.
- `resolved_view(server_args)` — snapshots the overlay when built, which is what
a post-process pass wants: it reads the state at *its* slot.
`test_resolution_reads_the_declarations` pins direct field reads at zero over the
two scopes it can derive exactly (every `arg_groups` function taking a config,
every `ServerArgs` handler the dispatcher reaches). Readers the pipeline calls
from elsewhere (`ModelConfig`, the platform defaults, the spec-algo hook) have
moved to the view as well — a field read there is the same bug, just one the
derivation cannot enumerate.
One consequence worth knowing: because the fields are the raw input, resolving a
bare `dataclasses.replace` copy lands in the same place as the parent — the
pipeline reads only its own input. `replace_resolved` is the way to copy a
resolved record (it carries the declarations and the `model_config` memo, so the
copy does not re-resolve at all).
### Adding a model-specific config adjustment ### Adding a model-specific config adjustment
@@ -473,8 +498,12 @@ ONE thread — do not design for TBO threads that don't exist.
them explicitly on the mock; `MagicMock(spec=...)` raises on attributes that only them explicitly on the mock; `MagicMock(spec=...)` raises on attributes that only
exist post-`__init__`, which is the fastest way to find a missed stub. exist post-`__init__`, which is the fastest way to find a missed stub.
- `reset_context()` in teardown when a test publishes outside a scoped override. - `reset_context()` in teardown when a test publishes outside a scoped override.
- `ServerArgs(model_path="dummy")` early-returns `__post_init__` (no materialization, no - `ServerArgs(model_path="dummy")` early-returns the pipeline (few declarations, no
strict guard) — fine for lightweight fixtures. strict guard) — fine for lightweight fixtures.
- **Asserting what resolution decided reads `resolution_result(sa, "field")`**, not
`sa.field`: the field is the raw input. Assert the field only when the point of
the case *is* that the record stayed pristine (the FA4 page-size and waterfill
cases do exactly that, and say so).
- **Run changed test files per-file** (own process), the way CI does: a monolithic local - **Run changed test files per-file** (own process), the way CI does: a monolithic local
pytest run lets a context published by an earlier file mask a missing-publish bug in a pytest run lets a context published by an earlier file mask a missing-publish bug in a
later one. later one.
+2 -1
View File
@@ -23,6 +23,7 @@ from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
from sglang import Engine, ServerArgs from sglang import Engine, ServerArgs
from sglang.srt.arg_groups.overrides import resolution_result
parser = ArgumentParser() parser = ArgumentParser()
ServerArgs.add_cli_args(parser) ServerArgs.add_cli_args(parser)
@@ -44,7 +45,7 @@ parser.add_argument(
def main(args): def main(args):
engine_args = ServerArgs.from_cli_args(args) engine_args = ServerArgs.from_cli_args(args)
engine_args.resolve_once() engine_args.resolve_once()
model_path = engine_args.model_path model_path = resolution_result(engine_args, "model_path")
if not Path(model_path).is_dir(): if not Path(model_path).is_dir():
raise ValueError("model path must be a local directory") raise ValueError("model path must be a local directory")
# Create LLM instance from arguments # Create LLM instance from arguments
@@ -28,6 +28,7 @@ from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
from sglang import Engine, ServerArgs from sglang import Engine, ServerArgs
from sglang.srt.arg_groups.overrides import resolution_result
parser = ArgumentParser() parser = ArgumentParser()
ServerArgs.add_cli_args(parser) ServerArgs.add_cli_args(parser)
@@ -49,7 +50,7 @@ parser.add_argument(
def main(args): def main(args):
engine_args = ServerArgs.from_cli_args(args) engine_args = ServerArgs.from_cli_args(args)
engine_args.resolve_once() engine_args.resolve_once()
model_path = engine_args.model_path model_path = resolution_result(engine_args, "model_path")
if not Path(model_path).is_dir(): if not Path(model_path).is_dir():
raise ValueError("model path must be a local directory") raise ValueError("model path must be a local directory")
# Create LLM instance from arguments # Create LLM instance from arguments
@@ -5,6 +5,7 @@ from transformers import AutoProcessor
from sglang import Engine from sglang import Engine
from sglang.lang.chat_template import get_chat_template_by_model_path from sglang.lang.chat_template import get_chat_template_by_model_path
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.test_utils import DEFAULT_IMAGE_URL from sglang.test.test_utils import DEFAULT_IMAGE_URL
@@ -36,12 +37,13 @@ def get_input_ids(
def token_in_out_example( def token_in_out_example(
server_args: ServerArgs, server_args: ServerArgs,
): ):
cfg = resolving_view(server_args)
input_ids, image_data = get_input_ids( input_ids, image_data = get_input_ids(
server_args, server_args,
ModelConfig( ModelConfig(
server_args.model_path, cfg.model_path,
trust_remote_code=server_args.trust_remote_code, trust_remote_code=cfg.trust_remote_code,
model_override_args=server_args.json_model_override_args, model_override_args=cfg.json_model_override_args,
), ),
) )
backend = Engine(server_args=server_args) backend = Engine(server_args=server_args)
+1 -1
View File
@@ -1005,7 +1005,7 @@ def main(server_args, bench_args):
# The decode phase has to capture the batch sizes this run benchmarks, and # The decode phase has to capture the batch sizes this run benchmarks, and
# the per-phase convenience knob loses to an explicit --cuda-graph-config # the per-phase convenience knob loses to an explicit --cuda-graph-config
# JSON (resolution applies that last), so the size is merged into that JSON. # JSON (resolution applies that last), so the size is merged into that JSON.
if getattr(server_args, "_declarations_materialized", False): if getattr(server_args, "_resolution_finished", False):
# A record the caller already resolved: nothing will parse a raw dict # A record the caller already resolved: nothing will parse a raw dict
# again, so the declaration has to be the finished typed config. # again, so the declaration has to be the finished typed config.
merged = resolution_result(server_args, "cuda_graph_config") merged = resolution_result(server_args, "cuda_graph_config")
+4 -4
View File
@@ -75,10 +75,10 @@ class Arg:
# When True, this field is skipped by add_cli_args_from_dataclass. # When True, this field is skipped by add_cli_args_from_dataclass.
# Use for fields that have no CLI surface (e.g. injected via Python only). # Use for fields that have no CLI surface (e.g. injected via Python only).
no_cli: bool = False no_cli: bool = False
# When True, this field may be written by config resolution (model # When True, config resolution (model overrides and post-process passes)
# overrides and post-process passes): it is part of the whitelist accepted # may decide this field: the declaration stash accepts the name, and
# by the declaration stash, and its resolved value materializes onto the # `resolution_result` and the config bags answer with the decision. The
# field at the end of __post_init__. # field keeps what the operator passed.
resolvable: bool = False resolvable: bool = False
+41 -60
View File
@@ -14,9 +14,13 @@
"""Declarative model-override registry. """Declarative model-override registry.
Model-identity adjustments to the server configuration are DECLARED here and Model-identity adjustments to the server configuration are DECLARED here and
materialized onto ``server_args`` at the end of ``__post_init__`` (gate appended to the record's declaration stash (gate order, last writer wins).
order, last writer wins) — model code never mutates ``ServerArgs`` fields Nothing here writes back onto ``ServerArgs``: the record holds the user's raw
imperatively. input, and a decision is read through ``resolution_result`` or the published
config bags — model code never mutates ``ServerArgs`` fields imperatively. The
one channel that still leaves a field changed is ``declare_direct_writes``,
which does not perform the write: it captures one an out-of-tree plugin already
made, and undoing it would surprise the plugin's own reads.
Two declaration forms, keyed on ``hf_config.architectures[0]``: Two declaration forms, keyed on ``hf_config.architectures[0]``:
@@ -206,10 +210,8 @@ def register_post_process(fn: Callable[..., dict]) -> Callable[..., dict]:
def _declaration_overlay(server_args: Any) -> Dict[str, Any]: def _declaration_overlay(server_args: Any) -> Dict[str, Any]:
"""What the declarations say so far, last writer wins. """What the declarations say so far, last writer wins.
Passes declare without touching the fields until Nothing writes the fields, so a mid-resolution reader needs this to see a
``materialize_declarations``, so a mid-resolution reader needs this to see decision at all; the fields keep what the caller supplied."""
them; handlers and hooks write as they declare, and for those the overlay
repeats what the field already holds."""
overlay: Dict[str, Any] = {} overlay: Dict[str, Any] = {}
for _source, declared in getattr(server_args, "_resolved_overrides", None) or (): for _source, declared in getattr(server_args, "_resolved_overrides", None) or ():
overlay.update(declared) overlay.update(declared)
@@ -222,9 +224,9 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
Evaluates the pass on the resolving state (a read-only view with the Evaluates the pass on the resolving state (a read-only view with the
accumulated declarations overlaid from the stash) and appends its accumulated declarations overlaid from the stash) and appends its
declaration to the stash. During ``__post_init__`` the fields stay declaration to the stash. During ``__post_init__`` the fields stay
untouched — ``materialize_declarations`` applies the whole stash once at untouched: the stash is what the config bags are projected from. A pass
the end of resolution; a pass invoked after materialization (a post-init invoked after resolution finished (a post-init slot) writes through
slot) writes through immediately. immediately, because there is no later projection to pick it up.
""" """
declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args))) declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args)))
if not isinstance(declared, dict): if not isinstance(declared, dict):
@@ -244,7 +246,7 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
stash = server_args._resolved_overrides = [] stash = server_args._resolved_overrides = []
stash.append(entry) stash.append(entry)
validate_declarations(server_args, [entry]) validate_declarations(server_args, [entry])
if getattr(server_args, "_declarations_materialized", False): if getattr(server_args, "_resolution_finished", False):
_apply_fields(server_args, declared) _apply_fields(server_args, declared)
@@ -260,29 +262,17 @@ def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None:
def declare_resolution(server_args: Any, source: str, **fields: Any) -> None: def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
"""Record a resolution write in the declaration stash, and apply it now. """Record a resolution write in the declaration stash.
The stash is what the projection reads, so a resolver that only assigns The stash *is* the resolution result: the bags are projected from it,
the field leaves that write invisible to it. The immediate write keeps the `resolution_result` answers from it, and no field is written. A resolver
resolver's successors seeing the value where they read the field directly. reading a field another resolver may have decided must read `resolving_view`
(or `ServerArgs._resolved()`), which
`test_resolution_reads_the_declarations` pins.
What it does change is which writer wins. A declaration is appended and For resolvers inside ``__post_init__``; launcher-stage resolution goes
replayed last, so a resolver that declares a field a *deferred* writer (a through ``declare_late_resolution``. A name that is not a field is rejected
post-process pass, a registry entry) also decides now beats it, where its here rather than becoming an attribute nothing reads.
bare assignment used to be overwritten by that writer's declaration. A
resolver that gates on such a field has to read the resolving view rather
than the raw field, or it decides from a value that is already stale.
For resolvers inside ``__post_init__``: the handlers on ``ServerArgs``
(through ``self._declare``) and the ``arg_groups`` hooks and hardware
defaults they call. Resolution that has to wait for the launcher stage
goes through ``declare_late_resolution`` instead.
Names arrive as keyword arguments, which accept anything; a misspelled one
would otherwise become a new attribute that nothing ever reads, so it is
rejected here. This is not the model-override whitelist: that one limits
which fields a *registry entry* may reach, while a resolver writing the
field it owns is the pipeline resolving by construction.
""" """
if dataclasses.is_dataclass(type(server_args)): if dataclasses.is_dataclass(type(server_args)):
unknown = sorted(set(fields) - field_names(type(server_args))) unknown = sorted(set(fields) - field_names(type(server_args)))
@@ -293,8 +283,6 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
stash = [] stash = []
object.__setattr__(server_args, "_resolved_overrides", stash) object.__setattr__(server_args, "_resolved_overrides", stash)
stash.append((source, dict(fields))) stash.append((source, dict(fields)))
for name, value in fields.items():
setattr(server_args, name, value)
def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> None: def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> None:
@@ -304,10 +292,10 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non
normalization and the auto-parser detection need the launcher's validation normalization and the auto-parser detection need the launcher's validation
stage (and, for the parsers, a tokenizer / chat-template load). They still stage (and, for the parsers, a tokenizer / chat-template load). They still
belong to the resolution pipeline — they decide what the process will run belong to the resolution pipeline — they decide what the process will run
with — so they write the fields in place, before anything publishes the with — so their decision goes to the stash like any other, and the record
object. Writing in place is the point: every holder of that instance (the keeps what the caller passed. Every holder of that instance reads the
HTTP server, the multi-tokenizer workers it serializes for, the schedulers decision the same way the rest of the pipeline does: the bags it publishes,
it forks) must see the resolved value. or ``resolution_result``, both of which survive the pickle to a child.
Refuses to touch the published instance: after publish the bags exist and a Refuses to touch the published instance: after publish the bags exist and a
field write would desync them, which is what ``get_context().override`` is field write would desync them, which is what ``get_context().override`` is
@@ -334,7 +322,6 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non
stash = [] stash = []
object.__setattr__(server_args, "_resolved_overrides", stash) object.__setattr__(server_args, "_resolved_overrides", stash)
stash.append((source, dict(fields))) stash.append((source, dict(fields)))
_apply_fields(server_args, fields)
def declare_direct_writes( def declare_direct_writes(
@@ -348,7 +335,10 @@ def declare_direct_writes(
Out-of-tree platform plugins are handed the record and set fields on it. Out-of-tree platform plugins are handed the record and set fields on it.
Their implementations live outside this tree, so they cannot be converted Their implementations live outside this tree, so they cannot be converted
by editing the resolver; and the raw snapshot is taken before the pipeline by editing the resolver; and the raw snapshot is taken before the pipeline
starts, so a plugin's default is neither declared nor raw. starts, so a plugin's default is neither declared nor raw. The write itself
stays: this captures it into the stash so the projection and the bags carry
it, but reverting the field would break the plugin's own reads of what it
just set. It is the only field a record still carries from resolution.
Rebinding is what the diff sees, and rebinding is all it needs to see: a Rebinding is what the diff sees, and rebinding is all it needs to see: a
plugin that mutates a value in place reaches the projection anyway, because plugin that mutates a value in place reaches the projection anyway, because
@@ -384,24 +374,12 @@ def declare_direct_writes(
return result return result
def materialize_declarations(server_args: Any) -> None:
"""Apply the accumulated declarations onto ``server_args`` once, at the
end of ``__post_init__`` (gate order: last writer wins). After this the
fields carry the resolved configuration — every post-init reader, in any
process, reads them directly; ``resolved_view`` remains an internal
helper for mid-resolution code only."""
for _source, declared in getattr(server_args, "_resolved_overrides", None) or ():
for field, value in declared.items():
setattr(server_args, field, value)
server_args._declarations_materialized = True
def resolution_result(server_args: Any, field: str, default: Any = None) -> Any: def resolution_result(server_args: Any, field: str, default: Any = None) -> Any:
"""What resolution decided for ``field``: the declaration if there is one, """What resolution decided for ``field``: the declaration if there is one,
otherwise what the caller supplied. otherwise what the caller supplied.
This is what the config projection reads. Reading the field instead would This is what the config projection reads. Reading the field instead would
work only for as long as declarations materialize onto the record -- and work whatever the caller passed onto the record -- and
the point of declaring is that they will not, so the projection must not the point of declaring is that they will not, so the projection must not
depend on it. A config that never ran the pipeline (a mock, a partial depend on it. A config that never ran the pipeline (a mock, a partial
fixture) carries no raw snapshot; its fields are all it has. fixture) carries no raw snapshot; its fields are all it has.
@@ -422,9 +400,8 @@ def resolution_projection(server_args: Any) -> Dict[str, Any]:
The whole-object shape of ``resolution_result``, for the exits that hand out The whole-object shape of ``resolution_result``, for the exits that hand out
the entire configuration (``/server_info``, the gRPC and engine readbacks). the entire configuration (``/server_info``, the gRPC and engine readbacks).
They used ``dataclasses.asdict``, which reads the fields -- correct only for They used ``dataclasses.asdict``, which reads the fields -- the operator's
as long as declarations materialize onto the record, and the point of input, not what resolution decided. Field values only: the private resolution
declaring is that they will not. Field values only: the private resolution
bookkeeping and the ``model_config`` memo that a ``vars()`` dump carried into bookkeeping and the ``model_config`` memo that a ``vars()`` dump carried into
the readback are not configuration. the readback are not configuration.
""" """
@@ -453,10 +430,14 @@ def _plain(value: Any) -> Any:
def resolved_view(server_args: Any) -> ResolvedView: def resolved_view(server_args: Any) -> ResolvedView:
"""Read-only view of the resolving configuration for mid-resolution code """Read-only view of the resolving configuration: the declarations
that is not a pass (``__post_init__`` handlers and hooks). Internal to overlaid on the fields, snapshotted per call.
the resolution pipeline: after ``materialize_declarations`` runs, the
fields themselves carry the resolved values — read them directly.""" For mid-resolution code that is not a pass (``__post_init__`` handlers and
hooks), and for the record's own members that must answer with what
resolution decided -- a declaration-only resolver (a model-specific
override, a registry entry) never writes the field, so a field read there
answers with the raw input."""
return ResolvedView(server_args, overlay=_declaration_overlay(server_args)) return ResolvedView(server_args, overlay=_declaration_overlay(server_args))
+2 -2
View File
@@ -251,7 +251,7 @@ class Engine(EngineScoreMixin, EngineBase):
kwargs["log_level"] = "error" kwargs["log_level"] = "error"
server_args = self.server_args_class(**kwargs) server_args = self.server_args_class(**kwargs)
self.server_args = server_args self.server_args = server_args
logger.info(f"{server_args=}") logger.info(f"server_args={server_args.resolved_dict()}")
# Rust Server is not supported with the offline Engine API # Rust Server is not supported with the offline Engine API
if envs.SGLANG_RUST_SERVER.get(): if envs.SGLANG_RUST_SERVER.get():
@@ -1076,7 +1076,7 @@ class Engine(EngineScoreMixin, EngineBase):
# Allocate ports for inter-process communications # Allocate ports for inter-process communications
if port_args is None: if port_args is None:
port_args = PortArgs.init_new(server_args) port_args = PortArgs.init_new(server_args)
logger.info(f"{server_args=}") logger.info(f"server_args={server_args.resolved_dict()}")
# Start the engine info bootstrap server if per-rank info is needed. # Start the engine info bootstrap server if per-rank info is needed.
engine_info_bootstrap_server = None engine_info_bootstrap_server = None
+2 -2
View File
@@ -803,8 +803,8 @@ async def get_server_info():
async def server_info(): async def server_info():
"""The startup configuration, plus live scheduler state. """The startup configuration, plus live scheduler state.
The `ServerArgs` fields here are the record: what the launcher was given, The values here are the resolution result: what the launcher was given,
with resolution written back into it. Fields the control plane changes with every decision resolution made applied over it. Fields the control plane changes
after publication -- the model a weight update swapped in, its load format, after publication -- the model a weight update swapped in, its load format,
an operator-set weight version -- are reported by `/model_info`, and the an operator-set weight version -- are reported by `/model_info`, and the
HiCache mirror by `GET /hicache/storage-backend`. HiCache mirror by `GET /hicache/storage-backend`.
@@ -247,11 +247,10 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
Phase, Phase,
check_cuda_graph_backend, check_cuda_graph_backend,
) )
from sglang.srt.runtime_context import get_spec
_sa = getattr(runner, "server_args", None) spec = get_spec()
self.speculative_num_draft_tokens = getattr( self.speculative_num_draft_tokens = spec.speculative_num_draft_tokens
_sa, "speculative_num_draft_tokens", None
)
_decode_cuda_graph = not check_cuda_graph_backend( _decode_cuda_graph = not check_cuda_graph_backend(
Phase.DECODE, Backend.DISABLED Phase.DECODE, Backend.DISABLED
) )
@@ -264,7 +263,7 @@ class MiniMaxSparseAttnBackend(AttentionBackend):
if ( if (
self.use_msa self.use_msa
and _decode_cuda_graph and _decode_cuda_graph
and getattr(_sa, "speculative_algorithm", None) is not None and spec.speculative_algorithm is not None
): ):
raise NotImplementedError( raise NotImplementedError(
"MiniMax-M3 MSA attention does not support speculative decoding under " "MiniMax-M3 MSA attention does not support speculative decoding under "
+2 -2
View File
@@ -315,8 +315,8 @@ def initialize_moe_config():
"""Seed the MoE runtime flags from the published configuration. """Seed the MoE runtime flags from the published configuration.
Reads the bags: `moe_a2a_backend` and its siblings are resolution's Reads the bags: `moe_a2a_backend` and its siblings are resolution's
answers, and the record carries them only while declarations materialize answers, and the record carries the operator's input. Called once per
onto it. Called once per process after publish process after publish
(scheduler init, the benchmark work functions). (scheduler init, the benchmark work functions).
""" """
exec_moe = get_exec().moe exec_moe = get_exec().moe
+1 -2
View File
@@ -39,7 +39,6 @@ from sglang.srt.runtime_context import (
get_observability, get_observability,
get_parallel, get_parallel,
get_schedule, get_schedule,
get_server_args,
get_serving, get_serving,
get_spec, get_spec,
) )
@@ -4412,7 +4411,7 @@ class Scheduler(
# Resolved config (pristine server_args + post-publish overrides) so a # Resolved config (pristine server_args + post-publish overrides) so a
# readback reflects values changed via /set_internal_state, not startup. # readback reflects values changed via /set_internal_state, not startup.
ret = get_context().resolved_server_args_dict() ret = get_context().resolved_server_args_dict()
ret["world_size"] = compute_world_size(get_server_args()) ret["world_size"] = compute_world_size(get_parallel().config)
ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput
draft_graph_memory_usage = ( draft_graph_memory_usage = (
None if self.draft_worker is None else self.draft_worker.graph_memory_usage None if self.draft_worker is None else self.draft_worker.graph_memory_usage
@@ -1380,7 +1380,7 @@ class KVCacheConfigurator:
""" """
full_pool_class = DSATokenToKVPool if is_dsa_model else MLATokenToKVPool full_pool_class = DSATokenToKVPool if is_dsa_model else MLATokenToKVPool
common = { common = {
"page_size": self.server_args.page_size, "page_size": get_schedule().page_size,
"device": self.device, "device": self.device,
"enable_memory_saver": False, "enable_memory_saver": False,
} }
@@ -1401,7 +1401,7 @@ class KVCacheConfigurator:
return SWAKVPool( return SWAKVPool(
size=full_max_total_num_tokens, size=full_max_total_num_tokens,
size_swa=swa_max_total_num_tokens, size_swa=swa_max_total_num_tokens,
page_size=self.server_args.page_size, page_size=get_schedule().page_size,
dtype=self.kv_cache_dtype, dtype=self.kv_cache_dtype,
head_num=0, head_num=0,
head_dim=0, head_dim=0,
@@ -702,12 +702,13 @@ def _architecture_auto_parsers(server_args, needs: Tuple[str, ...]) -> Dict[str,
def resolve_auto_parsers(server_args) -> None: def resolve_auto_parsers(server_args) -> None:
"""Resolve ``--reasoning-parser=auto`` / ``--tool-call-parser=auto`` from the """Resolve ``--reasoning-parser=auto`` / ``--tool-call-parser=auto`` from the
chat template, in place, before anything publishes ``server_args``. chat template, before anything publishes ``server_args``.
Performs a lightweight tokenizer load, so it runs once in engine init. In Performs a lightweight tokenizer load, so it runs once in engine init. The
place because everyone who holds this instance must see the resolved value: decision goes to this instance's declaration stash, so every holder of it
the schedulers it forks, the HTTP server, and the tokenizer workers it is carries it -- the schedulers it forks, the HTTP server, the tokenizer
serialized for. workers it is serialized for -- and each publishes bags projected from it.
The fields stay what the operator passed.
""" """
cfg = resolving_view(server_args) cfg = resolving_view(server_args)
needs = tuple( needs = tuple(
+2 -6
View File
@@ -35,7 +35,7 @@ from sglang.srt.ray.scheduler_actor import SchedulerActor
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_parallel, get_parallel,
) )
from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.server_args import PortArgs, ServerArgs, compute_world_size
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -108,14 +108,10 @@ def _get_bundle_node_ip(placement_group: PlacementGroup, bundle_idx: int) -> str
def _compute_world_size() -> int: def _compute_world_size() -> int:
"""Compute world_size (total number of scheduler actors/GPUs needed). """Compute world_size (total number of scheduler actors/GPUs needed).
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
Reads the published parallel leaves: the driver is sizing the actors that Reads the published parallel leaves: the driver is sizing the actors that
will hold the process groups, so there is nothing live to ask. will hold the process groups, so there is nothing live to ask.
""" """
parallel = get_parallel().config return compute_world_size(get_parallel().config)
if parallel.enable_dp_attention:
return parallel.tp_size * parallel.pp_size
return parallel.dp_size * parallel.tp_size * parallel.pp_size
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
+21 -14
View File
@@ -24,9 +24,10 @@ getters. The resolved parallel **configuration** is the same object's ``config``
hop (``get_parallel().config.tp_size``), which reads the published ``parallel`` hop (``get_parallel().config.tp_size``), which reads the published ``parallel``
bag: bare is the live group, ``config`` is what was configured. bag: bare is the live group, ``config`` is what was configured.
``get_server_args()`` returns the process-wide ``ServerArgs``. This is the pristine / resolved-at-startup **read-only** record kept ``get_server_args()`` returns the process-wide ``ServerArgs``. This is the
for debug and reproduction; business code reads resolved config from the user's raw input, kept **read-only** for debug and reproduction; what
namespace bags below, not from this object. The context owns the storage: resolution decided lives in the declarations (``resolution_result``) and, for
business code, in the namespace bags below -- never on this object's fields. The context owns the storage:
publishing goes through ``RuntimeContext.set_server_args`` (the legacy publishing goes through ``RuntimeContext.set_server_args`` (the legacy
``set_global_server_args_for_scheduler`` / ``get_global_server_args`` are thin ``set_global_server_args_for_scheduler`` / ``get_global_server_args`` are thin
shims over this slot). shims over this slot).
@@ -440,8 +441,8 @@ class DpFlags(_FlagGroupBase):
class Flags(_FlagGroupBase): class Flags(_FlagGroupBase):
"""Root of the runtime-flags tier. """Root of the runtime-flags tier.
Resolved configuration lives on ``server_args`` fields (materialized at Resolved configuration lives in the config bags below (projected from the
the end of ``__post_init__``) — this tier only carries genuine runtime declarations at publish) — this tier only carries genuine runtime
state whose value is not a function of the configuration alone, grouped state whose value is not a function of the configuration alone, grouped
by lifecycle (``capture``) or subsystem (``moe`` / ``dp``). by lifecycle (``capture``) or subsystem (``moe`` / ``dp``).
""" """
@@ -700,8 +701,7 @@ def _build_config_bags(server_args: Any) -> dict:
"""Snapshot the resolution result into the namespace bag tree, driven by """Snapshot the resolution result into the namespace bag tree, driven by
the ``NS(...)`` metadata on the dataclass fields. Each leaf comes from the ``NS(...)`` metadata on the dataclass fields. Each leaf comes from
``resolution_result`` -- the declaration if resolution made one, else what ``resolution_result`` -- the declaration if resolution made one, else what
the caller supplied -- rather than from the field, which carries the same the caller supplied. Returns
value only while declarations still materialize. Returns
``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``). ``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``).
Only dataclass fields carry ``NS`` markers, so derived properties/methods are Only dataclass fields carry ``NS`` markers, so derived properties/methods are
naturally excluded (they stay on the bag). A name used as both a leaf and a naturally excluded (they stay on the bag). A name used as both a leaf and a
@@ -783,7 +783,13 @@ class RuntimeContext:
if stream is None: if stream is None:
import torch import torch
device = self._server_args.device if self._server_args else "cuda" from sglang.srt.arg_groups.overrides import resolution_result
device = (
resolution_result(self._server_args, "device")
if self._server_args
else "cuda"
)
stream = torch.get_device_module(device).Stream() stream = torch.get_device_module(device).Stream()
self.resources.streams[name] = stream self.resources.streams[name] = stream
return stream return stream
@@ -819,8 +825,8 @@ 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 published the draft-worker guard in ``ModelRunner.__init__``). The published
object already carries the resolved configuration (declarations object is the raw input; the resolution it carries is its declaration
materialize at the end of ``__post_init__``). stash, which is what the bags are projected from.
""" """
# Seed the capture tier for the new lifecycle (defaults for sentinel # Seed the capture tier for the new lifecycle (defaults for sentinel
# and mock publishes, which carry no config). # and mock publishes, which carry no config).
@@ -1078,10 +1084,11 @@ class _ServerArgsOverride:
} }
if declared: if declared:
declare_late_resolution(server_args, "override_server_args", **declared) declare_late_resolution(server_args, "override_server_args", **declared)
_apply_fields( # This hook stands in for a launch: the caller's values are both what
server_args, # the operator passed and what resolution decided, so they go on the
{name: value for name, value in self._fields.items() if name[0] == "_"}, # record as well as into the stash. Production late resolution declares
) # only -- there the record stays the operator's input.
_apply_fields(server_args, self._fields)
ctx.set_server_args(server_args) ctx.set_server_args(server_args)
self._installed = True self._installed = True
return server_args return server_args
+47 -42
View File
@@ -3709,7 +3709,7 @@ class ServerArgs:
arrived by pickle and brought its declarations along, so the child has arrived by pickle and brought its declarations along, so the child has
nothing left to derive and projects what the parent decided. nothing left to derive and projects what the parent decided.
""" """
if getattr(self, "_declarations_materialized", False): if getattr(self, "_resolution_finished", False):
return return
if getattr(self, "_resolution_failed", False): if getattr(self, "_resolution_failed", False):
raise RuntimeError( raise RuntimeError(
@@ -3721,23 +3721,22 @@ class ServerArgs:
try: try:
self._run_resolution_pipeline() self._run_resolution_pipeline()
except BaseException: except BaseException:
# The handlers that ran already wrote to the record, and they are # The handlers that ran already declared, and they are not
# not idempotent over their own output. # idempotent over their own output.
object.__setattr__(self, "_resolution_failed", True) object.__setattr__(self, "_resolution_failed", True)
raise raise
# Set here too, because the dummy/absent-model path returns before the # Set here too, because the dummy/absent-model path returns before the
# materialization that normally sets it: the gate is about whether the # end of the pipeline that normally sets it: the gate is about whether
# handlers ran, not how far they got. # the handlers ran, not how far they got.
self._declarations_materialized = True self._resolution_finished = True
def resolved_dict(self) -> Dict[str, Any]: def resolved_dict(self) -> Dict[str, Any]:
"""This configuration as a plain dict of resolved field values. """This configuration as a plain dict of resolved field values.
What the whole-object readbacks report (`/server_info` and its gRPC and What the whole-object readbacks report (`/server_info` and its gRPC and
in-process twins). `dataclasses.asdict(self)` reads the fields, which in-process twins). `dataclasses.asdict(self)` reads the fields, which
carry resolution's result only while declarations materialize onto the carry the raw input; this reads the declarations, so it answers with what
record; this reads the declarations, so it keeps answering with what resolution decided. Nested dataclass fields are expanded
resolution decided once they stop. Nested dataclass fields are expanded
the way `asdict` expands them; the private resolution bookkeeping and the the way `asdict` expands them; the private resolution bookkeeping and the
`model_config` memo are not fields and do not appear. `model_config` memo are not fields and do not appear.
""" """
@@ -3750,12 +3749,11 @@ class ServerArgs:
`dataclasses.replace` builds a new instance, so the copy carries none of `dataclasses.replace` builds a new instance, so the copy carries none of
what makes a record resolved: no raw snapshot, no declarations, no what makes a record resolved: no raw snapshot, no declarations, no
materialization. The next publish therefore finds an unmaterialized finished flag. The next publish therefore resolves it again, which
record and runs the pipeline over values it already decided -- DP drops every decision the stash held -- the late ones (the auto-detected
attention halves `chunked_prefill_size` a second time (8192 -> 4096 -> parsers) and the direct ones alike -- and re-runs the device probes in
2048) and the schedule conservativeness is scaled again (0.3 -> 0.09). whatever process opened the copy. The Ray paths replace
The Ray paths replace `dist_init_addr` on a resolved record, which is `dist_init_addr` on a resolved record, which is how they reach this.
how they hit it.
The change is appended to the stash rather than left on the field: the The change is appended to the stash rather than left on the field: the
projection reads the raw snapshot plus the declarations, so a field the projection reads the raw snapshot plus the declarations, so a field the
@@ -3770,7 +3768,7 @@ class ServerArgs:
copy's deep structure in-process mutates the parent's too. copy's deep structure in-process mutates the parent's too.
""" """
replacement = dataclasses.replace(self, **changes) replacement = dataclasses.replace(self, **changes)
if not getattr(self, "_declarations_materialized", False): if not getattr(self, "_resolution_finished", False):
# Not resolved yet: the copy goes through the gate itself. # Not resolved yet: the copy goes through the gate itself.
return replacement return replacement
@@ -3780,7 +3778,7 @@ class ServerArgs:
# (the read-only guard refuses the write). # (the read-only guard refuses the write).
field_names = {field.name for field in dataclasses.fields(self)} field_names = {field.name for field in dataclasses.fields(self)}
for name, value in vars(self).items(): for name, value in vars(self).items():
if name in field_names or name == "_declarations_materialized": if name in field_names or name == "_resolution_finished":
continue continue
if isinstance(value, (dict, list, set)): if isinstance(value, (dict, list, set)):
value = copy.copy(value) value = copy.copy(value)
@@ -3791,7 +3789,7 @@ class ServerArgs:
object.__setattr__(replacement, "_resolved_overrides", stash) object.__setattr__(replacement, "_resolved_overrides", stash)
if changes: if changes:
stash.append((source, dict(changes))) stash.append((source, dict(changes)))
object.__setattr__(replacement, "_declarations_materialized", True) object.__setattr__(replacement, "_resolution_finished", True)
return replacement return replacement
def _declare(self, source: str, **fields: Any) -> None: def _declare(self, source: str, **fields: Any) -> None:
@@ -4022,13 +4020,7 @@ class ServerArgs:
# time; last declarations of the resolution, mirroring that order. # time; last declarations of the resolution, mirroring that order.
self._handle_model_capability_adjustments() self._handle_model_capability_adjustments()
# End of resolution: apply the accumulated declarations onto the self._resolution_finished = True
# fields once (gate order). From here on server_args carries the
# resolved configuration — post-init readers, in any process, read
# the fields directly.
from sglang.srt.arg_groups.overrides import materialize_declarations
materialize_declarations(self)
def _handle_return_hidden_states_mode(self): def _handle_return_hidden_states_mode(self):
cfg = resolving_view(self) cfg = resolving_view(self)
@@ -9794,22 +9786,22 @@ class ServerArgs:
def _late_resolution(self, source: str, **fields) -> None: def _late_resolution(self, source: str, **fields) -> None:
"""Resolve fields at the launcher's validation stage (pre-publish). """Resolve fields at the launcher's validation stage (pre-publish).
See ``arg_groups.overrides.declare_late_resolution``: in place, because See ``arg_groups.overrides.declare_late_resolution``: the decision goes
every holder of this instance must see the resolved value, and refused to this instance's declaration stash, so every holder of it carries the
outright once the config is published. decision and publishes bags that answer with it. Refused outright once
the config is published.
""" """
from sglang.srt.arg_groups.overrides import declare_late_resolution from sglang.srt.arg_groups.overrides import declare_late_resolution
declare_late_resolution(self, source, **fields) declare_late_resolution(self, source, **fields)
def __setattr__(self, name, value): def __setattr__(self, name, value):
# After materialization the fields are the resolved startup # Once resolution has finished the record is the READ-ONLY raw input
# configuration -- the pristine, READ-ONLY record that the config bags # the config bags were projected from. Resolved config changes go to the bags via
# were projected from. Resolved config changes go to the bags via
# get_context().override(source, ...); a value one runner or worker # get_context().override(source, ...); a value one runner or worker
# owns travels as a constructor argument to it. # owns travels as a constructor argument to it.
if ( if (
getattr(self, "_declarations_materialized", False) getattr(self, "_resolution_finished", False)
and not getattr(self, "_internal_write", False) and not getattr(self, "_internal_write", False)
and name not in _CACHE_SLOTS and name not in _CACHE_SLOTS
and (not name.startswith("_") or name in _underscore_field_names()) and (not name.startswith("_") or name in _underscore_field_names())
@@ -9832,7 +9824,13 @@ class ServerArgs:
return attention_backends_of(resolved_view(self)) return attention_backends_of(resolved_view(self))
def get_attention_backends(self): def get_attention_backends(self):
return attention_backends_of(self) """The (prefill, decode) pair resolution decided.
Reads through the declaration stash, not the fields: the model-specific
overrides declare into the stash without writing the fields, so a field
read answers with what the operator typed.
"""
return attention_backends_of(resolved_view(self))
def use_mla_backend(self): def use_mla_backend(self):
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
@@ -9884,7 +9882,7 @@ class ServerArgs:
# state needs steps + 1 draft-token slots. Revisit this if topk>1 # state needs steps + 1 draft-token slots. Revisit this if topk>1
# is supported. # is supported.
result = max(candidate_steps) + 1 result = max(candidate_steps) + 1
if getattr(self, "_declarations_materialized", False): if getattr(self, "_resolution_finished", False):
object.__setattr__(self, "_max_speculative_num_draft_tokens", result) object.__setattr__(self, "_max_speculative_num_draft_tokens", result)
return result return result
@@ -9913,7 +9911,7 @@ class ServerArgs:
assert ( assert (
max(chunk_size, page_size) % min(chunk_size, page_size) == 0 max(chunk_size, page_size) % min(chunk_size, page_size) == 0
), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}" ), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}"
if not getattr(self, "_declarations_materialized", False): if not getattr(self, "_resolution_finished", False):
return max(chunk_size, page_size) return max(chunk_size, page_size)
self._mamba_cache_chunk_size = max(chunk_size, page_size) self._mamba_cache_chunk_size = max(chunk_size, page_size)
return self._mamba_cache_chunk_size return self._mamba_cache_chunk_size
@@ -10574,8 +10572,8 @@ class ServerArgs:
"endpoint_host": host, "endpoint_host": host,
"endpoint_port_base": port, "endpoint_port_base": port,
"topic": cfg.topic, "topic": cfg.topic,
"block_size": self.kv_event_block_size, "block_size": resolved.kv_event_block_size,
"dp_size": self.dp_size, "dp_size": resolved.dp_size,
} }
def should_report_expert_balancedness(self) -> bool: def should_report_expert_balancedness(self) -> bool:
@@ -10593,12 +10591,19 @@ class ServerArgs:
return cfg.expert_balancedness_report_mode in ("prometheus", "both") return cfg.expert_balancedness_report_mode in ("prometheus", "both")
def compute_world_size(server_args: ServerArgs) -> int: def compute_world_size(config) -> int:
"""Return the total GPU count across all data-parallel replicas.""" """Return the total GPU count across all data-parallel replicas.
Takes the resolved topology -- the published `parallel` bag, or a view over
the declarations. `enable_dp_attention` and `dp_size` are both resolution's
answers (`_handle_dwdp` fills the pair, DeepSeek MLA context parallelism
turns DP attention on), so a raw-record read would size the world from what
the operator typed.
"""
return ( return (
(1 if server_args.enable_dp_attention else server_args.dp_size) (1 if config.enable_dp_attention else config.dp_size)
* server_args.tp_size * config.tp_size
* server_args.pp_size * config.pp_size
) )
+4
View File
@@ -22,6 +22,7 @@ from transformers import AutoTokenizer
from sglang.bench_serving import benchmark, set_global_args from sglang.bench_serving import benchmark, set_global_args
from sglang.benchmark.datasets import DatasetRow from sglang.benchmark.datasets import DatasetRow
from sglang.benchmark.datasets.mmmu import sample_mmmu_requests from sglang.benchmark.datasets.mmmu import sample_mmmu_requests
from sglang.srt.arg_groups.overrides import resolution_projection
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.test_utils import ( from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
@@ -135,6 +136,9 @@ def send_one_batch(base_url, num_prompts, batch_size, processor, is_multimodal):
def main(args, server_args): def main(args, server_args):
from sglang.srt.arg_groups.overrides import resolution_projection
server_args = SimpleNamespace(**resolution_projection(server_args))
base_url = "http://127.0.0.1:20000" base_url = "http://127.0.0.1:20000"
configs = [] configs = []
@@ -87,22 +87,25 @@ def launch_server_process(
server_args: ServerArgs, worker_port: int, dp_id: int server_args: ServerArgs, worker_port: int, dp_id: int
) -> mp.Process: ) -> mp.Process:
"""Launch a single server process with the given args and port.""" """Launch a single server process with the given args and port."""
# This binding is installed against a released sglang, so it cannot call
# into helpers newer than that wheel. Copy first, then write through the
# sanctioned channel if the record is resolved (a resolved record refuses
# plain assignment), else assign.
worker_args = copy.deepcopy(server_args)
changes = { changes = {
"port": worker_port, "port": worker_port,
"base_gpu_id": dp_id * server_args.tp_size, "base_gpu_id": dp_id * server_args.tp_size,
"dp_size": 1, "dp_size": 1,
} }
late = getattr(worker_args, "_late_resolution", None) # Three channels, newest first. A wheel that has the read-only record but not
if late is not None and getattr(worker_args, "_declarations_materialized", False): # `replace_resolved` still has `_late_resolution`, and plain assignment there
late("sglang_router.launch_server_process", **changes) # raises; only a wheel with neither accepts `setattr`.
replace_resolved = getattr(server_args, "replace_resolved", None)
if replace_resolved is not None:
worker_args = replace_resolved("sglang_router.launch_server_process", **changes)
else: else:
for field, value in changes.items(): worker_args = copy.deepcopy(server_args)
setattr(worker_args, field, value) late = getattr(worker_args, "_late_resolution", None)
if late is not None:
late("sglang_router.launch_server_process", **changes)
else:
for field, value in changes.items():
setattr(worker_args, field, value)
server_args = worker_args server_args = worker_args
proc = mp.Process(target=run_server, args=(server_args, dp_id)) proc = mp.Process(target=run_server, args=(server_args, dp_id))
@@ -188,7 +191,11 @@ def main():
server_args.resolve_once() server_args.resolve_once()
router_args = RouterArgs.from_cli_args(args, use_router_prefix=True) router_args = RouterArgs.from_cli_args(args, use_router_prefix=True)
# Find available ports for workers # Find available ports for workers. The count is the operator's requested
# replica count, which is the raw field on purpose: `--dwdp-size` makes
# resolution declare a `dp_size` that describes one multi-rank server's
# internal topology, and spawning that many single-rank children would ask
# for dp_size^2 GPUs.
worker_ports = find_available_ports( worker_ports = find_available_ports(
args.router_dp_worker_base_port, server_args.dp_size args.router_dp_worker_base_port, server_args.dp_size
) )
@@ -856,6 +856,62 @@ def test_launch_server_process_and_cleanup(monkeypatch):
assert (p1.pid, _sig.SIGTERM) in calls and (p2.pid, _sig.SIGTERM) in calls assert (p1.pid, _sig.SIGTERM) in calls and (p2.pid, _sig.SIGTERM) in calls
assert (p2.pid, _sig.SIGKILL) in calls assert (p2.pid, _sig.SIGKILL) in calls
def test_launch_server_process_declares_on_a_resolved_record(monkeypatch):
"""A record that carries its resolution takes the declaration channel.
The stub above has no `replace_resolved`, so it exercises the older wheels'
path. A current `ServerArgs` refuses plain assignment once resolution has
finished; the per-worker values reach the child as a declaration on a copy,
and the parent keeps what the operator passed.
"""
_install_sglang_stubs(monkeypatch)
import importlib
ls = importlib.import_module("sglang_router.launch_server")
created = {}
class FakeProcess:
def __init__(self, target, args):
created["target"] = target
created["args"] = args
self.pid = 4243
def start(self):
created["started"] = True
monkeypatch.setattr(ls.mp, "Process", FakeProcess)
calls = []
class ResolvedServerArgs:
def __init__(self, **fields):
self.port = fields.get("port", 30000)
self.base_gpu_id = fields.get("base_gpu_id", 0)
self.dp_size = fields.get("dp_size", 4)
self.tp_size = fields.get("tp_size", 2)
def replace_resolved(self, source, **changes):
calls.append((source, dict(changes)))
fields = dict(vars(self))
fields.update(changes)
return ResolvedServerArgs(**fields)
parent = ResolvedServerArgs()
proc = ls.launch_server_process(parent, worker_port=31002, dp_id=3)
assert created.get("started") is True
assert proc.pid == 4243
assert len(calls) == 1
source, changes = calls[0]
assert source == "sglang_router.launch_server_process"
assert changes == {"port": 31002, "base_gpu_id": 6, "dp_size": 1}
worker = created["args"][0]
assert (worker.port, worker.base_gpu_id, worker.dp_size) == (31002, 6, 1)
assert (parent.port, parent.base_gpu_id, parent.dp_size) == (30000, 0, 4)
def test_validation_error_handling(self): def test_validation_error_handling(self):
"""Test error handling when validation fails.""" """Test error handling when validation fails."""
args = RouterArgs( args = RouterArgs(
@@ -679,7 +679,7 @@ mod tests {
let mut workers: Vec<Arc<dyn Worker>> = Vec::new(); let mut workers: Vec<Arc<dyn Worker>> = Vec::new();
for j in 0..num_workers { for j in 0..num_workers {
workers.push(Arc::new( workers.push(Arc::new(
BasicWorkerBuilder::new(&format!("http://w{}:8000", j)) BasicWorkerBuilder::new(format!("http://w{}:8000", j))
.worker_type(WorkerType::Regular) .worker_type(WorkerType::Regular)
.build(), .build(),
)); ));
@@ -738,7 +738,7 @@ mod tests {
let mut workers: Vec<Arc<dyn Worker>> = Vec::new(); let mut workers: Vec<Arc<dyn Worker>> = Vec::new();
for j in 0..num_workers { for j in 0..num_workers {
workers.push(Arc::new( workers.push(Arc::new(
BasicWorkerBuilder::new(&format!("http://w{}:8000", j)) BasicWorkerBuilder::new(format!("http://w{}:8000", j))
.worker_type(WorkerType::Regular) .worker_type(WorkerType::Regular)
.build(), .build(),
)); ));
@@ -4,6 +4,7 @@ import tempfile
import unittest import unittest
from unittest.mock import patch from unittest.mock import patch
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -25,8 +26,10 @@ class TestServerArgsCPUBackend(unittest.TestCase):
ServerArgs._handle_cpu_backends(server_args) ServerArgs._handle_cpu_backends(server_args)
self.assertEqual(server_args.attention_backend, "torch_native") self.assertEqual(
self.assertEqual(server_args.sampling_backend, "pytorch") resolution_result(server_args, "attention_backend"), "torch_native"
)
self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch")
@patch("sglang.srt.server_args.is_host_cpu_arm64", return_value=False) @patch("sglang.srt.server_args.is_host_cpu_arm64", return_value=False)
def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64):
@@ -34,8 +37,10 @@ class TestServerArgsCPUBackend(unittest.TestCase):
ServerArgs._handle_cpu_backends(server_args) ServerArgs._handle_cpu_backends(server_args)
self.assertEqual(server_args.attention_backend, "intel_amx") self.assertEqual(
self.assertEqual(server_args.sampling_backend, "pytorch") resolution_result(server_args, "attention_backend"), "intel_amx"
)
self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch")
class TestServerArgsIBDeviceValidation(unittest.TestCase): class TestServerArgsIBDeviceValidation(unittest.TestCase):
@@ -48,10 +48,11 @@ class TestSchedulerInternalStateEnvVars(unittest.TestCase):
), patch( ), patch(
"sglang.srt.managers.scheduler.get_exec", "sglang.srt.managers.scheduler.get_exec",
return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)),
), patch(
"sglang.srt.managers.scheduler.get_server_args", return_value=None
), patch( ), patch(
"sglang.srt.managers.scheduler.compute_world_size", return_value=1 "sglang.srt.managers.scheduler.compute_world_size", return_value=1
), patch(
"sglang.srt.managers.scheduler.get_parallel",
return_value=SimpleNamespace(config=SimpleNamespace()),
): ):
output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) output = scheduler.get_internal_state(recv_req=GetInternalStateReq())
@@ -9,16 +9,16 @@ maybe_stub_sgl_kernel()
from sglang.srt.managers.io_struct import GetInternalStateReq from sglang.srt.managers.io_struct import GetInternalStateReq
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.server_args import ServerArgs, compute_world_size from sglang.srt.server_args import compute_world_size
register_cpu_ci(est_time=5, suite="base-a-test-cpu") register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _make_server_args( def _make_parallel_config(
*, tp_size: int, pp_size: int, dp_size: int, enable_dp_attention: bool *, tp_size: int, pp_size: int, dp_size: int, enable_dp_attention: bool
) -> ServerArgs: ) -> SimpleNamespace:
return ServerArgs( """The four `parallel` leaves the world size is computed from."""
model_path="dummy", return SimpleNamespace(
tp_size=tp_size, tp_size=tp_size,
pp_size=pp_size, pp_size=pp_size,
dp_size=dp_size, dp_size=dp_size,
@@ -29,39 +29,39 @@ def _make_server_args(
class TestComputeWorldSize(unittest.TestCase): class TestComputeWorldSize(unittest.TestCase):
def test_a_single_gpu_server_holds_one_gpu(self): def test_a_single_gpu_server_holds_one_gpu(self):
"""The default shape has to come out as one, or every consumer is off by a factor.""" """The default shape has to come out as one, or every consumer is off by a factor."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False
) )
self.assertEqual(compute_world_size(server_args), 1) self.assertEqual(compute_world_size(config), 1)
def test_tensor_and_pipeline_stages_multiply(self): def test_tensor_and_pipeline_stages_multiply(self):
"""Each (pp_rank, tp_rank) pair is its own scheduler process on its own gpu.""" """Each (pp_rank, tp_rank) pair is its own scheduler process on its own gpu."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False
) )
self.assertEqual(compute_world_size(server_args), 6) self.assertEqual(compute_world_size(config), 6)
def test_plain_data_parallel_replicas_each_hold_their_own_gpus(self): def test_plain_data_parallel_replicas_each_hold_their_own_gpus(self):
"""Without dp attention every replica launches a full tensor-parallel group of its own.""" """Without dp attention every replica launches a full tensor-parallel group of its own."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False
) )
self.assertEqual(compute_world_size(server_args), 4) self.assertEqual(compute_world_size(config), 4)
def test_data_parallel_attention_shares_the_tensor_parallel_gpus(self): def test_data_parallel_attention_shares_the_tensor_parallel_gpus(self):
"""With dp attention the dp ranks live inside the tensor-parallel world, not beside it.""" """With dp attention the dp ranks live inside the tensor-parallel world, not beside it."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True
) )
self.assertEqual(compute_world_size(server_args), 4) self.assertEqual(compute_world_size(config), 4)
class TestSchedulerInternalStateWorldSize(unittest.TestCase): class TestSchedulerInternalStateWorldSize(unittest.TestCase):
def _get_internal_state(self, server_args: ServerArgs) -> dict: def _get_internal_state(self, config: SimpleNamespace) -> dict:
scheduler = Scheduler.__new__(Scheduler) scheduler = Scheduler.__new__(Scheduler)
scheduler.metrics_reporter = SimpleNamespace( scheduler.metrics_reporter = SimpleNamespace(
last_gen_throughput=1.0, last_gen_throughput=1.0,
@@ -94,8 +94,8 @@ class TestSchedulerInternalStateWorldSize(unittest.TestCase):
"sglang.srt.managers.scheduler.get_exec", "sglang.srt.managers.scheduler.get_exec",
return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)),
), patch( ), patch(
"sglang.srt.managers.scheduler.get_server_args", "sglang.srt.managers.scheduler.get_parallel",
return_value=server_args, return_value=SimpleNamespace(config=config),
): ):
output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) output = scheduler.get_internal_state(recv_req=GetInternalStateReq())
@@ -103,24 +103,24 @@ class TestSchedulerInternalStateWorldSize(unittest.TestCase):
def test_the_internal_state_reports_the_whole_server(self): def test_the_internal_state_reports_the_whole_server(self):
"""A consumer sizing an external fleet reads the gpus the server occupies, not the declared sizes.""" """A consumer sizing an external fleet reads the gpus the server occupies, not the declared sizes."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False
) )
internal_state = self._get_internal_state(server_args) internal_state = self._get_internal_state(config)
self.assertEqual(internal_state["world_size"], 4) self.assertEqual(internal_state["world_size"], 4)
def test_the_reported_size_is_not_one_replica_of_a_data_parallel_server(self): def test_the_reported_size_is_not_one_replica_of_a_data_parallel_server(self):
"""Each plain dp replica has its own process group, so no scheduler can report the whole server from it.""" """Each plain dp replica has its own process group, so no scheduler can report the whole server from it."""
server_args = _make_server_args( config = _make_parallel_config(
tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False
) )
internal_state = self._get_internal_state(server_args) internal_state = self._get_internal_state(config)
self.assertNotEqual( self.assertNotEqual(
internal_state["world_size"], server_args.tp_size * server_args.pp_size internal_state["world_size"], config.tp_size * config.pp_size
) )
@@ -765,6 +765,18 @@ class TestToolCallParserDetection(unittest.TestCase):
self.assertEqual(result, "minicpm5") self.assertEqual(result, "minicpm5")
def _declared(server_args, field):
"""What late resolution decided for `field` on this record.
`resolve_auto_parsers` declares; the field keeps what the operator passed,
so the decision is read through the resolution result -- the same surface
the config bags are projected from.
"""
from sglang.srt.arg_groups.overrides import resolution_result
return resolution_result(server_args, field)
class TestResolveAutoParsers(unittest.TestCase): class TestResolveAutoParsers(unittest.TestCase):
"""Tests for resolve_auto_parsers().""" """Tests for resolve_auto_parsers()."""
@@ -790,8 +802,8 @@ class TestResolveAutoParsers(unittest.TestCase):
with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "qwen3") self.assertEqual(_declared(args, "reasoning_parser"), "qwen3")
self.assertEqual(args.tool_call_parser, "qwen") self.assertEqual(_declared(args, "tool_call_parser"), "qwen")
def test_resolves_reasoning_parser_only(self): def test_resolves_reasoning_parser_only(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser=None) args = self._make_server_args(reasoning_parser="auto", tool_call_parser=None)
@@ -800,8 +812,8 @@ class TestResolveAutoParsers(unittest.TestCase):
with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "qwen3") self.assertEqual(_declared(args, "reasoning_parser"), "qwen3")
self.assertIsNone(args.tool_call_parser) self.assertIsNone(_declared(args, "tool_call_parser"))
def test_resolves_tool_call_parser_only(self): def test_resolves_tool_call_parser_only(self):
args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="auto")
@@ -810,14 +822,14 @@ class TestResolveAutoParsers(unittest.TestCase):
with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "qwen3") self.assertEqual(_declared(args, "reasoning_parser"), "qwen3")
self.assertEqual(args.tool_call_parser, "qwen") self.assertEqual(_declared(args, "tool_call_parser"), "qwen")
def test_neither_auto_is_noop(self): def test_neither_auto_is_noop(self):
args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="qwen") args = self._make_server_args(reasoning_parser="qwen3", tool_call_parser="qwen")
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "qwen3") self.assertEqual(_declared(args, "reasoning_parser"), "qwen3")
self.assertEqual(args.tool_call_parser, "qwen") self.assertEqual(_declared(args, "tool_call_parser"), "qwen")
def test_nonexistent_model_disables_both_parsers(self): def test_nonexistent_model_disables_both_parsers(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -832,8 +844,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertIsNone(args.reasoning_parser) self.assertIsNone(_declared(args, "reasoning_parser"))
self.assertIsNone(args.tool_call_parser) self.assertIsNone(_declared(args, "tool_call_parser"))
def test_none_chat_template_disables_both_parsers(self): def test_none_chat_template_disables_both_parsers(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -842,8 +854,8 @@ class TestResolveAutoParsers(unittest.TestCase):
with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertIsNone(args.reasoning_parser) self.assertIsNone(_declared(args, "reasoning_parser"))
self.assertIsNone(args.tool_call_parser) self.assertIsNone(_declared(args, "tool_call_parser"))
def test_deepseek_v32_arch_without_chat_template_uses_custom_encoder(self): def test_deepseek_v32_arch_without_chat_template_uses_custom_encoder(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -855,8 +867,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "deepseek-v3") self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3")
self.assertEqual(args.tool_call_parser, "deepseekv32") self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32")
def test_deepseek_v4_arch_without_chat_template_uses_custom_encoder(self): def test_deepseek_v4_arch_without_chat_template_uses_custom_encoder(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -868,8 +880,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "deepseek-v4") self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v4")
self.assertEqual(args.tool_call_parser, "deepseekv4") self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv4")
def test_kimi_k3_arch_without_chat_template_uses_custom_encoder(self): def test_kimi_k3_arch_without_chat_template_uses_custom_encoder(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -883,8 +895,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "kimi_k3") self.assertEqual(_declared(args, "reasoning_parser"), "kimi_k3")
self.assertEqual(args.tool_call_parser, "kimi_k3") self.assertEqual(_declared(args, "tool_call_parser"), "kimi_k3")
def test_kimi_k3_model_type_without_architecture_uses_custom_encoder(self): def test_kimi_k3_model_type_without_architecture_uses_custom_encoder(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -896,8 +908,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "kimi_k3") self.assertEqual(_declared(args, "reasoning_parser"), "kimi_k3")
self.assertEqual(args.tool_call_parser, "kimi_k3") self.assertEqual(_declared(args, "tool_call_parser"), "kimi_k3")
def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self): def test_deepseek_arch_fallback_runs_when_tokenizer_load_fails(self):
args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto") args = self._make_server_args(reasoning_parser="auto", tool_call_parser="auto")
@@ -909,8 +921,8 @@ class TestResolveAutoParsers(unittest.TestCase):
): ):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "deepseek-v3") self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3")
self.assertEqual(args.tool_call_parser, "deepseekv32") self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32")
def test_explicit_non_jinja_template_skips_architecture_fallback(self): def test_explicit_non_jinja_template_skips_architecture_fallback(self):
args = self._make_server_args( args = self._make_server_args(
@@ -926,8 +938,8 @@ class TestResolveAutoParsers(unittest.TestCase):
resolve_auto_parsers(args) resolve_auto_parsers(args)
get_config.assert_not_called() get_config.assert_not_called()
self.assertIsNone(args.reasoning_parser) self.assertIsNone(_declared(args, "reasoning_parser"))
self.assertIsNone(args.tool_call_parser) self.assertIsNone(_declared(args, "tool_call_parser"))
def test_explicit_jinja_template_takes_precedence(self): def test_explicit_jinja_template_takes_precedence(self):
tokenizer = _DummyTokenizer([], chat_template=None) tokenizer = _DummyTokenizer([], chat_template=None)
@@ -947,8 +959,8 @@ class TestResolveAutoParsers(unittest.TestCase):
with _patch_hf_transformers_utils(Mock(return_value=tokenizer)): with _patch_hf_transformers_utils(Mock(return_value=tokenizer)):
resolve_auto_parsers(args) resolve_auto_parsers(args)
self.assertEqual(args.reasoning_parser, "deepseek-v3") self.assertEqual(_declared(args, "reasoning_parser"), "deepseek-v3")
self.assertEqual(args.tool_call_parser, "deepseekv32") self.assertEqual(_declared(args, "tool_call_parser"), "deepseekv32")
if __name__ == "__main__": if __name__ == "__main__":
@@ -26,8 +26,8 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" _SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt"
# Read for what the caller asked for: the constructor passes it through and # Read for what the caller asked for: the constructor passes it through and
# never stores it, while resolution later overwrites the field with the value # never stores it, while resolution declares the value the architecture implies.
# the architecture implies. Two quantities sharing one name. # Two quantities sharing one name.
_READ_BEFORE_RESOLUTION = frozenset({"is_embedding"}) _READ_BEFORE_RESOLUTION = frozenset({"is_embedding"})
# Declared after the first `get_model_config()`, so the cached configuration # Declared after the first `get_model_config()`, so the cached configuration
@@ -5,13 +5,11 @@ so a resolution write that only assigns the field is invisible to it. Every
resolver declares now -- the record's handlers through `self._declare`, the resolver declares now -- the record's handlers through `self._declare`, the
hooks and hardware defaults through `declare_resolution` -- and that is pinned hooks and hardware defaults through `declare_resolution` -- and that is pinned
two ways: no bare assignment to a field survives anywhere a ServerArgs instance two ways: no bare assignment to a field survives anywhere a ServerArgs instance
is in reach, and after resolution every declared field agrees with what the is in reach, and after resolution `resolution_result` answers for every declared
stash says. The second check is the one that keeps the transition honest -- field with what the stash holds. The second check is what the stash is measured
while a declaration still writes the field immediately, a stash entry and a against: the two can disagree only if something wrote behind the stash's back. A
field can only disagree if something assigned the field behind the stash's third check runs the other way -- every field resolution moved has to be
back. A third check runs the other way: every field resolution moved has to explained by the stash, which covers the spellings a source scan cannot see.
be explained by the stash, which covers the spellings a source scan cannot
see.
""" """
import ast import ast
@@ -233,6 +231,11 @@ def _bare_assignments():
return sorted(found) return sorted(found)
def shape_key(shape):
"""A shape rendered short enough for a failure message."""
return ",".join(f"{k}={v}" for k, v in sorted(shape.items())) or "defaults"
def _stash_overlay(server_args): def _stash_overlay(server_args):
"""What the declarations say, last writer wins -- the projection's input.""" """What the declarations say, last writer wins -- the projection's input."""
overlay = {} overlay = {}
@@ -345,35 +348,43 @@ class TestResolutionDeclarations(CustomTestCase):
+ "\n ".join(unexplained), + "\n ".join(unexplained),
) )
def test_the_projection_input_is_the_resolved_configuration(self): def test_a_declaration_only_resolver_leaves_the_field_alone(self):
"""What the bags are built from equals what the record ends up holding. """The direction of travel: resolution decides, the record does not move.
The projection reads `raw input + declarations` rather than the A resolver that only declares -- a model-specific override, a registry
fields, so that it keeps working when the declarations stop entry -- writes nothing onto the record. The projection carries its
materializing. While they still do, the two have to agree leaf for answer and the field still holds what the caller passed.
leaf -- a difference means the projection would publish something the
record does not say, which is the failure this whole transition is
meant to avoid.
""" """
from sglang.srt.arg_groups.arg_utils import namespace_of from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.overrides import resolution_result
differences = [] found = []
for shape in _SHAPES: for shape in _SHAPES:
server_args = self._resolve(shape) server_args = self._resolve(shape)
raw = getattr(server_args, "_raw_input", None) or {}
for field in namespace_of(type(server_args)): for field in namespace_of(type(server_args)):
projected = resolution_result(server_args, field) if field not in raw:
continue
decided = resolution_result(server_args, field)
on_record = getattr(server_args, field) on_record = getattr(server_args, field)
if projected != on_record: if decided == on_record:
differences.append( continue
f"{shape} -> {field}: projection={projected!r} " # It moved away from the record's value, so the record must
f"record={on_record!r}" # still hold exactly what the caller passed.
) self.assertEqual(
self.assertEqual( on_record,
differences, raw[field],
f"{shape} -> {field}: the record holds {on_record!r}, which "
f"is neither the raw input {raw[field]!r} nor what "
f"resolution decided ({decided!r})",
)
found.append((shape_key(shape), field))
self.assertNotEqual(
found,
[], [],
"the projection and the record disagree about a config leaf:\n " "no field is resolved by declaration alone any more, so this check "
+ "\n ".join(differences), "no longer covers anything -- either the shapes stopped reaching "
"one or the declarations are writing the fields again",
) )
def test_the_whole_object_readback_carries_only_fields(self): def test_the_whole_object_readback_carries_only_fields(self):
@@ -561,10 +572,9 @@ class TestResolutionDeclarations(CustomTestCase):
The parser detection and the LoRA normalization run at launcher stage -- The parser detection and the LoRA normalization run at launcher stage --
they need a tokenizer, a chat template, an adapter directory -- and they they need a tokenizer, a chat template, an adapter directory -- and they
write through `declare_late_resolution`. If those writes only reached declare through `declare_late_resolution`. The declaration is the only
the fields, the bags would describe the *unresolved* value: a server home for what they decide: the record keeps `--reasoning-parser auto`,
launched with `--reasoning-parser auto` would advertise and apply and the bags a process publishes carry the detected parser.
`auto` after detection had already replaced it.
A real model path, not the dummy one: a dummy record never materializes, A real model path, not the dummy one: a dummy record never materializes,
so its `resolve_once` re-runs and re-snapshots the raw input from so its `resolve_once` re-runs and re-snapshots the raw input from
@@ -586,15 +596,22 @@ class TestResolutionDeclarations(CustomTestCase):
) )
publish(server_args, role="tokenizer") publish(server_args, role="tokenizer")
self.assertEqual(get_serving().reasoning_parser, "qwen3") self.assertEqual(get_serving().reasoning_parser, "qwen3")
self.assertEqual(server_args.reasoning_parser, get_serving().reasoning_parser) self.assertEqual(
server_args.reasoning_parser,
"auto",
"the record is the operator's input; late resolution declares, it "
"does not write back",
)
def test_validation_can_still_resolve_before_the_record_is_published(self): def test_validation_can_still_resolve_before_the_record_is_published(self):
"""The LoRA checks normalize in place, so they must precede publish. """The LoRA checks resolve, so they must precede publish.
`check_server_args` is not read-only: it infers `enable_lora`, parses `check_server_args` is not read-only: it infers `enable_lora`, parses
adapter paths and normalizes target modules through late resolution, adapter paths and normalizes target modules through late resolution,
which a published record refuses. The launcher order is what keeps this which a published record refuses. The launcher order is what keeps this
legal, and this is the assertion that notices if it moves. legal, and this is the assertion that notices if it moves. What those
declarations decide reaches the bags; the record keeps the raw form the
operator passed.
""" """
from sglang.srt.runtime_context import get_lora, publish, reset_context from sglang.srt.runtime_context import get_lora, publish, reset_context
@@ -608,9 +625,17 @@ class TestResolutionDeclarations(CustomTestCase):
self.addCleanup(reset_context) self.addCleanup(reset_context)
server_args.check_server_args() server_args.check_server_args()
publish(server_args, role="tokenizer") publish(server_args, role="tokenizer")
self.assertEqual(get_lora().enable_lora, server_args.enable_lora)
self.assertEqual( self.assertEqual(
get_lora().lora_target_modules, server_args.lora_target_modules get_lora().enable_lora, resolution_result(server_args, "enable_lora")
)
self.assertEqual(
get_lora().lora_target_modules,
resolution_result(server_args, "lora_target_modules"),
)
self.assertEqual(
server_args.lora_target_modules,
["q_proj"],
"normalization is a declaration; the record keeps what was passed",
) )
def test_the_launcher_finishes_resolving_before_it_publishes(self): def test_the_launcher_finishes_resolving_before_it_publishes(self):
@@ -661,24 +686,37 @@ class TestResolutionDeclarations(CustomTestCase):
f"written:\n " + "\n ".join(too_late), f"written:\n " + "\n ".join(too_late),
) )
def test_the_stash_agrees_with_the_fields_it_declared(self): def test_an_undeclared_field_still_holds_the_raw_input(self):
mismatches = [] """Nothing writes a field behind the stash's back.
Comparing the stash against `resolution_result` would agree by
construction -- both are the same last-writer-wins walk over
`_resolved_overrides`, spelled forwards and backwards. The independent
source is the record's own `_raw_input` snapshot: a field with no
declaration has to still equal what the caller passed, because the only
sanctioned way to move one is to declare it.
"""
moved = []
for shape in _SHAPES: for shape in _SHAPES:
server_args = self._resolve(shape) server_args = self._resolve(shape)
overlay = _stash_overlay(server_args) overlay = _stash_overlay(server_args)
for field, declared in overlay.items(): raw_input = getattr(server_args, "_raw_input", None)
if field not in _RESOLVED_FIELDS: self.assertTrue(raw_input, f"{shape}: the record kept no raw snapshot")
for field in dataclasses.fields(server_args):
name = field.name
if name in overlay or name not in raw_input:
continue continue
actual = getattr(server_args, field) current = getattr(server_args, name, None)
if actual != declared: if current != raw_input[name]:
mismatches.append( moved.append(
f"{shape} -> {field}: field={actual!r} stash={declared!r}" f"{shape} -> {name}: raw={raw_input[name]!r} "
f"field={current!r}"
) )
self.assertEqual( self.assertEqual(
mismatches, moved,
[], [],
"a declared field and its stash entry disagree, so something " "these fields moved without a declaration, so the bags publish one "
"assigned the field behind the declaration:\n " + "\n ".join(mismatches), "value while the record shows another:\n " + "\n ".join(moved),
) )
def test_no_immediate_writer_overrides_a_deferred_one(self): def test_no_immediate_writer_overrides_a_deferred_one(self):
@@ -768,31 +768,43 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
server_args.resolve_once() server_args.resolve_once()
return server_args return server_args
def test_a_bare_replace_would_resolve_a_second_time(self): def test_a_bare_replace_resolves_again_and_lands_in_the_same_place(self):
"""Why the helper exists. If this stops drifting, the pipeline became """A bare copy resolves to the same place: the fields are the raw input.
idempotent and the helper's reason is gone -- read it again before
deleting either.""" `dataclasses.replace` copies the fields, so a bare copy re-runs
resolution over the *same input* the parent got -- the DP-attention
halving and the conservativeness scaling apply once. `replace_resolved`
buys something else: it carries the parent's declarations and its
`model_config`, so the copy answers without resolving at all.
"""
parent = self._resolved() parent = self._resolved()
bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000") bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000")
self.assertFalse( self.assertFalse(
getattr(bare, "_declarations_materialized", False), getattr(bare, "_resolution_finished", False),
"a bare replace carried the flag; then this test proves nothing", "a bare replace carried the flag; then this test proves nothing",
) )
bare.resolve_once() bare.resolve_once()
drifted = {
field.name: (
resolution_result(parent, field.name),
resolution_result(bare, field.name),
)
for field in dataclasses.fields(parent)
if field.name not in ("dist_init_addr", "random_seed")
and repr(resolution_result(parent, field.name))
!= repr(resolution_result(bare, field.name))
}
self.assertEqual( self.assertEqual(
(bare.chunked_prefill_size, round(bare.schedule_conservativeness, 4)), drifted,
( {},
parent.chunked_prefill_size // 2, "resolving a bare copy landed somewhere else, so the pipeline is "
round(parent.schedule_conservativeness * 0.3, 4), "reading its own output again",
),
"the second pass no longer drifts; this is the drift the copy "
"helper exists to avoid",
) )
def test_replace_resolved_keeps_the_parents_resolution(self): def test_replace_resolved_keeps_the_parents_resolution(self):
parent = self._resolved() parent = self._resolved()
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000") copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
self.assertTrue(getattr(copy_, "_declarations_materialized", False)) self.assertTrue(getattr(copy_, "_resolution_finished", False))
drifted = { drifted = {
field.name: (getattr(parent, field.name), getattr(copy_, field.name)) field.name: (getattr(parent, field.name), getattr(copy_, field.name))
for field in dataclasses.fields(parent) for field in dataclasses.fields(parent)
@@ -0,0 +1,612 @@
"""Resolution reads its own decisions, not the record's fields.
`declare_resolution` records a decision in the declaration stash and writes
nothing. The fields keep what the caller passed, so a resolver that reads a
field another resolver may have decided reads the raw input -- silently, and
only on the configurations where that other resolver fires. The whole pipeline
therefore reads through `resolving_view` (or `ServerArgs._resolved()`, which is
the same view spelled as the record's own member), and this pins that there is
nothing left reading a field directly.
Subjects: every function in `arg_groups/` that takes a config, every
`ServerArgs` handler the dispatcher reaches, and every member of `ServerArgs` /
`PortArgs` -- the members are reached from the hooks and from business code,
which the handler walk cannot see, and a member that recomputes from a raw field
decides from what was typed. All three
are derived -- a new hook file, a new handler or a new member is covered the
moment it is written. Readers *outside* those
two -- the platform defaults, `ModelConfig`, the spec-algo hook -- are reached by
resolution too and have moved to the view as well, but enumerating them needs
the call-graph derivation `test_resolution_reads_no_bag` owns; this file pins
the two scopes it can derive exactly.
"""
import ast
import dataclasses
import pathlib
import re
import sglang
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt"
_FIELDS = frozenset(field.name for field in dataclasses.fields(ServerArgs))
# Names a config travels under. `args` is included because the platform hooks
# use it; a false positive would be a function taking an argparse Namespace and
# reading an attribute that happens to be a ServerArgs field name, which the
# allowlist below would then have to carry.
_HOLDER_NAMES = frozenset({"server_args", "sa", "args"})
def _holders(fn):
names = {
arg.arg
for arg in list(fn.args.posonlyargs)
+ list(fn.args.args)
+ list(fn.args.kwonlyargs)
if arg.arg in _HOLDER_NAMES
}
for arg in (
list(fn.args.posonlyargs) + list(fn.args.args) + list(fn.args.kwonlyargs)
):
annotation = arg.annotation
text = (
annotation.value
if isinstance(annotation, ast.Constant)
else (
annotation.id
if isinstance(annotation, ast.Name)
else annotation.attr if isinstance(annotation, ast.Attribute) else None
)
)
if text == "ServerArgs":
names.add(arg.arg)
return names
def _field_reads(fn, holders):
for node in ast.walk(fn):
if (
isinstance(node, ast.Attribute)
and node.attr in _FIELDS
and isinstance(node.value, ast.Name)
and node.value.id in holders
and isinstance(node.ctx, ast.Load)
):
yield node.lineno, node.attr
def _resolution_handlers():
"""The `ServerArgs` methods the dispatcher reaches, transitively."""
tree = ast.parse((_SRT / "server_args.py").read_text(encoding="utf-8-sig"))
cls = next(
node
for node in ast.walk(tree)
if isinstance(node, ast.ClassDef) and node.name == "ServerArgs"
)
methods = {
node.name: node
for node in cls.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
}
assert "_run_resolution_pipeline" in methods, "the dispatcher was renamed"
seen, stack = set(), ["_run_resolution_pipeline"]
while stack:
name = stack.pop()
if name in seen or name not in methods:
continue
seen.add(name)
for node in ast.walk(methods[name]):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and isinstance(node.func.value, ast.Name)
and node.func.value.id == "self"
):
stack.append(node.func.attr)
return {name: methods[name] for name in seen}
_DECLARERS = frozenset(
{
"_declare",
"declare_resolution",
"declare_late_resolution",
"declare_direct_writes",
}
)
def _declared_fields():
"""The fields resolution decides, read off every shape that reaches the stash.
A keyword on a `declare_*` call is only one shape: the model-override and
post-process passes build a mapping instead (`MODEL_OVERRIDES` literals,
`overrides["dtype"] = ...`, a returned dict), and late resolution splats a
variable-keyed one. Deriving from keywords alone leaves nineteen fields
outside the subject set, `dtype` and `reasoning_parser` among them.
"""
fields = set()
# The declaration calls live wherever a resolver does; the mapping channels
# only exist where the override providers and post-process passes are.
keyword_sources = [_SRT / "server_args.py"]
for sub in ("arg_groups", "hardware_backend", "parser"):
keyword_sources += sorted((_SRT / sub).rglob("*.py"))
mapping_sources = {_SRT / "server_args.py", *(_SRT / "arg_groups").rglob("*.py")}
for path in keyword_sources:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
for node in ast.walk(tree):
# 1. `declare_resolution(sa, src, page_size=64)` and its siblings
if isinstance(node, ast.Call):
name = (
node.func.id
if isinstance(node.func, ast.Name)
else getattr(node.func, "attr", None)
)
if name in _DECLARERS:
for keyword in node.keywords:
if keyword.arg:
fields.add(keyword.arg)
elif isinstance(keyword.value, ast.Dict):
fields.update(_string_keys(keyword.value))
# 2. every mapping literal in the files that declare through one:
# the MODEL_OVERRIDES tables, the dicts the override providers
# return, the ones the post-process passes build. Scanning
# unrelated files here would collect a plain kwarg dict
# (`tokenizer_config={"trust_remote_code": ...}`) and turn a
# passthrough read into a violation.
if isinstance(node, ast.Dict) and path in mapping_sources:
fields.update(_string_keys(node))
# 3. `overrides["field"] = ...`
if (
path in mapping_sources
and isinstance(node, ast.Assign)
and isinstance(node.targets[0], ast.Subscript)
and isinstance(node.targets[0].slice, ast.Constant)
and isinstance(node.targets[0].slice.value, str)
):
fields.add(node.targets[0].slice.value)
return frozenset(fields & _FIELDS)
def _string_keys(node: ast.Dict) -> set:
return {
key.value
for key in node.keys
if isinstance(key, ast.Constant) and isinstance(key.value, str)
}
def _record_members():
"""Every member of `ServerArgs` / `PortArgs`, by class and name."""
tree = ast.parse((_SRT / "server_args.py").read_text(encoding="utf-8-sig"))
members = {}
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef) and node.name in ("ServerArgs", "PortArgs"):
for member in node.body:
if isinstance(member, (ast.FunctionDef, ast.AsyncFunctionDef)):
members[f"{node.name}.{member.name}"] = member
return members
def _config_reading_helpers():
"""Module functions that load a decided field off the config they are handed.
A member that hands them `self`, or a call site that hands them a record,
reads the raw input through the callee -- the shape neither an attribute
scan nor a `getattr` scan can see, because the field name is spelled in the
helper and the record is spelled at the call site.
"""
decided = _declared_fields()
helpers = {}
sources = [_SRT / "server_args.py"] + sorted((_SRT / "arg_groups").rglob("*.py"))
for path in sources:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
for fn in ast.walk(tree):
if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
params = {
arg.arg
for arg in list(fn.args.posonlyargs)
+ list(fn.args.args)
+ list(fn.args.kwonlyargs)
} - {"self", "cls"}
if not params:
continue
reads = {
node.attr
for node in ast.walk(fn)
if isinstance(node, ast.Attribute)
and isinstance(node.value, ast.Name)
and node.value.id in params
and isinstance(node.ctx, ast.Load)
and node.attr in decided
}
if reads:
helpers[fn.name] = sorted(reads)
return helpers
# The accessors that hand back the process-global record itself. A helper that
# is handed one of these reads the raw input exactly as a bare `self` would.
_RECORD_ACCESSORS = frozenset({"get_server_args", "global_server_args"})
# `self._server_args` is the same record under a private name; the scan has to
# see it or a reader inside the context object escapes every shape above.
_RECORD_ATTR = re.compile(r"^_*(server_args|sa)$")
def _record_arguments(node, aliases=frozenset()):
"""The bare-record arguments of a call.
Four spellings reach a helper with a record: the bare name (`self`, `sa`),
an attribute (`runner.server_args`), the process-global accessor called
inline (`get_server_args()`), and a local bound to either of the last two
earlier in the same function.
"""
out = []
for arg in node.args:
if isinstance(arg, ast.Name) and arg.id in ("self", "server_args", "sa"):
out.append(arg.id)
elif isinstance(arg, ast.Attribute) and _RECORD_ATTR.match(arg.attr or ""):
out.append(ast.unparse(arg))
elif (
isinstance(arg, ast.Call)
and isinstance(arg.func, ast.Name)
and arg.func.id in _RECORD_ACCESSORS
):
out.append(ast.unparse(arg))
elif isinstance(arg, ast.Name) and arg.id in aliases:
out.append(arg.id)
return out
def _record_aliases(function):
"""Locals bound to the record under a name of their own.
`_sa = getattr(runner, "server_args", None)`, `cfg = get_server_args()` and
`engine_args = ServerArgs.from_cli_args(args)` all put the record behind a
name the argument scan does not recognise, so a later
`getattr(_sa, "<decided leaf>")` reads what the operator typed.
"""
aliases = set()
for node in ast.walk(function):
if not isinstance(node, ast.Assign) or len(node.targets) != 1:
continue
target = node.targets[0]
if not isinstance(target, ast.Name):
continue
value = node.value
if isinstance(value, ast.Attribute):
if _RECORD_ATTR.match(value.attr or ""):
aliases.add(target.id)
continue
if not isinstance(value, ast.Call):
continue
func = value.func
if isinstance(func, ast.Name):
if func.id in _RECORD_ACCESSORS or func.id == "ServerArgs":
aliases.add(target.id)
elif (
func.id == "getattr"
and len(value.args) >= 2
and isinstance(value.args[1], ast.Constant)
and _RECORD_ATTR.match(str(value.args[1].value))
):
aliases.add(target.id)
elif (
isinstance(func, ast.Attribute)
and func.attr in ("from_cli_args", "replace_resolved")
and isinstance(func.value, ast.Name)
and func.value.id == "ServerArgs"
):
aliases.add(target.id)
return aliases
# The one reader for which the raw field is the right answer. The gateway sizes
# its worker pool from the operator's requested replica count; `--dwdp-size`
# makes resolution declare a `dp_size` describing one multi-rank server's
# internal topology, so reading the decision there would spawn dp_size
# single-rank children and ask for dp_size^2 GPUs. A new entry here needs that
# kind of reason next to it.
_NO_RESOLVED_SURFACE = frozenset(
{
(
"sgl-model-gateway/bindings/python/src/sglang_router/launch_server.py",
"server_args.dp_size",
),
}
)
def _is_record_base(node, aliases):
"""Is this expression the record itself?
A local bound to one, a parameter that carries one (`server_args`, `sa`,
`engine_args`), or an attribute holding one (`self._server_args`).
"""
if isinstance(node, ast.Name):
return node.id in aliases
if isinstance(node, ast.Attribute):
return bool(_RECORD_ATTR.match(node.attr or ""))
return False
def _record_handoff_offenders(rel, tree, helpers, decided, is_record=False):
"""Every way a decided leaf is reached through a record in one module.
Two shapes, both scanned under the record aliases the function binds:
handing the record to a helper that loads a decided field, and loading one
off the alias directly (`alias.<leaf>` or `getattr(alias, "<leaf>")`). The
second is what the MiniMax backend spelled, and an argument scan cannot see
it -- the leaf never appears at a call site.
"""
offenders, seen = [], set()
def record(lineno, text):
if (lineno, text) in seen:
return
seen.add((lineno, text))
offenders.append(f"{rel}:{lineno} {text}")
scopes = [(tree, frozenset())] + [
(fn, _record_aliases(fn))
for fn in ast.walk(tree)
if isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef))
]
for scope, aliases in scopes:
for node in ast.walk(scope):
if isinstance(node, ast.Call):
name = (
node.func.id
if isinstance(node.func, ast.Name)
else getattr(node.func, "attr", None)
)
if name in helpers:
for arg in _record_arguments(node, aliases):
if arg == "self" and not is_record:
continue
record(
node.lineno,
f"{name}({arg}) reads {', '.join(helpers[name])}",
)
if (
isinstance(node.func, ast.Name)
and node.func.id == "getattr"
and len(node.args) >= 2
and isinstance(node.args[0], ast.Name)
and node.args[0].id in aliases
and isinstance(node.args[1], ast.Constant)
and node.args[1].value in decided
):
record(
node.lineno,
f'getattr({node.args[0].id}, "{node.args[1].value}")',
)
elif (
isinstance(node, ast.Attribute)
and isinstance(node.ctx, ast.Load)
and node.attr in decided
and _is_record_base(node.value, aliases)
):
record(node.lineno, f"{ast.unparse(node.value)}.{node.attr}")
return offenders
# Source the scanner must read the same way whether or not the tree happens to
# contain these shapes today. The first four are the spellings that reached
# production and were converted; the last two are the legal forms next to them,
# which have to stay quiet or the guard is unusable.
_SPELLINGS = """
def hands_the_alias_to_a_helper(runner):
_sa = getattr(runner, "server_args", None)
return m3_fp8_attn_gemm_enabled(_sa)
def loads_a_leaf_off_the_alias(runner):
_sa = getattr(runner, "server_args", None)
return getattr(_sa, "speculative_num_draft_tokens", None)
def reads_a_leaf_through_the_alias(runner):
sa_local = runner.server_args
return sa_local.attention_backend
def hands_the_accessor_to_a_helper():
return compute_world_size(get_server_args())
def reads_the_view(runner):
cfg = resolving_view(runner.server_args)
return cfg.attention_backend
def reads_an_undecided_leaf(runner):
_sa = runner.server_args
return _sa.tp_size
def reads_a_leaf_off_a_private_attribute(self):
return self._server_args.attention_backend
def reads_a_leaf_off_a_constructed_record(cli):
engine_args = ServerArgs.from_cli_args(cli)
engine_args.resolve_once()
return engine_args.attention_backend
"""
class TestResolutionReadsTheDeclarations(CustomTestCase):
def test_no_hook_reads_a_field_off_the_record(self):
offenders = []
files = sorted((_SRT / "arg_groups").glob("*.py"))
self.assertGreater(len(files), 5, "the hook scan found almost nothing")
for path in files:
rel = f"arg_groups/{path.name}"
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
for fn in ast.walk(tree):
if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
holders = _holders(fn)
if not holders:
continue
for lineno, field in _field_reads(fn, holders):
offenders.append(f"{rel}:{lineno} {fn.name} reads .{field}")
self.assertEqual(
offenders,
[],
"a resolution hook reads a field off the record; the field holds the "
"raw input, so this decides from what was typed rather than from "
"what resolution decided. Read `resolving_view(server_args)`:\n "
+ "\n ".join(offenders),
)
def test_no_handler_reads_a_field_off_self(self):
handlers = _resolution_handlers()
self.assertGreater(
len(handlers), 50, f"only {len(handlers)} handlers were reached"
)
offenders = []
for name, fn in sorted(handlers.items()):
for lineno, field in _field_reads(fn, {"self"}):
offenders.append(f"server_args.py:{lineno} {name} reads self.{field}")
self.assertEqual(
offenders,
[],
"a resolution handler reads its own field; the field holds the raw "
"input. Bind `cfg = resolving_view(self)` and read that:\n "
+ "\n ".join(offenders),
)
def test_no_member_recomputes_from_a_raw_field(self):
decided = _declared_fields()
self.assertGreater(
len(decided), 100, f"the declaration set derived only {len(decided)} fields"
)
members = _record_members()
self.assertGreater(len(members), 100, f"only {len(members)} members were found")
offenders = []
for name, fn in sorted(members.items()):
holders = _holders(fn) | {"self"}
for lineno, field in _field_reads(fn, holders):
if field in decided:
offenders.append(f"server_args.py:{lineno} {name} reads .{field}")
self.assertEqual(
offenders,
[],
"a record member recomputes from a field resolution decides; the "
"field holds the raw input, so the member answers for what was "
"typed. Bind `cfg = resolving_view(self)` and read that:\n "
+ "\n ".join(offenders),
)
def test_no_reader_hands_the_record_to_a_config_helper(self):
helpers = _config_reading_helpers()
self.assertGreater(
len(helpers), 5, f"the helper derivation found only {len(helpers)}"
)
decided = _declared_fields()
offenders = []
# `scripts/`, `examples/` and the gateway binding are outside the
# package but hold records they resolve themselves, and every reader
# this scan found in them was reading a field resolution fills in.
_REPO = _SRT.parent.parent.parent
roots = (
[_SRT]
+ [_SRT.parent / d for d in ("benchmark", "lang")]
+ [
_REPO / d
for d in (
"scripts",
"examples",
"sgl-model-gateway/bindings/python/src",
)
]
)
for root in roots:
if not root.exists():
continue
for path in sorted(root.rglob("*.py")):
# Package files keep their `srt/...` spelling (the skips below
# key on it); the repo-level roots are named from the repo.
try:
rel = path.relative_to(_SRT.parent).as_posix()
except ValueError:
rel = path.relative_to(_REPO).as_posix()
if rel.startswith(("srt/arg_groups/", "multimodal_gen/")):
continue
try:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
except SyntaxError:
continue
offenders += _record_handoff_offenders(
rel, tree, helpers, decided, is_record=rel == "srt/server_args.py"
)
offenders = [
line
for line in offenders
if (line.split(":", 1)[0], line.split(" ", 1)[1])
not in _NO_RESOLVED_SURFACE
]
self.assertEqual(
offenders,
[],
"a caller hands the record to a helper that loads a field "
"resolution decides; the helper then reads the raw input. Hand it "
"`resolving_view(record)` (or the published bag):\n "
+ "\n ".join(offenders),
)
def test_the_scan_sees_every_spelling_that_reached_production(self):
"""Every spelling that reached production, pinned next to the scanner.
A shape the scan stops seeing is a silent hole, so each one is listed
here with the legal forms beside it and the flagged set compared
exactly.
What it does not reach: a record that arrives as a *parameter* and was
resolved by the caller (`scripts/playground/bench_speculative.py` hands
`main(args, server_args)` one). Binding that would need the call graph,
and naming a parameter `server_args` is also how the resolution-time
readers spell a view.
"""
helpers = _config_reading_helpers()
decided = _declared_fields()
for name in ("m3_fp8_attn_gemm_enabled", "compute_world_size"):
self.assertIn(name, helpers, f"the helper derivation lost {name}")
for field in ("speculative_num_draft_tokens", "attention_backend"):
self.assertIn(field, decided, f"the declared set lost {field}")
offenders = _record_handoff_offenders(
"sample.py", ast.parse(_SPELLINGS), helpers, decided
)
flagged = {line.split(" ", 1)[1] for line in offenders}
self.assertEqual(
flagged,
{
"m3_fp8_attn_gemm_enabled(_sa)"
" reads " + ", ".join(helpers["m3_fp8_attn_gemm_enabled"]),
'getattr(_sa, "speculative_num_draft_tokens")',
"sa_local.attention_backend",
"self._server_args.attention_backend",
"engine_args.attention_backend",
"compute_world_size(get_server_args())"
" reads " + ", ".join(helpers["compute_world_size"]),
},
"the scan lost a spelling, or started flagging a legal one:\n "
+ "\n ".join(sorted(flagged)),
)
if __name__ == "__main__":
import unittest
unittest.main()
+102 -53
View File
@@ -20,6 +20,7 @@ from sglang.srt.arg_groups.arg_utils import A, Arg, resolvable_fields
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
collect_model_override_declarations, collect_model_override_declarations,
register_model_override, register_model_override,
resolution_result,
validate_declarations, validate_declarations,
) )
from sglang.srt.configs.minicpm import MiniCPMHybridConfig from sglang.srt.configs.minicpm import MiniCPMHybridConfig
@@ -294,9 +295,28 @@ class TestPublishInstallsSlot(_IsolatedPublish):
class TestGoldenModelOverrides(_IsolatedPublish): class TestGoldenModelOverrides(_IsolatedPublish):
"""Per-arch golden diff for migrated families: the declarative path must """Per-arch golden diff for migrated families: the declarative path must
reproduce the legacy imperative writes byte-identically on the reproduce the legacy imperative writes byte-identically in the resolution
materialized server_args fields; the publish round-trip returns the same result; the publish round-trip returns the same object.
object."""
`_resolved` is how the assertions read it. A model-specific override only
declares -- it does not write the field -- so the record keeps what the
caller passed and the projection carries the override.
"""
def _resolved(self, server_args, field):
from sglang.srt.arg_groups.overrides import resolution_result
return resolution_result(server_args, field)
def _leaf(self, field):
"""The published value of `field`, whichever bag owns it.
The publish round-trip is checked on the bags: the record the process
publishes is the raw input, and the leaf is what every reader reads.
"""
from sglang.srt.runtime_context import get_context
return get_context().config_leaf(field)
_MINI_CONFIG = { _MINI_CONFIG = {
"hidden_size": 64, "hidden_size": 64,
@@ -571,40 +591,42 @@ class TestGoldenModelOverrides(_IsolatedPublish):
def test_mistral_large3_forces_bfloat16(self): def test_mistral_large3_forces_bfloat16(self):
sa = self._construct("MistralLarge3ForCausalLM", "mistral") sa = self._construct("MistralLarge3ForCausalLM", "mistral")
self.assertEqual(sa.dtype, "bfloat16") # materialized at end of resolution self.assertEqual(
self._resolved(sa, "dtype"), "bfloat16"
) # materialized at end of resolution
self.assertIn( self.assertIn(
("MODEL_OVERRIDES['MistralLarge3ForCausalLM']", {"dtype": "bfloat16"}), ("MODEL_OVERRIDES['MistralLarge3ForCausalLM']", {"dtype": "bfloat16"}),
sa._resolved_overrides, sa._resolved_overrides,
) )
self.assertEqual(self._publish(sa).dtype, "bfloat16") self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16")
def test_user_requested_dtype_is_still_overridden(self): def test_user_requested_dtype_is_still_overridden(self):
# Legacy fidelity: the arch branch overwrote dtype unconditionally, # Legacy fidelity: the arch branch overwrote dtype unconditionally,
# so the declaration must too. The pristine request survives on # so the declaration must too. The request survives on the record; the
# provenance; the materialized field carries the override. # projection carries the override.
sa = self._construct("MistralLarge3ForCausalLM", "mistral", dtype="float16") sa = self._construct("MistralLarge3ForCausalLM", "mistral", dtype="float16")
self.assertEqual(sa.dtype, "bfloat16") # materialized self.assertEqual(self._resolved(sa, "dtype"), "bfloat16")
self.assertEqual(self._publish(sa).dtype, "bfloat16") self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16")
def test_control_arch_keeps_pristine_dtype(self): def test_control_arch_keeps_pristine_dtype(self):
sa = self._construct("LlamaForCausalLM", "llama") sa = self._construct("LlamaForCausalLM", "llama")
self.assertEqual(sa.dtype, "auto") self.assertEqual(self._resolved(sa, "dtype"), "auto")
declared = {f for _s, d in sa._resolved_overrides for f in d} declared = {f for _s, d in sa._resolved_overrides for f in d}
self.assertNotIn("dtype", declared) # no arch declaration for Llama self.assertNotIn("dtype", declared) # no arch declaration for Llama
# publish still materializes the whitelisted leaf with the pristine # publish still projects the whitelisted leaf with the pristine
# value: readers only ever read flags. # value: readers only ever read flags.
self.assertEqual(self._publish(sa).dtype, "auto") self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "auto")
def test_minimax_m2_enables_tf32_matmul(self): def test_minimax_m2_enables_tf32_matmul(self):
sa = self._construct("MiniMaxM2ForCausalLM", "llama") sa = self._construct("MiniMaxM2ForCausalLM", "llama")
self.assertTrue(sa.enable_tf32_matmul) # materialized self.assertTrue(self._resolved(sa, "enable_tf32_matmul"))
self.assertIn( self.assertIn(
("_minimax_m2_overrides", {"enable_tf32_matmul": True}), ("_minimax_m2_overrides", {"enable_tf32_matmul": True}),
sa._resolved_overrides, sa._resolved_overrides,
) )
flags = self._publish(sa) flags = self._publish(sa)
self.assertTrue(flags.enable_tf32_matmul) self.assertTrue(self._leaf("enable_tf32_matmul"))
self.assertFalse(flags.enable_multi_layer_eagle) # pristine materialize self.assertFalse(self._leaf("enable_multi_layer_eagle")) # the pristine value
def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self): def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self):
"""MiniMax-M2 NVFP4 auto must avoid the unsupported plain TRT-LLM path.""" """MiniMax-M2 NVFP4 auto must avoid the unsupported plain TRT-LLM path."""
@@ -622,10 +644,14 @@ class TestGoldenModelOverrides(_IsolatedPublish):
"MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4" "MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4"
) )
self.assertEqual(explicit.moe_runner_backend, "flashinfer_cutlass") self.assertEqual(
self.assertEqual(non_nvfp4.moe_runner_backend, "auto") self._resolved(explicit, "moe_runner_backend"), "flashinfer_cutlass"
self.assertEqual(nvfp4.moe_runner_backend, "flashinfer_trtllm_routed") )
self.assertTrue(nvfp4.disable_shared_experts_fusion) self.assertEqual(self._resolved(non_nvfp4, "moe_runner_backend"), "auto")
self.assertEqual(
self._resolved(nvfp4, "moe_runner_backend"), "flashinfer_trtllm_routed"
)
self.assertTrue(self._resolved(nvfp4, "disable_shared_experts_fusion"))
self.assertIn( self.assertIn(
( (
"_minimax_m2_overrides", "_minimax_m2_overrides",
@@ -649,7 +675,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
non_sm10x = self._construct( non_sm10x = self._construct(
"MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4" "MiniMaxM2ForCausalLM", "llama", quantization="modelopt_fp4"
) )
self.assertEqual(non_sm10x.moe_runner_backend, "auto") self.assertEqual(self._resolved(non_sm10x, "moe_runner_backend"), "auto")
self._publish(nvfp4) self._publish(nvfp4)
self.assertEqual(get_exec().moe.moe_runner_backend, "flashinfer_trtllm_routed") self.assertEqual(get_exec().moe.moe_runner_backend, "flashinfer_trtllm_routed")
@@ -1000,25 +1026,25 @@ class TestGoldenModelOverrides(_IsolatedPublish):
enable_hierarchical_cache=True, enable_hierarchical_cache=True,
) )
# materialized at the end of resolution # materialized at the end of resolution
self.assertEqual(sa.swa_full_tokens_ratio, 1.0) self.assertEqual(self._resolved(sa, "swa_full_tokens_ratio"), 1.0)
self.assertTrue(sa.disable_hybrid_swa_memory) self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory"))
flags = self._publish(sa) flags = self._publish(sa)
self.assertEqual(flags.swa_full_tokens_ratio, 1.0) self.assertEqual(self._leaf("swa_full_tokens_ratio"), 1.0)
self.assertTrue(flags.disable_hybrid_swa_memory) self.assertTrue(self._leaf("disable_hybrid_swa_memory"))
def test_gemma2_disables_hybrid_swa_memory(self): def test_gemma2_disables_hybrid_swa_memory(self):
sa = self._construct("Gemma2ForCausalLM", "llama") sa = self._construct("Gemma2ForCausalLM", "llama")
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized
self.assertIn( self.assertIn(
("_gemma2_gemma3_overrides", {"disable_hybrid_swa_memory": True}), ("_gemma2_gemma3_overrides", {"disable_hybrid_swa_memory": True}),
sa._resolved_overrides, sa._resolved_overrides,
) )
self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1])
def test_olmo2_disables_hybrid_swa_memory(self): def test_olmo2_disables_hybrid_swa_memory(self):
sa = self._construct("Olmo2ForCausalLM", "llama") sa = self._construct("Olmo2ForCausalLM", "llama")
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized
self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1])
def test_exaone_conditional_on_sliding_window_pattern(self): def test_exaone_conditional_on_sliding_window_pattern(self):
# With the pattern the branch also asserts an explicit backend. # With the pattern the branch also asserts an explicit backend.
@@ -1028,8 +1054,8 @@ class TestGoldenModelOverrides(_IsolatedPublish):
config_extra={"sliding_window_pattern": "LLLG"}, config_extra={"sliding_window_pattern": "LLLG"},
attention_backend="fa3", attention_backend="fa3",
) )
self.assertTrue(sa.disable_hybrid_swa_memory) # materialized self.assertTrue(self._resolved(sa, "disable_hybrid_swa_memory")) # materialized
self.assertTrue(self._publish(sa).disable_hybrid_swa_memory) self.assertTrue((self._publish(sa), self._leaf("disable_hybrid_swa_memory"))[1])
def test_exaone_without_pattern_declares_nothing(self): def test_exaone_without_pattern_declares_nothing(self):
from sglang.srt.arg_groups.overrides import _exaone_overrides from sglang.srt.arg_groups.overrides import _exaone_overrides
@@ -1051,13 +1077,13 @@ class TestGoldenModelOverrides(_IsolatedPublish):
"llama", "llama",
config_extra={"quantization_config": {"quant_method": "mxfp4"}}, config_extra={"quantization_config": {"quant_method": "mxfp4"}},
) )
self.assertEqual(sa.dtype, "bfloat16") # materialized self.assertEqual(self._resolved(sa, "dtype"), "bfloat16")
self.assertEqual(self._publish(sa).dtype, "bfloat16") self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "bfloat16")
def test_gpt_oss_without_mxfp4_keeps_pristine_dtype(self): def test_gpt_oss_without_mxfp4_keeps_pristine_dtype(self):
sa = self._construct("GptOssForCausalLM", "llama") sa = self._construct("GptOssForCausalLM", "llama")
self.assertEqual(sa.dtype, "auto") self.assertEqual(self._resolved(sa, "dtype"), "auto")
self.assertEqual(self._publish(sa).dtype, "auto") self.assertEqual((self._publish(sa), self._leaf("dtype"))[1], "auto")
def test_gpt_oss_xpu_dtype_validation_reads_pristine(self): def test_gpt_oss_xpu_dtype_validation_reads_pristine(self):
from sglang.srt.arg_groups.overrides import _gpt_oss_overrides from sglang.srt.arg_groups.overrides import _gpt_oss_overrides
@@ -1077,28 +1103,36 @@ class TestGoldenModelOverrides(_IsolatedPublish):
sa = self._construct("LlamaForCausalLM", "llama") sa = self._construct("LlamaForCausalLM", "llama")
expected = "flashinfer" if is_flashinfer_available() else "pytorch" expected = "flashinfer" if is_flashinfer_available() else "pytorch"
self.assertEqual(sa.sampling_backend, expected) # materialized self.assertEqual(
self._resolved(sa, "sampling_backend"), expected
) # materialized
self.assertIn( self.assertIn(
("_sampling_backend_default", {"sampling_backend": expected}), ("_sampling_backend_default", {"sampling_backend": expected}),
sa._resolved_overrides, sa._resolved_overrides,
) )
self.assertEqual(self._publish(sa).sampling_backend, expected) self.assertEqual(
(self._publish(sa), self._leaf("sampling_backend"))[1], expected
)
def test_sampling_backend_user_choice_survives(self): def test_sampling_backend_user_choice_survives(self):
sa = self._construct("LlamaForCausalLM", "llama", sampling_backend="pytorch") sa = self._construct("LlamaForCausalLM", "llama", sampling_backend="pytorch")
self.assertEqual(sa.sampling_backend, "pytorch") self.assertEqual(self._resolved(sa, "sampling_backend"), "pytorch")
# the pass declared nothing; publish materializes the pristine choice # the pass declared nothing; publish materializes the pristine choice
self.assertEqual(self._publish(sa).sampling_backend, "pytorch") self.assertEqual(
(self._publish(sa), self._leaf("sampling_backend"))[1], "pytorch"
)
def test_deterministic_inference_forces_pytorch_sampling(self): def test_deterministic_inference_forces_pytorch_sampling(self):
sa = self._construct( sa = self._construct(
"LlamaForCausalLM", "llama", enable_deterministic_inference=True "LlamaForCausalLM", "llama", enable_deterministic_inference=True
) )
# two pass writers chain: default fill, then the deterministic force — # two pass writers chain: default fill, then the deterministic force --
# last writer wins; materialization lands the end state on the fields. # last writer wins. The end state lives in the stash, which is what the
self.assertEqual(sa.sampling_backend, "pytorch") # projection reads and the bags are built from; the field still holds
# what the caller passed.
self.assertEqual(resolution_result(sa, "sampling_backend"), "pytorch")
flags = self._publish(sa) flags = self._publish(sa)
self.assertEqual(flags.sampling_backend, "pytorch") self.assertEqual(self._leaf("sampling_backend"), "pytorch")
# the deterministic attention fill declared a compatible backend and # the deterministic attention fill declared a compatible backend and
# the compatibility default-fill then had nothing to do # the compatibility default-fill then had nothing to do
deterministic_fills = [ deterministic_fills = [
@@ -1107,8 +1141,10 @@ class TestGoldenModelOverrides(_IsolatedPublish):
if source == "_deterministic_attention_backend" if source == "_deterministic_attention_backend"
] ]
self.assertEqual(len(deterministic_fills), 1) self.assertEqual(len(deterministic_fills), 1)
self.assertEqual(sa.attention_backend, deterministic_fills[0]) self.assertEqual(
self.assertEqual(flags.attention_backend, deterministic_fills[0]) resolution_result(sa, "attention_backend"), deterministic_fills[0]
)
self.assertEqual(self._leaf("attention_backend"), deterministic_fills[0])
def test_deterministic_incompatible_backend_raises(self): def test_deterministic_incompatible_backend_raises(self):
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
@@ -1148,13 +1184,17 @@ class TestGoldenModelOverrides(_IsolatedPublish):
disable_radix_cache=True, disable_radix_cache=True,
attention_backend="triton", attention_backend="triton",
) )
self.assertEqual(sa.attention_backend, "flashinfer") # materialized self.assertEqual(
self._resolved(sa, "attention_backend"), "flashinfer"
) # materialized
self.assertIn( self.assertIn(
("_dllm_attention_backend", {"attention_backend": "flashinfer"}), ("_dllm_attention_backend", {"attention_backend": "flashinfer"}),
sa._resolved_overrides, sa._resolved_overrides,
) )
# the deterministic fill lands on the attention_backend field # the deterministic fill lands on the attention_backend field
self.assertEqual(self._publish(sa).attention_backend, "flashinfer") self.assertEqual(
(self._publish(sa), self._leaf("attention_backend"))[1], "flashinfer"
)
def test_attention_backend_leaf_materializes_end_state(self): def test_attention_backend_leaf_materializes_end_state(self):
# The default-fill pass declares the platform-selected backend; the # The default-fill pass declares the platform-selected backend; the
@@ -1167,8 +1207,12 @@ class TestGoldenModelOverrides(_IsolatedPublish):
if "attention_backend" in d if "attention_backend" in d
] ]
self.assertTrue(declared_values) # default fill declared self.assertTrue(declared_values) # default fill declared
self.assertEqual(sa.attention_backend, declared_values[-1]) # materialized self.assertEqual(
self.assertEqual(self._publish(sa).attention_backend, declared_values[-1]) self._resolved(sa, "attention_backend"), declared_values[-1]
) # materialized
self.assertEqual(
(self._publish(sa), self._leaf("attention_backend"))[1], declared_values[-1]
)
def test_post_materialize_pass_writes_through(self): def test_post_materialize_pass_writes_through(self):
from sglang.srt.arg_groups.overrides import run_post_process_pass from sglang.srt.arg_groups.overrides import run_post_process_pass
@@ -1177,7 +1221,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
# legacy runner-side adjustments) declares AND writes through, so # legacy runner-side adjustments) declares AND writes through, so
# field readers and the publish see the same end state. # field readers and the publish see the same end state.
sa = self._construct("LlamaForCausalLM", "llama") sa = self._construct("LlamaForCausalLM", "llama")
resolved_before = sa.attention_backend resolved_before = self._resolved(sa, "attention_backend")
def _force_triton(view): def _force_triton(view):
if view.attention_backend != "triton": if view.attention_backend != "triton":
@@ -1186,13 +1230,18 @@ class TestGoldenModelOverrides(_IsolatedPublish):
run_post_process_pass(sa, _force_triton) run_post_process_pass(sa, _force_triton)
if resolved_before != "triton": if resolved_before != "triton":
self.assertEqual(sa.attention_backend, "triton") self.assertEqual(self._resolved(sa, "attention_backend"), "triton")
self.assertEqual(self._publish(sa).attention_backend, sa.attention_backend) self.assertEqual(
(self._publish(sa), self._leaf("attention_backend"))[1],
self._resolved(sa, "attention_backend"),
)
def test_attention_backend_user_choice_declares_nothing_extra(self): def test_attention_backend_user_choice_declares_nothing_extra(self):
sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton") sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton")
self.assertEqual(sa.attention_backend, "triton") self.assertEqual(self._resolved(sa, "attention_backend"), "triton")
self.assertEqual(self._publish(sa).attention_backend, "triton") self.assertEqual(
(self._publish(sa), self._leaf("attention_backend"))[1], "triton"
)
def test_compatibility_passes_at_callable_level(self): def test_compatibility_passes_at_callable_level(self):
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
@@ -11,6 +11,7 @@ from unittest import mock
from sglang.srt import runtime_context as rc from sglang.srt import runtime_context as rc
from sglang.srt.arg_groups.arg_utils import NS, A from sglang.srt.arg_groups.arg_utils import NS, A
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -96,16 +97,15 @@ class TestConfigBags(CustomTestCase):
none-flags) the first resolution may have written -- so the assertion none-flags) the first resolution may have written -- so the assertion
is "bag == what resolution produces", not "bag == the instance publish is "bag == what resolution produces", not "bag == the instance publish
copied from". Reproducibility (`test_resolution_is_reproducible`) copied from". Reproducibility (`test_resolution_is_reproducible`)
licenses the sibling as a stand-in for the pipeline's output. The licenses the sibling as a stand-in for the pipeline's output.
raw-differs guard keeps the comparison meaningful: every sampled leaf
must have moved off its dataclass default, so each equality compares a The reference's resolved values are read through `resolution_result`,
value resolution demonstrably wrote. Supplied construction inputs because a record holds the user's raw input: the decision lives in the
(`model_path`, `device`, `random_seed`) and leaves resolution leaves declarations, and the bags are where a process reads it. The
alone never enter the sample -- projection coverage for those lives in raw-differs guard keeps the comparison meaningful -- every sampled leaf
`test_passthrough_leaves_project_into_their_namespaces`. Step 12 keeps must have moved off its dataclass default -- and the last assertion is
records at the user's raw input; then the sibling goes raw and this the other half of that invariant: the record still answers the raw
assertion starts failing for every sampled leaf, which is the signal input for a leaf resolution decided.
the bags became the only home of the effective value.
""" """
import dataclasses import dataclasses
@@ -124,12 +124,13 @@ class TestConfigBags(CustomTestCase):
# The raw-differs guard: a sampled leaf that still sits on its # The raw-differs guard: a sampled leaf that still sits on its
# default (or has none to differ from) proves nothing. # default (or has none to differ from) proves nothing.
self.assertIsNot(defaults[leaf], dataclasses.MISSING) self.assertIsNot(defaults[leaf], dataclasses.MISSING)
self.assertNotEqual(getattr(reference, leaf), defaults[leaf]) resolved = resolution_result(reference, leaf)
self.assertEqual(accessor(), getattr(reference, leaf)) self.assertNotEqual(resolved, defaults[leaf])
# And the record agrees today, which is what step 12 changes: when this self.assertEqual(accessor(), resolved)
# assertion starts failing for a resolution-written leaf, the flip # The record is the raw input, so the field still reads as the default
# landed and the bag is the only place the effective value lives. # for a leaf the bag now answers for.
self.assertEqual(rc.get_schedule().page_size, sa.page_size) self.assertEqual(sa.page_size, defaults["page_size"])
self.assertNotEqual(rc.get_schedule().page_size, sa.page_size)
def test_passthrough_leaves_project_into_their_namespaces(self): def test_passthrough_leaves_project_into_their_namespaces(self):
"""Thin projection smoke over leaves resolution does not move. """Thin projection smoke over leaves resolution does not move.
@@ -141,7 +142,7 @@ class TestConfigBags(CustomTestCase):
sa = self._publish() sa = self._publish()
sampled = ( sampled = (
(lambda: rc.get_serving().host, "host"), (lambda: rc.get_serving().host, "host"),
(lambda: rc.get_memory().hicache_ratio, "hicache_ratio"), (lambda: rc.get_memory().hicache_write_policy, "hicache_write_policy"),
(lambda: rc.get_exec().moe.moe_runner_backend, "moe_runner_backend"), (lambda: rc.get_exec().moe.moe_runner_backend, "moe_runner_backend"),
(lambda: rc.get_model().model_path, "model_path"), (lambda: rc.get_model().model_path, "model_path"),
) )
@@ -230,8 +231,8 @@ class TestConfigBags(CustomTestCase):
rc.get_memory().hicache_ratio = 9.0 rc.get_memory().hicache_ratio = 9.0
def test_scoped_override_restores(self): def test_scoped_override_restores(self):
sa = self._publish() self._publish()
original = sa.hicache_ratio original = rc.get_memory().hicache_ratio
with rc.get_memory().override(hicache_ratio=original + 1.0): with rc.get_memory().override(hicache_ratio=original + 1.0):
self.assertEqual(rc.get_memory().hicache_ratio, original + 1.0) self.assertEqual(rc.get_memory().hicache_ratio, original + 1.0)
self.assertEqual(rc.get_memory().hicache_ratio, original) self.assertEqual(rc.get_memory().hicache_ratio, original)
@@ -110,7 +110,7 @@ class TestContextOverride(CustomTestCase):
# server_args is read-only after resolution: resolved config changes go # server_args is read-only after resolution: resolved config changes go
# to the bags, a per-runner config to a derived variant. # to the bags, a per-runner config to a derived variant.
sa = ServerArgs(model_path="dummy") sa = ServerArgs(model_path="dummy")
object.__setattr__(sa, "_declarations_materialized", True) object.__setattr__(sa, "_resolution_finished", True)
with self.assertRaises(AttributeError): with self.assertRaises(AttributeError):
sa.page_size = 999 sa.page_size = 999
@@ -6,9 +6,9 @@ resolution the fields are the record the config bags were projected from, so a
write desyncs every namespace reader, and a copy invites publishing stale write desyncs every namespace reader, and a copy invites publishing stale
variants. Both are gone: post-publish changes go to the bags variants. Both are gone: post-publish changes go to the bags
(``get_context().override``), a value one runner or worker owns travels as a (``get_context().override``), a value one runner or worker owns travels as a
constructor argument, and late launcher-stage resolution writes in place constructor argument, and late launcher-stage resolution declares through
through ``arg_groups.overrides.declare_late_resolution``, which refuses the ``arg_groups.overrides.declare_late_resolution``, which writes no field and
published instance. refuses the published instance.
The textual half of this guard matters because the resolution pipeline's own file The textual half of this guard matters because the resolution pipeline's own file
is exempt from the mutation ratchet: a ``self.override(...)`` there — exactly is exempt from the mutation ratchet: a ``self.override(...)`` there — exactly
@@ -535,9 +535,9 @@ class TestSuppliedInstanceExposure(CustomTestCase):
``MODEL_OVERRIDES`` maps arch -> {field: value}, and the ``MODEL_OVERRIDES`` maps arch -> {field: value}, and the
``@register_model_override``(-``_predicate``) providers return (or ``@register_model_override``(-``_predicate``) providers return (or
build by subscript) {field: value} dicts; ``materialize_declarations`` build by subscript) {field: value} dicts, which go straight into the
applies them all via setattr, so no assignment scan sees these writes declaration stash, so no assignment scan sees these writes and a
and a llama-only matrix never triggers them. Keys must be llama-only matrix never triggers them. Keys must be
string literals; anything else fails loudly. string literals; anything else fails loudly.
""" """
tree = ast.parse( tree = ast.parse(