[Config] One writer for the declaration stash; no exception to the write seal (#38752)
This commit is contained in:
@@ -91,8 +91,11 @@ with what the operator typed, not with what resolution decided.**
|
|||||||
- **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
|
tokenizer/chat-template load. They are resolution, not mutation, and they
|
||||||
**declare** via `arg_groups.overrides.declare_late_resolution(server_args,
|
**declare** via `arg_groups.overrides.declare_resolution(server_args, source,
|
||||||
source, **fields)`, which refuses the published instance. The declaration lands
|
**fields)`, the same call the rest of the pipeline makes; there is no
|
||||||
|
`declare_late_resolution` any more. *When* a declaration is made is not
|
||||||
|
something the code marks — the guardrails that used to read that marker
|
||||||
|
cover these sites through the ordinary keyword scan instead. The declaration lands
|
||||||
in the stash on that very object, so every holder of it carries the decision —
|
in the stash on that very object, so every holder of it carries the decision —
|
||||||
the HTTP server, the multi-tokenizer workers it is serialized for, the
|
the HTTP server, the multi-tokenizer workers it is serialized for, the
|
||||||
schedulers it forks — and each of them publishes bags projected from it. The
|
schedulers it forks — and each of them publishes bags projected from it. The
|
||||||
@@ -408,9 +411,38 @@ derivation cannot enumerate.
|
|||||||
|
|
||||||
One consequence worth knowing: because the fields are the raw input, resolving a
|
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
|
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
|
pipeline reads only its own input. **So a resolved record is not copied at
|
||||||
resolved record (it carries the declarations and the `model_config` memo, so the
|
all.** A caller that needs one field different for the process it is about to
|
||||||
copy does not re-resolve at all).
|
hand the record to — the Ray paths and their `dist_init_addr` — declares it on
|
||||||
|
the record it holds (`declare_resolution`) and hands that over: the declaration
|
||||||
|
travels inside the object, the receiving process projects its bags from it, and
|
||||||
|
nothing re-resolves. There is no `ServerArgs.replace_resolved` any more, and the
|
||||||
|
`model_config`-memo bug that copying used to cause (a copy marked resolved but
|
||||||
|
arriving without the memo cannot refill it, because the guard refuses the write)
|
||||||
|
is gone by construction rather than guarded.
|
||||||
|
|
||||||
|
A bag `override` cannot stand in for this. It is *not* because overriding needs
|
||||||
|
a publish — `set_server_args` is what projects the bags and `override` works as
|
||||||
|
soon as the context holds a record — but because `override` writes bag leaves
|
||||||
|
and by contract never touches the record, so its effect cannot travel inside an
|
||||||
|
object to another process.
|
||||||
|
|
||||||
|
### The declaration stash has one writer
|
||||||
|
|
||||||
|
Everything that decides configuration goes through
|
||||||
|
`declare_resolution(server_args, source, **fields)`. It validates the names,
|
||||||
|
refuses the published config (the stash is projected at publish and never
|
||||||
|
again, so a later declaration is a silent no-op), and appends. The other names
|
||||||
|
around it are spellings, not mechanisms:
|
||||||
|
|
||||||
|
| name | what it adds |
|
||||||
|
|---|---|
|
||||||
|
| `run_post_process_pass` | runs a pass at its slot and validates its return; declares through `declare_resolution`. A pass returning an **empty** dict is a validation, not a declaration, and stays legal on the published instance — `Engine(server_args=sa)` after `Engine.shutdown()` re-runs `check_server_args` on the very instance the context holds |
|
||||||
|
| `record_foreign_defaults` | for a resolver this tree does not own (an out-of-tree platform plugin, a registered speculative algorithm), whose interface is to *assign* fields. It gets a stand-in whose reads fall through to `resolving_view`; what it assigned is declared. The record is never written, so the write seal has no exception. In-tree code does not go through it — `handle_platform_defaults` wraps the platform hook, and the in-tree speculative dispatcher is called directly, because handed the stand-in its own `declare_resolution` calls would stash on that instead |
|
||||||
|
|
||||||
|
`resolution_projection` is gone; the whole-object readback is
|
||||||
|
`ServerArgs.resolved_dict()`, which is what `/server_info` and its gRPC and
|
||||||
|
in-process twins report.
|
||||||
|
|
||||||
### Adding a model-specific config adjustment
|
### Adding a model-specific config adjustment
|
||||||
|
|
||||||
@@ -531,8 +563,10 @@ ONE thread — do not design for TBO threads that don't exist.
|
|||||||
|
|
||||||
## Guardrails (these fail CI; what to do when they fire)
|
## Guardrails (these fail CI; what to do when they fire)
|
||||||
|
|
||||||
1. **Strict mutation guard** (always on): bare `server_args.x = ...` after resolution
|
1. **Strict mutation guard** (always on, and with no exception): bare
|
||||||
raises unconditionally in `ServerArgs.__setattr__` — this *is* the guarantee that
|
`server_args.x = ...` after resolution raises unconditionally in
|
||||||
|
`ServerArgs.__setattr__` — the named lift that out-of-tree plugins used to
|
||||||
|
ask for is gone, they assign onto a stand-in instead — this *is* the guarantee that
|
||||||
no writer can desync the bags, so there is no writer ratchet any more. Change
|
no writer can desync the bags, so there is no writer ratchet any more. Change
|
||||||
resolved config with `get_context().override`; hand a per-runner value to its
|
resolved config with `get_context().override`; hand a per-runner value to its
|
||||||
runner as a constructor argument. Projected bags are sealed the same way (leaf
|
runner as a constructor argument. Projected bags are sealed the same way (leaf
|
||||||
@@ -621,7 +655,7 @@ Never module-skip a test "until the migration settles" — seed the context inst
|
|||||||
Key source files: `python/sglang/srt/runtime_context.py` (the container, every tier,
|
Key source files: `python/sglang/srt/runtime_context.py` (the container, every tier,
|
||||||
`publish`, `_ConfigBag`, `override_server_args`),
|
`publish`, `_ConfigBag`, `override_server_args`),
|
||||||
`python/sglang/srt/arg_groups/overrides.py` (override registry, passes,
|
`python/sglang/srt/arg_groups/overrides.py` (override registry, passes,
|
||||||
`declare_late_resolution`), `python/sglang/srt/server_args.py` (`NS` metadata,
|
`declare_resolution` and the spellings around it), `python/sglang/srt/server_args.py` (`NS` metadata,
|
||||||
`Arg(..., resolvable=True)`, `__setattr__` strict guard), and the guardrail tests under
|
`Arg(..., resolvable=True)`, `__setattr__` strict guard), and the guardrail tests under
|
||||||
`test/registered/unit/` (`test_server_args_mutation_ratchet.py`,
|
`test/registered/unit/` (`test_server_args_mutation_ratchet.py`,
|
||||||
`test_global_config_read_ratchet.py`,
|
`test_global_config_read_ratchet.py`,
|
||||||
|
|||||||
@@ -64,7 +64,11 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
import torch.distributed as dist
|
import torch.distributed as dist
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import resolution_result, resolving_view
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
declare_resolution,
|
||||||
|
resolution_result,
|
||||||
|
resolving_view,
|
||||||
|
)
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.distributed.parallel_state import (
|
from sglang.srt.distributed.parallel_state import (
|
||||||
destroy_distributed_environment,
|
destroy_distributed_environment,
|
||||||
@@ -1034,8 +1038,8 @@ def main(server_args, bench_args):
|
|||||||
decode = dict(graph_config.get(Phase.DECODE) or {})
|
decode = dict(graph_config.get(Phase.DECODE) or {})
|
||||||
decode["max_bs"] = max(bench_args.batch_size)
|
decode["max_bs"] = max(bench_args.batch_size)
|
||||||
graph_config[Phase.DECODE] = decode
|
graph_config[Phase.DECODE] = decode
|
||||||
server_args = server_args.replace_resolved(
|
declare_resolution(
|
||||||
"benchmark.one_batch", cuda_graph_config=graph_config
|
server_args, "benchmark.one_batch", cuda_graph_config=graph_config
|
||||||
)
|
)
|
||||||
server_args.resolve_once()
|
server_args.resolve_once()
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import logging
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
declare_late_resolution,
|
declare_resolution,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -24,9 +24,7 @@ def check_lora_server_args(server_args: Any):
|
|||||||
# Enable LoRA if any LoRA paths are provided for backward compatibility.
|
# Enable LoRA if any LoRA paths are provided for backward compatibility.
|
||||||
if cfg.lora_paths:
|
if cfg.lora_paths:
|
||||||
if cfg.enable_lora is None:
|
if cfg.enable_lora is None:
|
||||||
declare_late_resolution(
|
declare_resolution(server_args, "check_lora_server_args", enable_lora=True)
|
||||||
server_args, "check_lora_server_args", enable_lora=True
|
|
||||||
)
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"--enable-lora is set to True because --lora-paths is provided."
|
"--enable-lora is set to True because --lora-paths is provided."
|
||||||
)
|
)
|
||||||
@@ -37,7 +35,7 @@ def check_lora_server_args(server_args: Any):
|
|||||||
|
|
||||||
if cfg.enable_lora:
|
if cfg.enable_lora:
|
||||||
if cfg.enable_lora_overlap_loading is None:
|
if cfg.enable_lora_overlap_loading is None:
|
||||||
declare_late_resolution(
|
declare_resolution(
|
||||||
server_args, "check_lora_server_args", enable_lora_overlap_loading=False
|
server_args, "check_lora_server_args", enable_lora_overlap_loading=False
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -93,11 +91,11 @@ def check_lora_server_args(server_args: Any):
|
|||||||
"Expected a string or a dictionary."
|
"Expected a string or a dictionary."
|
||||||
)
|
)
|
||||||
parsed_lora_paths.append(lora_ref)
|
parsed_lora_paths.append(lora_ref)
|
||||||
declare_late_resolution(
|
declare_resolution(
|
||||||
server_args, "check_lora_server_args", lora_paths=parsed_lora_paths
|
server_args, "check_lora_server_args", lora_paths=parsed_lora_paths
|
||||||
)
|
)
|
||||||
elif isinstance(cfg.lora_paths, dict):
|
elif isinstance(cfg.lora_paths, dict):
|
||||||
declare_late_resolution(
|
declare_resolution(
|
||||||
server_args,
|
server_args,
|
||||||
"check_lora_server_args",
|
"check_lora_server_args",
|
||||||
lora_paths=[
|
lora_paths=[
|
||||||
@@ -111,9 +109,7 @@ def check_lora_server_args(server_args: Any):
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
elif cfg.lora_paths is None:
|
elif cfg.lora_paths is None:
|
||||||
declare_late_resolution(
|
declare_resolution(server_args, "check_lora_server_args", lora_paths=[])
|
||||||
server_args, "check_lora_server_args", lora_paths=[]
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Invalid type for --lora-paths: {type(cfg.lora_paths)}. "
|
f"Invalid type for --lora-paths: {type(cfg.lora_paths)}. "
|
||||||
@@ -123,7 +119,7 @@ def check_lora_server_args(server_args: Any):
|
|||||||
# Normalize target modules to a set; keep {"all"} as a sentinel
|
# Normalize target modules to a set; keep {"all"} as a sentinel
|
||||||
# that gets resolved model-awarely in lora_manager.init_lora_shapes().
|
# that gets resolved model-awarely in lora_manager.init_lora_shapes().
|
||||||
if cfg.lora_target_modules:
|
if cfg.lora_target_modules:
|
||||||
declare_late_resolution(
|
declare_resolution(
|
||||||
server_args,
|
server_args,
|
||||||
"check_lora_server_args",
|
"check_lora_server_args",
|
||||||
lora_target_modules=set(cfg.lora_target_modules),
|
lora_target_modules=set(cfg.lora_target_modules),
|
||||||
|
|||||||
@@ -17,10 +17,9 @@ Model-identity adjustments to the server configuration are DECLARED here and
|
|||||||
appended to the record's declaration stash (gate order, last writer wins).
|
appended to the record's declaration stash (gate order, last writer wins).
|
||||||
Nothing here writes back onto ``ServerArgs``: the record holds the user's raw
|
Nothing here writes back onto ``ServerArgs``: the record holds the user's raw
|
||||||
input, and a decision is read through ``resolution_result`` or the published
|
input, and a decision is read through ``resolution_result`` or the published
|
||||||
config bags — model code never mutates ``ServerArgs`` fields imperatively. The
|
config bags — model code never mutates ``ServerArgs`` fields imperatively. That
|
||||||
one channel that still leaves a field changed is ``declare_direct_writes``,
|
holds without exception: a resolver this tree does not own assigns onto a
|
||||||
which does not perform the write: it captures one an out-of-tree plugin already
|
stand-in (``record_foreign_defaults``), and what it set is declared.
|
||||||
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]``:
|
||||||
|
|
||||||
@@ -33,7 +32,6 @@ Two declaration forms, keyed on ``hf_config.architectures[0]``:
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import copy
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
@@ -116,9 +114,9 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
|
|||||||
an empty dict is a validation, and it may run on the published instance --
|
an empty dict is a validation, and it may run on the published instance --
|
||||||
it has to, because ``Engine(server_args=sa)`` after ``Engine.shutdown()``
|
it has to, because ``Engine(server_args=sa)`` after ``Engine.shutdown()``
|
||||||
re-runs ``check_server_args`` on the very instance the context still holds.
|
re-runs ``check_server_args`` on the very instance the context still holds.
|
||||||
A pass that returns a non-empty dict there is refused, as
|
A pass that returns a non-empty dict there is refused by the guard in
|
||||||
``declare_late_resolution`` is -- post-publish changes go to the bags through
|
``declare_resolution``, as a late declaration is -- post-publish changes go
|
||||||
``get_context().override(...)``.
|
to the bags through ``get_context().override(...)``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args)))
|
declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args)))
|
||||||
@@ -133,29 +131,12 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
|
|||||||
# a rebuild: `Engine(server_args=sa)` after `Engine.shutdown()` hands
|
# a rebuild: `Engine(server_args=sa)` after `Engine.shutdown()` hands
|
||||||
# back the same instance while the context still holds it, and
|
# back the same instance while the context still holds it, and
|
||||||
# refusing on identity alone would fail that launch.
|
# refusing on identity alone would fail that launch.
|
||||||
try:
|
# Only a non-empty return is a declaration. An empty one is a
|
||||||
published = get_context().server_args
|
# validation and may run on the published instance -- see above -- so it
|
||||||
except ValueError:
|
# must not reach the guard in `declare_resolution`.
|
||||||
published = None
|
if declared:
|
||||||
if published is server_args:
|
declare_resolution(server_args, fn.__qualname__, **declared)
|
||||||
raise ValueError(
|
validate_declarations(server_args, [(fn.__qualname__, dict(declared))])
|
||||||
f"run_post_process_pass({fn.__qualname__!r}) declared "
|
|
||||||
f"{sorted(declared)} on the published config; the stash is "
|
|
||||||
"projected at publish and never again, so this would be a "
|
|
||||||
"silent no-op -- post-publish changes go to the bags via "
|
|
||||||
"get_context().override(...)"
|
|
||||||
)
|
|
||||||
entry = (fn.__qualname__, dict(declared))
|
|
||||||
stash = getattr(server_args, "_resolved_overrides", None)
|
|
||||||
if stash is None:
|
|
||||||
# Handlers hosting pass slots may be invoked directly on fixtures
|
|
||||||
# that never ran the monolith dispatch (which owns the stash);
|
|
||||||
# create it lazily. Real publishes always pass through the
|
|
||||||
# dispatch first — the dispatch ASSIGNS the stash, so pass slots
|
|
||||||
# must sit at or after it in __post_init__ order.
|
|
||||||
stash = server_args._resolved_overrides = []
|
|
||||||
stash.append(entry)
|
|
||||||
validate_declarations(server_args, [entry])
|
|
||||||
|
|
||||||
|
|
||||||
def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
|
def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
|
||||||
@@ -167,52 +148,32 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
|
|||||||
(or `resolved_view(server_args)`), which
|
(or `resolved_view(server_args)`), which
|
||||||
`test_resolution_reads_the_declarations` pins.
|
`test_resolution_reads_the_declarations` pins.
|
||||||
|
|
||||||
For resolvers inside ``__post_init__``; launcher-stage resolution goes
|
Every declaration goes through here, whenever it is made: inside
|
||||||
through ``declare_late_resolution``. A name that is not a field is rejected
|
``__post_init__``, at launcher stage (LoRA normalization, the auto-detected
|
||||||
here rather than becoming an attribute nothing reads.
|
parsers -- they decide what the process will run with, so they belong to the
|
||||||
|
pipeline even though they run after it), and on a copy about to cross a
|
||||||
|
process boundary. A name that is not a field is rejected here rather than
|
||||||
|
becoming an attribute nothing reads.
|
||||||
|
|
||||||
|
Refuses the published config. The stash is projected at publish and never
|
||||||
|
again, so a declaration afterwards is a silent no-op; post-publish changes
|
||||||
|
go to the bags through ``get_context().override(...)``.
|
||||||
"""
|
"""
|
||||||
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)))
|
||||||
if unknown:
|
if unknown:
|
||||||
raise AttributeError(f"{source}: {unknown} are not ServerArgs fields")
|
raise AttributeError(f"{source}: {unknown} are not ServerArgs fields")
|
||||||
stash = getattr(server_args, "_resolved_overrides", None)
|
|
||||||
if stash is None:
|
|
||||||
stash = []
|
|
||||||
server_args._resolved_overrides = stash
|
|
||||||
stash.append((source, dict(fields)))
|
|
||||||
|
|
||||||
|
|
||||||
def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> None:
|
|
||||||
"""Resolve fields on a config that is **not published yet**.
|
|
||||||
|
|
||||||
A few resolution rules cannot run inside ``__post_init__``: LoRA
|
|
||||||
normalization and the auto-parser detection need the launcher's validation
|
|
||||||
stage (and, for the parsers, a tokenizer / chat-template load). They still
|
|
||||||
belong to the resolution pipeline — they decide what the process will run
|
|
||||||
with — so their decision goes to the stash like any other, and the record
|
|
||||||
keeps what the caller passed. Every holder of that instance reads the
|
|
||||||
decision the same way the rest of the pipeline does: the bags it publishes,
|
|
||||||
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
|
|
||||||
field write would desync them, which is what ``get_context().override`` is
|
|
||||||
for.
|
|
||||||
"""
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
published = get_context().server_args
|
published = get_context().server_args
|
||||||
except ValueError:
|
except ValueError:
|
||||||
published = None
|
published = None
|
||||||
if published is server_args:
|
if published is server_args:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"declare_late_resolution({source!r}) called on the published config; "
|
f"{source}: declared on the published config; the stash is "
|
||||||
"post-publish changes go to the bags via get_context().override(...)"
|
"projected at publish and never again, so this would be a silent "
|
||||||
|
"no-op -- post-publish changes go to the bags via "
|
||||||
|
"get_context().override(...)"
|
||||||
)
|
)
|
||||||
log = getattr(server_args, "_runtime_mutations", None)
|
|
||||||
if log is None:
|
|
||||||
log = []
|
|
||||||
server_args._runtime_mutations = log
|
|
||||||
log.append((source, dict(fields)))
|
|
||||||
stash = getattr(server_args, "_resolved_overrides", None)
|
stash = getattr(server_args, "_resolved_overrides", None)
|
||||||
if stash is None:
|
if stash is None:
|
||||||
stash = []
|
stash = []
|
||||||
@@ -220,59 +181,65 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non
|
|||||||
stash.append((source, dict(fields)))
|
stash.append((source, dict(fields)))
|
||||||
|
|
||||||
|
|
||||||
def declare_direct_writes(
|
class _ForeignDefaults:
|
||||||
|
"""The stand-in handed to a resolver this tree does not own.
|
||||||
|
|
||||||
|
Reads fall through to the resolving view, so a plugin sees what resolution
|
||||||
|
has decided so far rather than the raw input -- better than what it used to
|
||||||
|
get, which was the record's own fields. Writes are captured here and
|
||||||
|
declared by the caller, so the record is never written and the write seal
|
||||||
|
has no exception.
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = ("_cfg", "_written")
|
||||||
|
|
||||||
|
def __init__(self, server_args: Any):
|
||||||
|
object.__setattr__(self, "_cfg", resolving_view(server_args))
|
||||||
|
object.__setattr__(self, "_written", {})
|
||||||
|
|
||||||
|
def __getattr__(self, name: str) -> Any:
|
||||||
|
written = object.__getattribute__(self, "_written")
|
||||||
|
if name in written:
|
||||||
|
return written[name]
|
||||||
|
return getattr(object.__getattribute__(self, "_cfg"), name)
|
||||||
|
|
||||||
|
def __setattr__(self, name: str, value: Any) -> None:
|
||||||
|
object.__getattribute__(self, "_written")[name] = value
|
||||||
|
|
||||||
|
|
||||||
|
def record_foreign_defaults(
|
||||||
server_args: Any, source: str, resolve: Callable[[Any], Any]
|
server_args: Any, source: str, resolve: Callable[[Any], Any]
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""Run a resolver that writes the fields directly, and declare what it moved.
|
"""Run a resolver this tree does not own, and declare what it set.
|
||||||
|
|
||||||
|
Out-of-tree platform plugins and registered speculative algorithms are
|
||||||
|
handed a configuration and assign fields on it. That interface is not ours
|
||||||
|
to change, so the assignment stays the contract -- it just lands on a
|
||||||
|
stand-in instead of the record, and what it set is declared like any other
|
||||||
|
decision. Nothing writes the record, which is why there is no longer a
|
||||||
|
named hole in the seal.
|
||||||
|
|
||||||
Returns whatever the resolver returned, so a provider with a return value
|
Returns whatever the resolver returned, so a provider with a return value
|
||||||
can go through the same capture.
|
goes through the same capture.
|
||||||
|
|
||||||
Out-of-tree platform plugins are handed the record and set fields on it.
|
Non-field names are dropped: a plugin scribbling on an attribute that is
|
||||||
Their implementations live outside this tree, so they cannot be converted
|
not configuration is not a decision, and it was invisible to the previous
|
||||||
by editing the resolver; and the raw snapshot is taken before the pipeline
|
diff for the same reason.
|
||||||
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
|
|
||||||
plugin that mutates a value in place reaches the projection anyway, because
|
|
||||||
the raw snapshot and the stash entries hold the same object it mutated.
|
|
||||||
|
|
||||||
A stand-in record (tests drive the hooks with a plain namespace) has no
|
A stand-in record (tests drive the hooks with a plain namespace) has no
|
||||||
fields to diff and no projection to feed, so the resolver runs uncaptured.
|
view to read, so the resolver runs against it directly and uncaptured.
|
||||||
"""
|
"""
|
||||||
if not dataclasses.is_dataclass(server_args):
|
if not dataclasses.is_dataclass(server_args):
|
||||||
return resolve(server_args)
|
return resolve(server_args)
|
||||||
before = {
|
recorder = _ForeignDefaults(server_args)
|
||||||
field.name: getattr(server_args, field.name)
|
result = resolve(recorder)
|
||||||
for field in dataclasses.fields(server_args)
|
written = {
|
||||||
|
name: value
|
||||||
|
for name, value in object.__getattribute__(recorder, "_written").items()
|
||||||
|
if name in field_names(type(server_args))
|
||||||
}
|
}
|
||||||
already = len(getattr(server_args, "_resolved_overrides", None) or ())
|
if written:
|
||||||
# The one place the input seal comes off. The plugin writes the record;
|
declare_resolution(server_args, source, **written)
|
||||||
# the diff below captures what it moved into the stash so the projection
|
|
||||||
# and the bags carry it.
|
|
||||||
from sglang.srt.server_args import record_writable
|
|
||||||
|
|
||||||
with record_writable(server_args):
|
|
||||||
result = resolve(server_args)
|
|
||||||
stash = getattr(server_args, "_resolved_overrides", None)
|
|
||||||
if stash is None:
|
|
||||||
stash = []
|
|
||||||
server_args._resolved_overrides = stash
|
|
||||||
# A resolver reached this way can also declare properly -- the in-tree
|
|
||||||
# implementations of these hooks do. Those fields are already explained, and
|
|
||||||
# recording them again would attribute them to the wrapper and bury an
|
|
||||||
# actual direct write among the echoes.
|
|
||||||
declared = {name for _source, fields in stash[already:] for name in fields}
|
|
||||||
changed = {
|
|
||||||
name: getattr(server_args, name)
|
|
||||||
for name, previous in before.items()
|
|
||||||
if name not in declared and getattr(server_args, name) is not previous
|
|
||||||
}
|
|
||||||
if changed:
|
|
||||||
stash.append((source, changed))
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
@@ -303,40 +270,6 @@ def resolution_result(server_args: Any, field: str, default: Any = None) -> Any:
|
|||||||
return with_fallback(type(server_args), field, getattr(server_args, field, default))
|
return with_fallback(type(server_args), field, getattr(server_args, field, default))
|
||||||
|
|
||||||
|
|
||||||
def resolution_projection(server_args: Any) -> Dict[str, Any]:
|
|
||||||
"""Every field's resolved value, nested dataclasses expanded.
|
|
||||||
|
|
||||||
The whole-object shape of ``resolution_result``, for the exits that hand out
|
|
||||||
the entire configuration (``/server_info``, the gRPC and engine readbacks).
|
|
||||||
They used ``dataclasses.asdict``, which reads the fields -- the operator's
|
|
||||||
input, not what resolution decided. Field values only: the private resolution
|
|
||||||
bookkeeping and the ``model_config`` memo that a ``vars()`` dump carried into
|
|
||||||
the readback are not configuration.
|
|
||||||
"""
|
|
||||||
return {
|
|
||||||
field.name: _plain(resolution_result(server_args, field.name))
|
|
||||||
for field in dataclasses.fields(server_args)
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def _plain(value: Any) -> Any:
|
|
||||||
"""``dataclasses.asdict``'s conversion, applied to one value: dataclasses
|
|
||||||
become dicts, containers recurse, everything else is deep-copied (a caller
|
|
||||||
mutating the dump must not reach the live configuration)."""
|
|
||||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
|
||||||
return {
|
|
||||||
field.name: _plain(getattr(value, field.name))
|
|
||||||
for field in dataclasses.fields(value)
|
|
||||||
}
|
|
||||||
if isinstance(value, tuple) and hasattr(value, "_fields"): # namedtuple
|
|
||||||
return type(value)(*(_plain(item) for item in value))
|
|
||||||
if isinstance(value, (list, tuple)):
|
|
||||||
return type(value)(_plain(item) for item in value)
|
|
||||||
if isinstance(value, dict):
|
|
||||||
return type(value)((_plain(k), _plain(v)) for k, v in value.items())
|
|
||||||
return copy.deepcopy(value)
|
|
||||||
|
|
||||||
|
|
||||||
def pre_capture_activation_reserve_mb_of(cfg: Any, gpu_mem: Optional[float]) -> float:
|
def pre_capture_activation_reserve_mb_of(cfg: Any, gpu_mem: Optional[float]) -> float:
|
||||||
"""The activation working-set reserve held back before cuda-graph capture.
|
"""The activation working-set reserve held back before cuda-graph capture.
|
||||||
|
|
||||||
|
|||||||
@@ -15,11 +15,9 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
_page_size_default,
|
_page_size_default,
|
||||||
_pipeline_parallel_overlap_disable,
|
_pipeline_parallel_overlap_disable,
|
||||||
_sampling_backend_default,
|
_sampling_backend_default,
|
||||||
declare_direct_writes,
|
|
||||||
resolving_view,
|
resolving_view,
|
||||||
run_post_process_pass,
|
run_post_process_pass,
|
||||||
)
|
)
|
||||||
from sglang.srt.platforms import current_platform
|
|
||||||
from sglang.srt.utils.common import get_device_memory_capacity
|
from sglang.srt.utils.common import get_device_memory_capacity
|
||||||
|
|
||||||
|
|
||||||
@@ -204,6 +202,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
|
|||||||
handle_mps_backends,
|
handle_mps_backends,
|
||||||
handle_nccl_pre_warm,
|
handle_nccl_pre_warm,
|
||||||
handle_npu_backends,
|
handle_npu_backends,
|
||||||
|
handle_platform_defaults,
|
||||||
handle_symm_mem_device_support,
|
handle_symm_mem_device_support,
|
||||||
handle_xpu_backends,
|
handle_xpu_backends,
|
||||||
)
|
)
|
||||||
@@ -217,13 +216,7 @@ def run_resolution_pipeline(server_args: Any) -> None:
|
|||||||
# keys off enable_symm_mem.
|
# keys off enable_symm_mem.
|
||||||
handle_symm_mem_device_support(server_args)
|
handle_symm_mem_device_support(server_args)
|
||||||
|
|
||||||
# OOT platform plugins set fields directly (an interface this tree
|
handle_platform_defaults(server_args)
|
||||||
# does not own); the diff records what they applied.
|
|
||||||
declare_direct_writes(
|
|
||||||
server_args,
|
|
||||||
f"platform:{current_platform.device_name}",
|
|
||||||
current_platform.apply_server_args_defaults,
|
|
||||||
)
|
|
||||||
|
|
||||||
gpu_mem = get_device_memory_capacity(cfg.device)
|
gpu_mem = get_device_memory_capacity(cfg.device)
|
||||||
|
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from typing import Any
|
|||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
declare_resolution,
|
declare_resolution,
|
||||||
|
record_foreign_defaults,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
|
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
|
||||||
@@ -84,6 +85,29 @@ def handle_nccl_pre_warm(server_args: Any):
|
|||||||
declare_resolution(server_args, "_handle_nccl_pre_warm", pre_warm_nccl=False)
|
declare_resolution(server_args, "_handle_nccl_pre_warm", pre_warm_nccl=False)
|
||||||
|
|
||||||
|
|
||||||
|
def handle_platform_defaults(server_args: Any):
|
||||||
|
"""An out-of-tree platform's defaults, declared like every rule beside it.
|
||||||
|
|
||||||
|
`Platform.apply_server_args_defaults` is a plugin interface: the platform is
|
||||||
|
handed a configuration and assigns the fields it wants defaulted. In-tree
|
||||||
|
platforms do not implement it -- the base is a no-op and nothing overrides
|
||||||
|
it -- so this captures nothing here and exists for the platforms that live
|
||||||
|
outside this tree.
|
||||||
|
|
||||||
|
Ordering: it must precede `handle_gpu_memory_settings`, whose symm-mem
|
||||||
|
prealloc default keys off `enable_symm_mem`.
|
||||||
|
"""
|
||||||
|
# `current_platform` is the plugin object; `get_platform()` is the facts
|
||||||
|
# view over it and carries neither the name nor the hook.
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
|
|
||||||
|
record_foreign_defaults(
|
||||||
|
server_args,
|
||||||
|
f"platform:{current_platform.device_name}",
|
||||||
|
current_platform.apply_server_args_defaults,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def handle_symm_mem_device_support(server_args: Any):
|
def handle_symm_mem_device_support(server_args: Any):
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
# The symm-mem allocator compiles a CUDA plugin and links -lnccl, so off
|
# The symm-mem allocator compiles a CUDA plugin and links -lnccl, so off
|
||||||
|
|||||||
@@ -8,9 +8,9 @@ from typing import TYPE_CHECKING, Optional
|
|||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
_speculative_moe_runner_default,
|
_speculative_moe_runner_default,
|
||||||
attention_backends_of,
|
attention_backends_of,
|
||||||
declare_direct_writes,
|
|
||||||
declare_resolution,
|
declare_resolution,
|
||||||
model_config_of,
|
model_config_of,
|
||||||
|
record_foreign_defaults,
|
||||||
resolved_view,
|
resolved_view,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
run_post_process_pass,
|
run_post_process_pass,
|
||||||
@@ -157,7 +157,7 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
|||||||
|
|
||||||
# TODO: move the per-algorithm validation below into spec module hooks.
|
# TODO: move the per-algorithm validation below into spec module hooks.
|
||||||
if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None:
|
if isinstance(algo, CustomSpecAlgo) and algo.validate_server_args is not None:
|
||||||
declare_direct_writes(
|
record_foreign_defaults(
|
||||||
server_args,
|
server_args,
|
||||||
"handle_speculative_decoding.custom_validate",
|
"handle_speculative_decoding.custom_validate",
|
||||||
algo.validate_server_args,
|
algo.validate_server_args,
|
||||||
@@ -175,13 +175,24 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
|||||||
_init_adaptive_speculative_params(server_args)
|
_init_adaptive_speculative_params(server_args)
|
||||||
|
|
||||||
if algo is not None:
|
if algo is not None:
|
||||||
# A registered algorithm's callback lives outside this tree and sets
|
# Imported here and not above: the name is only bound inside the
|
||||||
# fields on the record, so the writes are captured around the call.
|
# `speculative_algorithm is not None` branch, and this runs either way.
|
||||||
declare_direct_writes(
|
from sglang.srt.speculative.spec_registry import CustomSpecAlgo
|
||||||
|
|
||||||
|
if isinstance(algo, CustomSpecAlgo):
|
||||||
|
# A registered algorithm's callback lives outside this tree and
|
||||||
|
# assigns fields, so it gets the stand-in and its writes are
|
||||||
|
# declared.
|
||||||
|
record_foreign_defaults(
|
||||||
server_args,
|
server_args,
|
||||||
"handle_speculative_decoding.custom_algo",
|
"handle_speculative_decoding.custom_algo",
|
||||||
algo.handle_server_args,
|
algo.handle_server_args,
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
# The in-tree dispatcher, which declares. It needs the record
|
||||||
|
# itself: handed the stand-in, its `declare_resolution` calls would
|
||||||
|
# stash on that instead.
|
||||||
|
algo.handle_server_args(server_args)
|
||||||
|
|
||||||
|
|
||||||
def _handle_dflash(server_args: ServerArgs) -> None:
|
def _handle_dflash(server_args: ServerArgs) -> None:
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ import jinja2.ext
|
|||||||
import jinja2.nodes
|
import jinja2.nodes
|
||||||
import jinja2.sandbox
|
import jinja2.sandbox
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import declare_late_resolution, resolving_view
|
from sglang.srt.arg_groups.overrides import declare_resolution, resolving_view
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -829,4 +829,4 @@ def resolve_auto_parsers(server_args) -> None:
|
|||||||
detected[attr] = _detect_auto_parser(attr, ctx, rules, label)
|
detected[attr] = _detect_auto_parser(attr, ctx, rules, label)
|
||||||
|
|
||||||
if detected:
|
if detected:
|
||||||
declare_late_resolution(server_args, "template-detection", **detected)
|
declare_resolution(server_args, "template-detection", **detected)
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ import ray
|
|||||||
from ray.util.placement_group import PlacementGroup
|
from ray.util.placement_group import PlacementGroup
|
||||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||||
|
|
||||||
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
from sglang.srt.entrypoints.engine import (
|
from sglang.srt.entrypoints.engine import (
|
||||||
Engine,
|
Engine,
|
||||||
SchedulerInitResult,
|
SchedulerInitResult,
|
||||||
@@ -471,12 +472,16 @@ class RayEngine(Engine):
|
|||||||
f"enable_dp_attention={parallel.enable_dp_attention}"
|
f"enable_dp_attention={parallel.enable_dp_attention}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set dist_init_addr on server_args so PortArgs.init_new() can compute
|
# Declared on the record itself so `PortArgs.init_new()` can compute
|
||||||
# TCP addresses correctly (required for DP attention path).
|
# TCP addresses (required for the DP attention path). No copy: this
|
||||||
dp_server_args = server_args.replace_resolved(
|
# process does not publish here, and every other reader of the field
|
||||||
|
# goes through the bags its own process projects.
|
||||||
|
declare_resolution(
|
||||||
|
server_args,
|
||||||
"ray.dp_controller",
|
"ray.dp_controller",
|
||||||
dist_init_addr=f"{rank0_node_ip}:{port_args.nccl_port}",
|
dist_init_addr=f"{rank0_node_ip}:{port_args.nccl_port}",
|
||||||
)
|
)
|
||||||
|
dp_server_args = server_args
|
||||||
# Create the DP controller in-process. This blocks until all actors
|
# Create the DP controller in-process. This blocks until all actors
|
||||||
# are initialized and their event loops have started.
|
# are initialized and their event loops have started.
|
||||||
controller = RayDataParallelController(
|
controller = RayDataParallelController(
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from typing import Any, Dict, Optional
|
|||||||
|
|
||||||
import ray
|
import ray
|
||||||
|
|
||||||
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
from sglang.srt.runtime_context import publish
|
from sglang.srt.runtime_context import publish
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
|
|
||||||
@@ -54,11 +55,13 @@ class SchedulerActor:
|
|||||||
numa_bind_to_node,
|
numa_bind_to_node,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Override dist_init_addr if provided (for multi-node), through
|
# Declared, not copied: Ray deserializes the argument per call, so this
|
||||||
# `replace_resolved` so the copy keeps the parent's resolution.
|
# record is the actor's own and nothing else in the process holds it.
|
||||||
|
# The field stays the operator's input; `PortArgs.init_new` and the bags
|
||||||
|
# this actor publishes read the decision.
|
||||||
if dist_init_addr:
|
if dist_init_addr:
|
||||||
server_args = server_args.replace_resolved(
|
declare_resolution(
|
||||||
"ray.scheduler_actor", dist_init_addr=dist_init_addr
|
server_args, "ray.scheduler_actor", dist_init_addr=dist_init_addr
|
||||||
)
|
)
|
||||||
|
|
||||||
# Get actual GPU IDs from Ray runtime context
|
# Get actual GPU IDs from Ray runtime context
|
||||||
|
|||||||
@@ -1282,7 +1282,7 @@ class _ServerArgsOverride:
|
|||||||
self._prev_parallel_config = ctx.parallel._config
|
self._prev_parallel_config = ctx.parallel._config
|
||||||
self._prev_capture = ctx.flags.capture.enable_torch_compile
|
self._prev_capture = ctx.flags.capture.enable_torch_compile
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
declare_late_resolution,
|
declare_resolution,
|
||||||
)
|
)
|
||||||
|
|
||||||
server_args = ServerArgs(model_path="dummy")
|
server_args = ServerArgs(model_path="dummy")
|
||||||
@@ -1306,7 +1306,7 @@ class _ServerArgsOverride:
|
|||||||
fields = set(type(server_args).__dataclass_fields__)
|
fields = set(type(server_args).__dataclass_fields__)
|
||||||
declared = {n: v for n, v in self._fields.items() if n in fields}
|
declared = {n: v for n, v in self._fields.items() if n in fields}
|
||||||
if declared:
|
if declared:
|
||||||
declare_late_resolution(server_args, "override_server_args", **declared)
|
declare_resolution(server_args, "override_server_args", **declared)
|
||||||
# What is left seeds the record's own private caches (`_model_config`
|
# What is left seeds the record's own private caches (`_model_config`
|
||||||
# and friends), which are not configuration and never were.
|
# and friends), which are not configuration and never were.
|
||||||
seeds = {n: v for n, v in self._fields.items() if n not in fields}
|
seeds = {n: v for n, v in self._fields.items() if n not in fields}
|
||||||
|
|||||||
@@ -40,7 +40,6 @@ import functools
|
|||||||
import logging
|
import logging
|
||||||
import tempfile
|
import tempfile
|
||||||
import uuid
|
import uuid
|
||||||
from contextlib import contextmanager
|
|
||||||
from typing import Any, NoReturn
|
from typing import Any, NoReturn
|
||||||
|
|
||||||
from sglang.kernels.ops.kv_canary.consts import RealKvHashMode
|
from sglang.kernels.ops.kv_canary.consts import RealKvHashMode
|
||||||
@@ -53,7 +52,7 @@ from sglang.srt.arg_groups.argparse_actions import (
|
|||||||
from sglang.srt.arg_groups.model_override_base import ep_joiner_of, ep_scale_joiner_of
|
from sglang.srt.arg_groups.model_override_base import ep_joiner_of, ep_scale_joiner_of
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
remote_instance_transfer_engine_of,
|
remote_instance_transfer_engine_of,
|
||||||
resolution_projection,
|
resolution_result,
|
||||||
resolving_view,
|
resolving_view,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -170,6 +169,24 @@ from sglang.srt.utils.common import ( # noqa: F401
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _plain(value: Any) -> Any:
|
||||||
|
"""``dataclasses.asdict``'s conversion, applied to one value: dataclasses
|
||||||
|
become dicts, containers recurse, everything else is deep-copied (a caller
|
||||||
|
mutating the dump must not reach the live configuration)."""
|
||||||
|
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||||
|
return {
|
||||||
|
field.name: _plain(getattr(value, field.name))
|
||||||
|
for field in dataclasses.fields(value)
|
||||||
|
}
|
||||||
|
if isinstance(value, tuple) and hasattr(value, "_fields"): # namedtuple
|
||||||
|
return type(value)(*(_plain(item) for item in value))
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return type(value)(_plain(item) for item in value)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return type(value)((_plain(k), _plain(v)) for k, v in value.items())
|
||||||
|
return copy.deepcopy(value)
|
||||||
|
|
||||||
|
|
||||||
class ServerArgs:
|
class ServerArgs:
|
||||||
"""Server-wide configuration for SGLang.
|
"""Server-wide configuration for SGLang.
|
||||||
|
|
||||||
@@ -252,9 +269,8 @@ class ServerArgs:
|
|||||||
from sglang.srt.arg_groups.pipeline import run_resolution_pipeline
|
from sglang.srt.arg_groups.pipeline import run_resolution_pipeline
|
||||||
|
|
||||||
# Sealed for the duration, not just afterwards: everything below this
|
# Sealed for the duration, not just afterwards: everything below this
|
||||||
# line reads the input and declares against it, and the one channel
|
# line reads the input and declares against it. No exceptions -- even a
|
||||||
# that still writes the record (`declare_direct_writes`, for
|
# resolver from outside this tree assigns onto a stand-in, not here.
|
||||||
# out-of-tree platform plugins) asks for the seal to be lifted by name.
|
|
||||||
self._input_frozen = True
|
self._input_frozen = True
|
||||||
try:
|
try:
|
||||||
run_resolution_pipeline(self)
|
run_resolution_pipeline(self)
|
||||||
@@ -298,58 +314,10 @@ class ServerArgs:
|
|||||||
`model_config` memo are not fields and do not appear.
|
`model_config` memo are not fields and do not appear.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
return resolution_projection(self)
|
return {
|
||||||
|
field.name: _plain(resolution_result(self, field.name))
|
||||||
def replace_resolved(self, source: str, **changes: Any) -> ServerArgs:
|
for field in dataclasses.fields(self)
|
||||||
"""A copy of this record that stays resolved, and says what it changed.
|
}
|
||||||
|
|
||||||
`dataclasses.replace` builds a new instance, so the copy carries none of
|
|
||||||
what makes a record resolved: no raw snapshot, no declarations, no
|
|
||||||
finished flag. The next publish therefore resolves it again, which
|
|
||||||
drops every decision the stash held -- the late ones (the auto-detected
|
|
||||||
parsers) and the direct ones alike -- and re-runs the device probes in
|
|
||||||
whatever process opened the copy. The Ray paths replace
|
|
||||||
`dist_init_addr` on a resolved record, which is how they reach this.
|
|
||||||
|
|
||||||
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
|
|
||||||
copy set on its own would publish the parent's raw value instead.
|
|
||||||
|
|
||||||
The carry is shallow. The containers are copied so the copy's own
|
|
||||||
declaration does not travel back into the parent, but everything inside
|
|
||||||
them -- the stash entries, the raw-input values, the memoized
|
|
||||||
`ModelConfig` -- is shared. That is fine for what this is for: a copy
|
|
||||||
that immediately crosses a process boundary (Ray actors, the gateway's
|
|
||||||
workers), where pickling severs the sharing. A caller that mutates the
|
|
||||||
copy's deep structure in-process mutates the parent's too.
|
|
||||||
"""
|
|
||||||
replacement = dataclasses.replace(self, **changes)
|
|
||||||
# Provenance, not resolution state: a copy was still launched by
|
|
||||||
# whatever launched its parent, resolved or not.
|
|
||||||
object.__setattr__(replacement, "_launch_command", self.launch_command)
|
|
||||||
if not getattr(self, "_resolution_finished", False):
|
|
||||||
# Not resolved yet: the copy goes through the gate itself.
|
|
||||||
return replacement
|
|
||||||
|
|
||||||
# Everything outside the fields, enumerated from the instance: the raw
|
|
||||||
# snapshot, the stash, and what resolution memoized -- including the
|
|
||||||
# model-configuration memo, which the copy carries over rather than
|
|
||||||
# rebuild.
|
|
||||||
field_names = {field.name for field in dataclasses.fields(self)}
|
|
||||||
for name, value in vars(self).items():
|
|
||||||
if name in field_names or name == "_resolution_finished":
|
|
||||||
continue
|
|
||||||
if isinstance(value, (dict, list, set)):
|
|
||||||
value = copy.copy(value)
|
|
||||||
object.__setattr__(replacement, name, value)
|
|
||||||
stash = getattr(replacement, "_resolved_overrides", None)
|
|
||||||
if stash is None:
|
|
||||||
stash = []
|
|
||||||
object.__setattr__(replacement, "_resolved_overrides", stash)
|
|
||||||
if changes:
|
|
||||||
stash.append((source, dict(changes)))
|
|
||||||
object.__setattr__(replacement, "_resolution_finished", True)
|
|
||||||
return replacement
|
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# CUDA graph configuration resolution
|
# CUDA graph configuration resolution
|
||||||
@@ -663,27 +631,6 @@ def get_global_server_args() -> NoReturn:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def record_writable(server_args: Any):
|
|
||||||
"""Lift the input seal for a resolver that genuinely writes the record.
|
|
||||||
|
|
||||||
There is exactly one: `declare_direct_writes`, which hands the record to an
|
|
||||||
out-of-tree platform plugin that sets fields on it. Those implementations
|
|
||||||
live outside this tree and cannot be converted by editing a resolver here,
|
|
||||||
so the write stays and is captured into the stash afterwards. Naming the
|
|
||||||
exception is the point -- an in-tree resolver that reaches for this is
|
|
||||||
doing something it should be declaring instead.
|
|
||||||
"""
|
|
||||||
frozen = getattr(server_args, "_input_frozen", False)
|
|
||||||
if frozen:
|
|
||||||
object.__setattr__(server_args, "_input_frozen", False)
|
|
||||||
try:
|
|
||||||
yield
|
|
||||||
finally:
|
|
||||||
if frozen:
|
|
||||||
object.__setattr__(server_args, "_input_frozen", True)
|
|
||||||
|
|
||||||
|
|
||||||
def prepare_server_args(argv: list[str]) -> ServerArgs:
|
def prepare_server_args(argv: list[str]) -> ServerArgs:
|
||||||
"""
|
"""
|
||||||
Prepare the server arguments from the command line arguments.
|
Prepare the server arguments from the command line arguments.
|
||||||
|
|||||||
@@ -148,16 +148,19 @@ class TestTheModelConfigCache(CustomTestCase):
|
|||||||
second_checkpoint = self._checkpoint()
|
second_checkpoint = self._checkpoint()
|
||||||
|
|
||||||
server_args = self._resolved(model_path=first_checkpoint)
|
server_args = self._resolved(model_path=first_checkpoint)
|
||||||
copy_ = server_args.replace_resolved(
|
self.assertEqual(model_config_of(server_args).model_path, first_checkpoint)
|
||||||
|
|
||||||
|
# Declaring a new `model_path` moves what the memo is keyed on, so the
|
||||||
|
# next read rebuilds rather than handing back a configuration that
|
||||||
|
# describes the previous checkpoint.
|
||||||
|
declare_resolution(
|
||||||
|
server_args,
|
||||||
"test_the_cache_refills_on_a_resolved_record",
|
"test_the_cache_refills_on_a_resolved_record",
|
||||||
model_path=second_checkpoint,
|
model_path=second_checkpoint,
|
||||||
)
|
)
|
||||||
|
rebuilt = model_config_of(server_args)
|
||||||
rebuilt = model_config_of(copy_)
|
|
||||||
self.assertEqual(rebuilt.model_path, second_checkpoint)
|
self.assertEqual(rebuilt.model_path, second_checkpoint)
|
||||||
self.assertIs(model_config_of(copy_), rebuilt)
|
self.assertIs(model_config_of(server_args), rebuilt)
|
||||||
# The parent keeps the configuration it resolved with.
|
|
||||||
self.assertEqual(model_config_of(server_args).model_path, first_checkpoint)
|
|
||||||
|
|
||||||
def test_a_supplied_configuration_is_handed_back(self):
|
def test_a_supplied_configuration_is_handed_back(self):
|
||||||
"""A configuration nothing in here built carries no key, so nothing
|
"""A configuration nothing in here built carries no key, so nothing
|
||||||
|
|||||||
@@ -436,7 +436,7 @@ 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
|
||||||
declare through `declare_late_resolution`. The declaration is the only
|
declare through `declare_resolution`. The declaration is the only
|
||||||
home for what they decide: the record keeps `--reasoning-parser auto`,
|
home for what they decide: the record keeps `--reasoning-parser auto`,
|
||||||
and the bags a process publishes carry the detected parser.
|
and the bags a process publishes carry the detected parser.
|
||||||
|
|
||||||
@@ -444,14 +444,12 @@ class TestResolutionDeclarations(CustomTestCase):
|
|||||||
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
|
||||||
already-late-resolved fields, which hides exactly this.
|
already-late-resolved fields, which hides exactly this.
|
||||||
"""
|
"""
|
||||||
from sglang.srt.arg_groups.overrides import declare_late_resolution
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
from sglang.srt.runtime_context import get_serving, publish, reset_context
|
from sglang.srt.runtime_context import get_serving, publish, reset_context
|
||||||
|
|
||||||
server_args = self._resolve({"reasoning_parser": "auto"})
|
server_args = self._resolve({"reasoning_parser": "auto"})
|
||||||
self.addCleanup(reset_context)
|
self.addCleanup(reset_context)
|
||||||
declare_late_resolution(
|
declare_resolution(server_args, "template-detection", reasoning_parser="qwen3")
|
||||||
server_args, "template-detection", reasoning_parser="qwen3"
|
|
||||||
)
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
resolution_result(server_args, "reasoning_parser"),
|
resolution_result(server_args, "reasoning_parser"),
|
||||||
"qwen3",
|
"qwen3",
|
||||||
@@ -469,10 +467,10 @@ class TestResolutionDeclarations(CustomTestCase):
|
|||||||
|
|
||||||
def test_pre_engine_late_resolution_reaches_the_projection(self):
|
def test_pre_engine_late_resolution_reaches_the_projection(self):
|
||||||
"""A launcher declaration survives the engine's first resolution pass."""
|
"""A launcher declaration survives the engine's first resolution pass."""
|
||||||
from sglang.srt.arg_groups.overrides import declare_late_resolution
|
from sglang.srt.arg_groups.overrides import declare_resolution
|
||||||
|
|
||||||
server_args = ServerArgs(model_path="dummy")
|
server_args = ServerArgs(model_path="dummy")
|
||||||
declare_late_resolution(
|
declare_resolution(
|
||||||
server_args,
|
server_args,
|
||||||
"launcher",
|
"launcher",
|
||||||
enable_forward_pass_metrics=True,
|
enable_forward_pass_metrics=True,
|
||||||
@@ -695,11 +693,13 @@ class TestResolutionDeclarations(CustomTestCase):
|
|||||||
server_args.attention_backend = "triton"
|
server_args.attention_backend = "triton"
|
||||||
server_args.schedule_conservativeness = 0.5
|
server_args.schedule_conservativeness = 0.5
|
||||||
|
|
||||||
from sglang.srt.arg_groups import pipeline as pipeline_module
|
from sglang.srt import platforms as platforms_module
|
||||||
|
|
||||||
# The write capture runs in the dispatcher, so that is the namespace the
|
# `handle_platform_defaults` imports `current_platform` when it runs, so
|
||||||
# plugin has to be installed in.
|
# the platform module is the namespace to install the plugin in.
|
||||||
with unittest.mock.patch.object(pipeline_module, "current_platform", _Plugin()):
|
with unittest.mock.patch.object(
|
||||||
|
platforms_module, "current_platform", _Plugin()
|
||||||
|
):
|
||||||
server_args = self._resolve({})
|
server_args = self._resolve({})
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
(
|
(
|
||||||
@@ -738,7 +738,7 @@ class TestDeclaredValuesAreNotEditedLater(CustomTestCase):
|
|||||||
|
|
||||||
The property is about the stash, so the seam is the stash: a list that
|
The property is about the stash, so the seam is the stash: a list that
|
||||||
snapshots on append. Every declaration path -- `declare_resolution`,
|
snapshots on append. Every declaration path -- `declare_resolution`,
|
||||||
`declare_late_resolution`, `declare_direct_writes` and the passes --
|
`declare_resolution`, `record_foreign_defaults` and the passes --
|
||||||
reaches it through `.append`, whatever it was imported as.
|
reaches it through `.append`, whatever it was imported as.
|
||||||
"""
|
"""
|
||||||
recorded = []
|
recorded = []
|
||||||
|
|||||||
@@ -35,7 +35,10 @@ import unittest.mock
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import model_config_of, resolution_result
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
declare_resolution,
|
||||||
|
resolution_result,
|
||||||
|
)
|
||||||
from sglang.srt.environ import EnvField, envs
|
from sglang.srt.environ import EnvField, envs
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
@@ -480,15 +483,15 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
|
|||||||
self.assertEqual(getattr(first, "_resolved_overrides", None), first_provenance)
|
self.assertEqual(getattr(first, "_resolved_overrides", None), first_provenance)
|
||||||
|
|
||||||
|
|
||||||
class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
|
class TestALateDeclarationKeepsTheResolution(_RestoresProcessState, CustomTestCase):
|
||||||
"""A resolved record copied with `dataclasses.replace` loses what makes it
|
"""A resolved record copied with `dataclasses.replace` loses what makes it
|
||||||
resolved, and the next publish resolves it a second time -- over values it
|
resolved, and the next publish resolves it a second time -- over values it
|
||||||
already decided. The Ray paths copy a resolved record to set
|
already decided. The Ray paths declare `dist_init_addr` on a record that
|
||||||
`dist_init_addr`, which is how they reach this.
|
has already resolved, which is how they reach this.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def _resolved(self):
|
def _resolved(self):
|
||||||
config_dir = tempfile.mkdtemp(prefix="replace_resolved_")
|
config_dir = tempfile.mkdtemp(prefix="late_declaration_")
|
||||||
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
|
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
|
||||||
with open(os.path.join(config_dir, "config.json"), "w") as handle:
|
with open(os.path.join(config_dir, "config.json"), "w") as handle:
|
||||||
json.dump(_MINI_CONFIG, handle)
|
json.dump(_MINI_CONFIG, handle)
|
||||||
@@ -509,9 +512,9 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
|
|||||||
|
|
||||||
`dataclasses.replace` copies the fields, so a bare copy re-runs
|
`dataclasses.replace` copies the fields, so a bare copy re-runs
|
||||||
resolution over the *same input* the parent got -- the DP-attention
|
resolution over the *same input* the parent got -- the DP-attention
|
||||||
halving and the conservativeness scaling apply once. `replace_resolved`
|
halving and the conservativeness scaling apply once. This is why the Ray
|
||||||
buys something else: it carries the parent's declarations and its
|
paths declare on the record they were handed instead of copying it: the
|
||||||
`model_config`, so the copy answers without resolving at all.
|
record arrives resolved, and a copy would throw that away.
|
||||||
"""
|
"""
|
||||||
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")
|
||||||
@@ -537,56 +540,34 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
|
|||||||
"reading its own output again",
|
"reading its own output again",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_replace_resolved_keeps_the_parents_resolution(self):
|
def test_a_late_change_leaves_the_rest_of_the_resolution_alone(self):
|
||||||
parent = self._resolved()
|
"""What the Ray paths do: declare one field on a record that has already
|
||||||
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
|
resolved, then hand it to the process that will publish it.
|
||||||
self.assertTrue(getattr(copy_, "_resolution_finished", False))
|
|
||||||
drifted = {
|
|
||||||
field.name: (getattr(parent, field.name), getattr(copy_, field.name))
|
|
||||||
for field in dataclasses.fields(parent)
|
|
||||||
if field.name != "dist_init_addr"
|
|
||||||
and getattr(parent, field.name) != getattr(copy_, field.name)
|
|
||||||
}
|
|
||||||
self.assertEqual(
|
|
||||||
drifted,
|
|
||||||
{},
|
|
||||||
f"the copy differs from its parent beyond the change: {drifted}",
|
|
||||||
)
|
|
||||||
self.assertEqual(copy_.dist_init_addr, "1.2.3.4:5000")
|
|
||||||
|
|
||||||
def test_the_copy_carries_what_resolution_left_on_the_record(self):
|
The record stays resolved, so nothing re-derives; the field stays the
|
||||||
"""Not just the stash and the flag.
|
operator's input, because resolution does not write fields; and the
|
||||||
|
decision is what `resolution_result` answers.
|
||||||
`model_config_of()` memoizes on the record, and that cache is filled
|
|
||||||
during resolution. A copy that is marked resolved but arrives without it
|
|
||||||
cannot fill it -- the read-only guard refuses the cache write -- so the
|
|
||||||
first `model_config_of()` raises. That is what killed the Ray
|
|
||||||
schedulers, and it is why the carry is enumerated from the instance
|
|
||||||
rather than from a list of names.
|
|
||||||
"""
|
"""
|
||||||
parent = self._resolved()
|
parent = self._resolved()
|
||||||
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
|
declare_resolution(parent, "ray.test", dist_init_addr="1.2.3.4:5000")
|
||||||
fields = {field.name for field in dataclasses.fields(parent)}
|
|
||||||
missing = sorted(
|
self.assertTrue(getattr(parent, "_resolution_finished", False))
|
||||||
name
|
self.assertIsNone(
|
||||||
for name in vars(parent)
|
parent.dist_init_addr,
|
||||||
if name not in fields and name not in vars(copy_)
|
"the declaration wrote the field; the record is the operator's input",
|
||||||
)
|
|
||||||
self.assertEqual(
|
|
||||||
missing,
|
|
||||||
[],
|
|
||||||
f"the copy did not carry what resolution left on the record: {missing}",
|
|
||||||
)
|
|
||||||
self.assertIsNotNone(model_config_of(copy_))
|
|
||||||
# Containers are copied, so the copy's declaration stays with it.
|
|
||||||
self.assertEqual(
|
|
||||||
len(parent._resolved_overrides) + 1, len(copy_._resolved_overrides)
|
|
||||||
)
|
)
|
||||||
|
self.assertEqual(resolution_result(parent, "dist_init_addr"), "1.2.3.4:5000")
|
||||||
|
|
||||||
def test_the_change_reaches_the_bags(self):
|
def test_the_change_reaches_the_bags(self):
|
||||||
"""The projection reads the raw snapshot plus the declarations, so a
|
"""The projection reads the raw snapshot plus the declarations, so a
|
||||||
change the copy only wrote to the field would publish the parent's raw
|
change written only to the field would publish the raw value instead.
|
||||||
value."""
|
|
||||||
|
This is the Ray hop: the actor receives the record by pickle, declares
|
||||||
|
its own `dist_init_addr`, and publishes. Nothing else may move --
|
||||||
|
publishing must not re-run resolution.
|
||||||
|
"""
|
||||||
|
import pickle
|
||||||
|
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
@@ -595,16 +576,17 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
parent = self._resolved()
|
parent = self._resolved()
|
||||||
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
|
arrived = pickle.loads(pickle.dumps(parent))
|
||||||
|
declare_resolution(arrived, "ray.test", dist_init_addr="1.2.3.4:5000")
|
||||||
self.addCleanup(reset_context)
|
self.addCleanup(reset_context)
|
||||||
reset_context()
|
reset_context()
|
||||||
publish(copy_, role="scheduler")
|
publish(arrived, role="scheduler")
|
||||||
self.assertEqual(get_parallel().dist_init_addr, "1.2.3.4:5000")
|
self.assertEqual(get_parallel().dist_init_addr, "1.2.3.4:5000")
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
get_schedule().chunked_prefill_size,
|
get_schedule().chunked_prefill_size,
|
||||||
resolution_result(parent, "chunked_prefill_size"),
|
resolution_result(parent, "chunked_prefill_size"),
|
||||||
"publishing the copy re-ran resolution; the bag disagrees with what "
|
"publishing re-ran resolution; the bag disagrees with what the "
|
||||||
"the parent's resolution decided",
|
"parent's resolution decided",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -86,8 +86,7 @@ def _field_reads(fn, holders):
|
|||||||
_DECLARERS = frozenset(
|
_DECLARERS = frozenset(
|
||||||
{
|
{
|
||||||
"declare_resolution",
|
"declare_resolution",
|
||||||
"declare_late_resolution",
|
"record_foreign_defaults",
|
||||||
"declare_direct_writes",
|
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -271,7 +270,7 @@ def _record_aliases(function):
|
|||||||
aliases.add(target.id)
|
aliases.add(target.id)
|
||||||
elif (
|
elif (
|
||||||
isinstance(func, ast.Attribute)
|
isinstance(func, ast.Attribute)
|
||||||
and func.attr in ("from_cli_args", "replace_resolved")
|
and func.attr == "from_cli_args"
|
||||||
and isinstance(func.value, ast.Name)
|
and isinstance(func.value, ast.Name)
|
||||||
and func.value.id == "ServerArgs"
|
and func.value.id == "ServerArgs"
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -2886,17 +2886,32 @@ class TestTheInputIsSealedDuringResolution(CustomTestCase):
|
|||||||
with self.assertRaisesRegex(AttributeError, "after resolution"):
|
with self.assertRaisesRegex(AttributeError, "after resolution"):
|
||||||
server_args.tp_size = 4
|
server_args.tp_size = 4
|
||||||
|
|
||||||
def test_the_named_exception_lifts_it(self):
|
def test_it_has_no_exception(self):
|
||||||
"""`declare_direct_writes` hands the record to an out-of-tree platform
|
"""A resolver from outside this tree assigns fields -- an interface this
|
||||||
plugin that sets fields on it; that is the only channel."""
|
tree does not own -- and it still does not reach the record.
|
||||||
from sglang.srt.server_args import record_writable
|
|
||||||
|
`record_foreign_defaults` hands it a stand-in: the assignment is
|
||||||
|
captured and declared, the field keeps the operator's input, and the
|
||||||
|
seal stays armed for the whole call. There used to be a named lift for
|
||||||
|
this, which made the record the one thing resolution could write.
|
||||||
|
"""
|
||||||
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
record_foreign_defaults,
|
||||||
|
resolution_result,
|
||||||
|
)
|
||||||
|
|
||||||
server_args = ServerArgs(model_path="dummy", device="cuda")
|
server_args = ServerArgs(model_path="dummy", device="cuda")
|
||||||
object.__setattr__(server_args, "_input_frozen", True)
|
object.__setattr__(server_args, "_input_frozen", True)
|
||||||
with record_writable(server_args):
|
|
||||||
server_args.tp_size = 4
|
def foreign(config):
|
||||||
self.assertEqual(server_args.tp_size, 4)
|
# What a plugin does: read what is decided, assign a default.
|
||||||
# and it goes back on afterwards
|
assert config.tp_size == 1
|
||||||
|
config.tp_size = 4
|
||||||
|
|
||||||
|
record_foreign_defaults(server_args, "platform:probe", foreign)
|
||||||
|
|
||||||
|
self.assertEqual(resolution_result(server_args, "tp_size"), 4)
|
||||||
|
self.assertEqual(server_args.tp_size, 1, "the record is the input")
|
||||||
with self.assertRaisesRegex(AttributeError, "during resolution"):
|
with self.assertRaisesRegex(AttributeError, "during resolution"):
|
||||||
server_args.tp_size = 8
|
server_args.tp_size = 8
|
||||||
|
|
||||||
@@ -2942,15 +2957,6 @@ class TestLaunchCommand(CustomTestCase):
|
|||||||
server_args.launch_command,
|
server_args.launch_command,
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_a_copy_keeps_it(self):
|
|
||||||
"""`replace_resolved` is how the Ray paths rewrite `dist_init_addr`;
|
|
||||||
the copy was launched by whatever launched its parent."""
|
|
||||||
server_args = prepare_server_args(["--model-path", "/tmp/x"])
|
|
||||||
self.assertEqual(
|
|
||||||
server_args.replace_resolved("test").launch_command,
|
|
||||||
server_args.launch_command,
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_it_is_not_a_config_field(self):
|
def test_it_is_not_a_config_field(self):
|
||||||
"""It describes how the configuration was asked for, so it is not part
|
"""It describes how the configuration was asked for, so it is not part
|
||||||
of the configuration: no CLI flag, no namespace, not in the bags."""
|
of the configuration: no CLI flag, no namespace, not in the bags."""
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ _OWNERS = ("server_args.py", "runtime_context.py", "arg_groups/")
|
|||||||
# startup default wherever it is written, and `benchmark/` ships too.
|
# startup default wherever it is written, and `benchmark/` ships too.
|
||||||
_READS_SCANNED = _PACKAGE
|
_READS_SCANNED = _PACKAGE
|
||||||
|
|
||||||
_DECLARERS = ("declare_resolution", "declare_late_resolution")
|
_DECLARERS = ("declare_resolution",)
|
||||||
|
|
||||||
|
|
||||||
def _declared_by_keyword():
|
def _declared_by_keyword():
|
||||||
@@ -239,27 +239,6 @@ def _declared_by_registry_and_passes():
|
|||||||
return fields
|
return fields
|
||||||
|
|
||||||
|
|
||||||
def _declared_by_late_resolution():
|
|
||||||
"""Keywords of `declare_late_resolution(record, ...)`, the late spelling.
|
|
||||||
|
|
||||||
The fields sit at the call sites rather than in the declarer, so a scan
|
|
||||||
that only knew the declarer's own definition would find none of them.
|
|
||||||
"""
|
|
||||||
# The record plus `arg_groups/`: a hook calls it on the record it was
|
|
||||||
# handed, so scanning the record's file alone finds nothing.
|
|
||||||
sources = [_SRT / "server_args.py", *sorted((_SRT / "arg_groups").rglob("*.py"))]
|
|
||||||
fields = set()
|
|
||||||
for source in sources:
|
|
||||||
for node in ast.walk(ast.parse(source.read_text(encoding="utf-8-sig"))):
|
|
||||||
if (
|
|
||||||
isinstance(node, ast.Call)
|
|
||||||
and isinstance(node.func, ast.Name)
|
|
||||||
and node.func.id == "declare_late_resolution"
|
|
||||||
):
|
|
||||||
fields |= {keyword.arg for keyword in node.keywords if keyword.arg}
|
|
||||||
return fields
|
|
||||||
|
|
||||||
|
|
||||||
def _written_after_publish():
|
def _written_after_publish():
|
||||||
"""Fields the runtime overrides once the bags exist.
|
"""Fields the runtime overrides once the bags exist.
|
||||||
|
|
||||||
@@ -286,7 +265,6 @@ def _resolution_written():
|
|||||||
return (
|
return (
|
||||||
_declared_by_keyword()
|
_declared_by_keyword()
|
||||||
| _declared_by_registry_and_passes()
|
| _declared_by_registry_and_passes()
|
||||||
| _declared_by_late_resolution()
|
|
||||||
| _written_after_publish()
|
| _written_after_publish()
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -561,7 +539,6 @@ class TestNoChainReadsOfResolvedConfig(CustomTestCase):
|
|||||||
|
|
||||||
by_keyword = _declared_by_keyword()
|
by_keyword = _declared_by_keyword()
|
||||||
by_data = _declared_by_registry_and_passes()
|
by_data = _declared_by_registry_and_passes()
|
||||||
by_late = _declared_by_late_resolution()
|
|
||||||
|
|
||||||
self.assertGreater(
|
self.assertGreater(
|
||||||
len(by_keyword),
|
len(by_keyword),
|
||||||
@@ -586,24 +563,10 @@ class TestNoChainReadsOfResolvedConfig(CustomTestCase):
|
|||||||
f"{len(overrides.POST_PROCESS_PASSES)} passes; the scan of the "
|
f"{len(overrides.POST_PROCESS_PASSES)} passes; the scan of the "
|
||||||
"dict-key channel broke",
|
"dict-key channel broke",
|
||||||
)
|
)
|
||||||
self.assertGreaterEqual(
|
|
||||||
len(by_late),
|
|
||||||
3,
|
|
||||||
f"only {len(by_late)} fields are declared late; the "
|
|
||||||
"`declare_late_resolution` keyword scan broke",
|
|
||||||
)
|
|
||||||
# The data channel is not the keyword scan's subset: if it became one,
|
# The data channel is not the keyword scan's subset: if it became one,
|
||||||
# that scan would be doing all the work and a regression here would be
|
# that scan would be doing all the work and a regression here would be
|
||||||
# invisible. The late channel *is* a subset, and deliberately so --
|
# invisible.
|
||||||
# `declare_late_resolution` is a keyword declarer like the others now
|
|
||||||
# that the record hosts no forwarding member, so its own floor above is
|
|
||||||
# what pins it.
|
|
||||||
self.assertTrue(by_data - by_keyword, "the data channel adds nothing")
|
self.assertTrue(by_data - by_keyword, "the data channel adds nothing")
|
||||||
self.assertTrue(
|
|
||||||
by_late <= by_keyword,
|
|
||||||
"late resolution declares outside the keyword channel; it is the "
|
|
||||||
"same spelling, so the two cannot disagree",
|
|
||||||
)
|
|
||||||
|
|
||||||
def test_nothing_reads_a_resolved_field_off_a_borrowed_record(self):
|
def test_nothing_reads_a_resolved_field_off_a_borrowed_record(self):
|
||||||
found = _chain_reads(_resolution_written())
|
found = _chain_reads(_resolution_written())
|
||||||
|
|||||||
@@ -112,7 +112,7 @@ _MATRIX = (
|
|||||||
{"enable_mis": True, "attention_backend": "flashinfer"},
|
{"enable_mis": True, "attention_backend": "flashinfer"},
|
||||||
)
|
)
|
||||||
|
|
||||||
# `declare_late_resolution` call sites whose keyword expansion is built
|
# `declare_resolution` call sites whose keyword expansion is built
|
||||||
# dynamically; the written fields are spelled out here and drift-guarded.
|
# dynamically; the written fields are spelled out here and drift-guarded.
|
||||||
_LATE_RESOLUTION_DYNAMIC_SITES = {
|
_LATE_RESOLUTION_DYNAMIC_SITES = {
|
||||||
"parser/template_detection.py": frozenset({"reasoning_parser", "tool_call_parser"}),
|
"parser/template_detection.py": frozenset({"reasoning_parser", "tool_call_parser"}),
|
||||||
@@ -341,10 +341,10 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
|||||||
union does not depend on matrix order; and the ambient CI marker is
|
union does not depend on matrix order; and the ambient CI marker is
|
||||||
cleared, so a runner's identity cannot leak into the measurement --
|
cleared, so a runner's identity cannot leak into the measurement --
|
||||||
the CI-conditioned writes come from `_ENV_MATRIX`'s explicit entry.
|
the CI-conditioned writes come from `_ENV_MATRIX`'s explicit entry.
|
||||||
Late resolution counts too: `declare_late_resolution` writers run at
|
Declarers outside `arg_groups/` count too: the parser auto-detection
|
||||||
launcher stage (LoRA normalization, parser auto-detection), so their
|
runs at launcher stage and the NPU helper is called by the pipeline, so
|
||||||
target fields are collected statically from the call sites -- they are
|
their target fields are collected statically from the call sites --
|
||||||
resolution writes by definition, just staged after `__post_init__`.
|
resolution writes by definition, just not reached by the matrix.
|
||||||
"""
|
"""
|
||||||
pristine = (dict(os.environ), self._env_field_flags())
|
pristine = (dict(os.environ), self._env_field_flags())
|
||||||
written = set()
|
written = set()
|
||||||
@@ -383,7 +383,7 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
|||||||
for extra, env in _ENV_MATRIX:
|
for extra, env in _ENV_MATRIX:
|
||||||
resolve_one(extra, env)
|
resolve_one(extra, env)
|
||||||
self._restore_process_state(pristine)
|
self._restore_process_state(pristine)
|
||||||
written |= self._late_resolution_written_fields()
|
written |= self._declared_outside_the_pipeline()
|
||||||
written |= self._hook_assignment_targets()
|
written |= self._hook_assignment_targets()
|
||||||
written |= self._record_method_assignment_targets()
|
written |= self._record_method_assignment_targets()
|
||||||
written |= self._declarative_override_fields()
|
written |= self._declarative_override_fields()
|
||||||
@@ -607,22 +607,32 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
|||||||
fields.add(key.value)
|
fields.add(key.value)
|
||||||
return fields
|
return fields
|
||||||
|
|
||||||
def _late_resolution_written_fields(self) -> set:
|
def _declared_outside_the_pipeline(self) -> set:
|
||||||
"""Fields `declare_late_resolution` writes, collected statically.
|
"""Fields declared by a `declare_resolution` caller outside
|
||||||
|
`arg_groups/`, collected statically.
|
||||||
|
|
||||||
These are resolution's launcher-stage writes (they need a tokenizer or
|
Resolution's launcher-stage writes live here -- the auto-detected
|
||||||
adapter load, so they cannot run in `__post_init__`), which the
|
parsers need a tokenizer or chat-template load, so they cannot run in
|
||||||
construct-and-diff pass above never sees. The keywords at the call
|
`__post_init__` -- alongside the NPU default helper and the expert-pack
|
||||||
sites are the written fields; an expansion this cannot resolve fails
|
loader, which the pipeline calls the same way. The construct-and-diff
|
||||||
loudly like the override collector's, except the named dynamic sites
|
pass above never sees any of them.
|
||||||
below, whose field sets are spelled out and drift-guarded (each name
|
|
||||||
must still appear as a constant in the file)."""
|
`arg_groups/` is deliberately excluded: `_hook_assignment_targets`
|
||||||
|
covers it exactly, and it resolves the pipeline's own computed
|
||||||
|
expansions (`record_foreign_defaults` declares a `**` dict this
|
||||||
|
collector's resolver cannot read). The keywords at the call sites are
|
||||||
|
the written fields; an expansion this cannot resolve fails loudly like
|
||||||
|
the override collector's, except the named dynamic sites below, whose
|
||||||
|
field sets are spelled out and drift-guarded (each name must still
|
||||||
|
appear as a constant in the file)."""
|
||||||
written = set()
|
written = set()
|
||||||
root = _PACKAGE_ROOT
|
root = _PACKAGE_ROOT
|
||||||
for path in sorted(root.rglob("*.py")):
|
for path in sorted(root.rglob("*.py")):
|
||||||
rel = path.relative_to(root).as_posix()
|
rel = path.relative_to(root).as_posix()
|
||||||
|
if rel.startswith("arg_groups/"):
|
||||||
|
continue
|
||||||
source = path.read_text(encoding="utf-8-sig")
|
source = path.read_text(encoding="utf-8-sig")
|
||||||
if "declare_late_resolution" not in source:
|
if "declare_resolution" not in source:
|
||||||
continue
|
continue
|
||||||
try:
|
try:
|
||||||
tree = ast.parse(source)
|
tree = ast.parse(source)
|
||||||
@@ -634,12 +644,12 @@ class TestSuppliedInstanceExposure(CustomTestCase):
|
|||||||
and (
|
and (
|
||||||
(
|
(
|
||||||
isinstance(node.func, ast.Name)
|
isinstance(node.func, ast.Name)
|
||||||
and node.func.id == "declare_late_resolution"
|
and node.func.id == "declare_resolution"
|
||||||
)
|
)
|
||||||
or (
|
or (
|
||||||
isinstance(node.func, ast.Attribute)
|
isinstance(node.func, ast.Attribute)
|
||||||
and node.func.attr
|
and node.func.attr
|
||||||
in ("declare_late_resolution", "_late_resolution")
|
in ("declare_resolution", "_declare_resolution")
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user