Files
sglang/python/sglang/srt/runtime_context.py
T

1576 lines
62 KiB
Python

# Copyright 2023-2026 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""A single structured accessor for process-static runtime state.
``get_parallel()`` returns a ``ParallelContext`` whose attributes — tp / dcp / pp /
moe / attn size and rank, plus the process-group handles — each delegate live to
the canonical getter in ``distributed.parallel_state`` / ``layers.dp_attention``.
Returned values are exactly what those getters return; this is a read-through
wrapper, not a cache. It gives call-sites one import and one naming scheme in
place of a dozen free functions, plus a test-only ``override()`` hook to force a
topology without monkeypatching the underlying getters.
``get_server_args()`` returns the process-wide ``ServerArgs``. This is the pristine / resolved-at-startup **read-only** record kept
for debug and reproduction; business code reads resolved config from the
namespace bags below, not from this object. The context owns the storage:
publishing goes through ``RuntimeContext.set_server_args`` (the legacy
``set_global_server_args_for_scheduler`` / ``get_global_server_args`` are thin
shims over this slot).
``get_exec()`` / ``get_memory()`` / ``get_schedule()`` / ``get_device()`` /
``get_model()`` / ``get_spec()`` / ``get_lora()`` / ``get_mm()`` /
``get_disagg()`` / ``get_serving()`` / ``get_observability()`` return the
resolved **config namespace bags** — the single source of truth for config,
snapshotted from ``server_args`` at publish and driven by the ``NS(...)``
metadata on each field (multi-level under ``exec.*``). Reads are attribute
chains (``get_exec().moe.moe_runner_backend``); bags are read-only by bare
assignment (written via ``override``).
``get_flags()`` returns the runtime-flags tier: state that is **not** a pure
function of config (the capture lifecycle, ACTIVE MoE backend, DP runtime) —
never a mirror of config. Flags live in typed dataclass groups; reads and
writes are plain attribute access, and each group offers a transactional,
test-only ``override(**kw)``.
"""
from __future__ import annotations
import dataclasses
import functools
import os
import sys
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
# Imported lazily so this module has no import-time dependencies: any module can
# import get_parallel at module level without risking an import cycle.
def _ps():
from sglang.srt.distributed import parallel_state
return parallel_state
def _dp():
from sglang.srt.layers import dp_attention
return dp_attention
_PARALLEL_FIELDS = frozenset(
{
"world_size",
"world_rank",
"tp_size",
"tp_rank",
"pp_size",
"pp_rank",
"moe_ep_size",
"moe_ep_rank",
"moe_dp_size",
"moe_dp_rank",
"moe_tp_size",
"moe_tp_rank",
"attn_tp_size",
"attn_tp_rank",
"attn_cp_size",
"attn_cp_rank",
"dcp_enabled",
"dcp_size",
"dcp_rank",
"attn_dcp_size",
"attn_dcp_rank",
"attn_dp_size",
"attn_dp_rank",
"world_group",
"tp_group",
"pp_group",
"moe_ep_group",
"moe_dp_group",
"moe_tp_group",
"attn_tp_group",
"attn_cp_group",
"dcp_group",
}
)
class ParallelContext:
"""Parallel-topology namespace.
Live topology (size / rank / group) is read-through via ``@property`` (the
canonical getters). Parallel **config** leaves (``nccl_port``,
``pp_max_micro_batch_size``, ``enable_dp_attention``, …) come from the
published ``parallel`` config bag via ``__getattr__``. Where a config leaf
shares a name with a live property (``tp_size`` …), the property (the live
fact) wins; the same-name==same-value invariant holds once dist is up.
"""
__slots__ = ("_overrides", "_config")
def __init__(self):
self._overrides = {}
self._config = None # parallel config bag, wired at publish
def __getattr__(self, name):
# Reached only for names that are neither a live @property nor a slot:
# serve parallel config leaves from the published bag. The body must
# stay dynamo-traceable — config-leaf reads such as
# ``get_parallel().moe_dense_tp_size`` run inside compiled model
# forwards, and ``object.__getattribute__`` graph-breaks.
if name.startswith("_"):
# No config leaf is underscored; this also breaks the recursion
# when the ``_config`` slot itself is still unset (pickle/copy
# protocols probe attributes before __init__ runs).
raise AttributeError(name)
config = self._config
# ``_fields`` is a plain ``__dict__`` entry on the bag; ``in`` on the
# dict avoids ``_ConfigBag.__contains__`` (not traceable).
if config is not None and name in config._fields:
return getattr(config, name)
detail = (
"not a published parallel config leaf"
if config is not None
else "config not published"
)
raise AttributeError(f"ParallelContext has no {name!r} ({detail})")
def _v(self, name, getter):
overrides = self._overrides
return overrides[name] if name in overrides else getter()
@contextmanager
def override(self, **kwargs):
"""Temporarily force parallel values, restoring on exit. Validates keys and
supports nesting."""
unknown = set(kwargs) - _PARALLEL_FIELDS
if unknown:
raise ValueError(f"unknown parallel field(s): {sorted(unknown)}")
saved = dict(self._overrides)
self._overrides.update(kwargs)
try:
yield self
finally:
self._overrides = saved
@property
def world_size(self) -> int:
return self._v("world_size", _ps().get_world_size)
@property
def world_rank(self) -> int:
return self._v("world_rank", _ps().get_world_rank)
@property
def tp_size(self) -> int:
return self._v("tp_size", _ps().get_tensor_model_parallel_world_size)
@property
def tp_rank(self) -> int:
return self._v("tp_rank", _ps().get_tensor_model_parallel_rank)
@property
def pp_size(self) -> int:
return self._v("pp_size", _ps().get_pipeline_model_parallel_world_size)
@property
def pp_rank(self) -> int:
return self._v("pp_rank", _ps().get_pipeline_model_parallel_rank)
@property
def moe_ep_size(self) -> int:
return self._v("moe_ep_size", _ps().get_moe_expert_parallel_world_size)
@property
def moe_ep_rank(self) -> int:
return self._v("moe_ep_rank", _ps().get_moe_expert_parallel_rank)
@property
def moe_dp_size(self) -> int:
return self._v("moe_dp_size", _ps().get_moe_data_parallel_world_size)
@property
def moe_dp_rank(self) -> int:
return self._v("moe_dp_rank", _ps().get_moe_data_parallel_rank)
@property
def moe_tp_size(self) -> int:
return self._v("moe_tp_size", _ps().get_moe_tensor_parallel_world_size)
@property
def moe_tp_rank(self) -> int:
return self._v("moe_tp_rank", _ps().get_moe_tensor_parallel_rank)
@property
def attn_tp_size(self) -> int:
return self._v("attn_tp_size", _ps().get_attn_tensor_model_parallel_world_size)
@property
def attn_tp_rank(self) -> int:
return self._v("attn_tp_rank", _ps().get_attn_tensor_model_parallel_rank)
@property
def attn_cp_size(self) -> int:
return self._v("attn_cp_size", _ps().get_attn_context_model_parallel_world_size)
@property
def attn_cp_rank(self) -> int:
return self._v("attn_cp_rank", _ps().get_attn_context_model_parallel_rank)
@property
def dcp_size(self) -> int:
return self._v("dcp_size", _ps().get_dcp_world_size)
@property
def dcp_rank(self) -> int:
return self._v("dcp_rank", _ps().get_dcp_rank)
@property
def dcp_enabled(self) -> bool:
def getter():
if _ps().get_dcp_group_no_assert() is None:
return False
return self.dcp_size > 1
return self._v("dcp_enabled", getter)
@property
def attn_dcp_size(self) -> int:
return self._v(
"attn_dcp_size", lambda: self.dcp_size if self.dcp_enabled else 1
)
@property
def attn_dcp_rank(self) -> int:
return self._v(
"attn_dcp_rank", lambda: self.dcp_rank if self.dcp_enabled else 0
)
@property
def attn_dp_size(self) -> int:
return self._v("attn_dp_size", _dp().get_attention_dp_size)
@property
def attn_dp_rank(self) -> int:
return self._v("attn_dp_rank", _dp().get_attention_dp_rank)
@property
def world_group(self) -> Any:
return self._v("world_group", _ps().get_world_group)
@property
def tp_group(self) -> Any:
return self._v("tp_group", _ps().get_tp_group)
@property
def pp_group(self) -> Any:
return self._v("pp_group", _ps().get_pp_group)
@property
def moe_ep_group(self) -> Any:
return self._v("moe_ep_group", _ps().get_moe_ep_group)
@property
def moe_dp_group(self) -> Any:
return self._v("moe_dp_group", _ps().get_moe_dp_group)
@property
def moe_tp_group(self) -> Any:
return self._v("moe_tp_group", _ps().get_moe_tp_group)
@property
def attn_tp_group(self) -> Any:
return self._v("attn_tp_group", _ps().get_attn_tp_group)
@property
def attn_cp_group(self) -> Any:
return self._v("attn_cp_group", _ps().get_attn_cp_group)
@property
def dcp_group(self) -> Any:
return self._v("dcp_group", _ps().get_dcp_group)
class _FlagGroupBase:
"""Shared flag-group behavior: typo-safe writes + transactional ``override()``.
Groups are plain dataclasses; ``__dataclass_fields__`` is the single source
of truth for which leaves exist, so a mistyped name fails loudly instead of
creating a stray attribute.
"""
def __setattr__(self, name: str, value: Any) -> None:
if name not in type(self).__dataclass_fields__:
raise AttributeError(
f"{type(self).__name__} has no flag '{name}' (leaves are "
"declared as dataclass fields; check for typos)"
)
object.__setattr__(self, name, value)
@contextmanager
def override(self, **kwargs):
"""Temporarily force flag values, restoring on exit. Transactional
(keys validated before any write) — the test-only injection
primitive."""
fields = type(self).__dataclass_fields__
unknown = set(kwargs) - set(fields)
if unknown:
raise ValueError(
f"unknown flag(s) for {type(self).__name__}: {sorted(unknown)}"
)
saved = {name: getattr(self, name) for name in kwargs}
for name, value in kwargs.items():
object.__setattr__(self, name, value)
try:
yield self
finally:
for name, value in saved.items():
object.__setattr__(self, name, value)
@dataclasses.dataclass
class CaptureFlags(_FlagGroupBase):
"""Capture-time flags; never frozen (written during cuda-graph capture)."""
# Seeded from server_args at publish; a model whose _can_torch_compile is
# False clears it during warmup (the only post-publish writer).
enable_torch_compile: bool = False
# Set for the duration of decode/spec graph capture (model_capture_mode).
# While set, dispose_tensor() is a no-op so deep_gemm's pre-permute does not
# free hidden_states that the dual-stream MoE shared expert reads afterward.
disable_dispose_tensor: bool = False
@dataclasses.dataclass
class MoeFlags(_FlagGroupBase):
"""MoE runtime flags, materialized by ``initialize_moe_config`` (scheduler
init, after distributed setup). ``a2a_backend`` / ``runner_backend`` /
``disable_fp4_allgather`` are the ACTIVE values: the speculative contexts
in ``layers.moe.utils`` swap them around draft-model forwards. Values are
the parsed enums from ``layers.moe.utils``; ``None`` means "not
initialized yet" and the accessors fall back lazily.
"""
a2a_backend: Any = None
runner_backend: Any = None
speculative_runner_backend: Any = None
speculative_a2a_backend: Any = None
deepep_mode: Any = None
deepep_config: str | None = None
tbo_enabled: bool | None = None
sbo_enabled: bool | None = None
tbo_token_distribution_threshold: float | None = None
disable_fp4_allgather: bool | None = None
quantization: str | None = None
# The shared-experts-fusion decision, per runner — the runner_backend /
# speculative_runner_backend shape. Both leaves are seeded from the config
# intent by ``initialize_moe_config``; each MoE model's gate
# (determine_num_fused_shared_experts) refines the ACTIVE leaf, both ways,
# before its layers build and read it. ``speculative_moe_backend_context``
# brackets a draft's build: on exit the draft's effective decision is
# persisted onto the speculative leaf (inspectable afterwards) and the
# target's ACTIVE value returns.
disable_shared_experts_fusion: bool | None = None
speculative_disable_shared_experts_fusion: bool | None = None
# Lifecycle marker (the capture.disable_dispose_tensor family): set while
# speculative_moe_backend_context is active, so a draft gate's write also
# lands on the speculative leaf.
in_speculative_scope: bool = False
@dataclasses.dataclass
class DpFlags(_FlagGroupBase):
"""DP-attention runtime flags, materialized by ``initialize_dp_attention``
(after distributed setup; reads the model config). Topology values
(sizes/ranks) stay on ``layers.dp_attention`` until the parallel vertical
migrates them."""
enabled: bool = False
use_world_group_for_gather: bool = False
joiner_skip_all_gather: bool = False
# Hybrid-SSM models materialize idle ranks via the MAX_LEN fabricated-row
# conversion (set when hf_config has hybrid_override_pattern).
max_len_with_idle: bool = False
# DP gathered-buffer allocation metadata (model hidden size / dtype /
# device), set by initialize_dp_attention alongside the flags above.
buffer_hidden_size: Any = None
buffer_dtype: Any = None
buffer_device: Any = None
@dataclasses.dataclass
class Flags(_FlagGroupBase):
"""Root of the runtime-flags tier.
Resolved configuration lives on ``server_args`` fields (materialized at
the end of ``__post_init__``) — this tier only carries genuine runtime
state whose value is not a function of the configuration alone, grouped
by lifecycle (``capture``) or subsystem (``moe`` / ``dp``).
"""
capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags)
moe: MoeFlags = dataclasses.field(default_factory=MoeFlags)
dp: DpFlags = dataclasses.field(default_factory=DpFlags)
@dataclasses.dataclass
class Resources(_FlagGroupBase):
"""Process-level resource handles: named slots with one reset lifecycle,
scoped test injection via ``override()``, and the creation/publish
semantics kept in the owning modules' accessors (which are thin shims
over these slots)."""
# CUDA graph memory pool shared across the prefill and decode graph
# backends (created lazily by model_executor.runner_utils.pool).
graph_memory_pool: Any = None
# EPLB: per-process recorder and the publish-once location metadata
# (owning accessors live in sglang.srt.eplb).
expert_distribution_recorder: Any = None
expert_location_metadata: Any = None
# LPLB: layer_id -> solver.
lplb_solvers: dict = dataclasses.field(default_factory=dict)
# Named side streams (see RuntimeContext.get_stream): name -> stream.
streams: dict = dataclasses.field(default_factory=dict)
# Named persistent buffers (see RuntimeContext.get_buffer): name -> tensor.
# Accessors with bespoke semantics (grow-only, per-device keys) manage
# their entries directly.
buffers: dict = dataclasses.field(default_factory=dict)
# Persistent reusable CUDA events for non-EP DP TBO, keyed by
# (kind, subbatch) — see dp_attention._tbo_event for why reuse matters.
tbo_event_pool: dict = dataclasses.field(default_factory=dict)
# State capturers (installed by their subsystems when capture is on).
indexer_capturer: Any = None
experts_capturer: Any = None
# The shared TCPStore created during distributed initialization.
tcp_store: Any = None
# Trace verbosity; the accessor seeds it lazily from SGLANG_TRACE_LEVEL.
trace_level: Any = None
class ForwardFlags:
"""Per-forward runtime flags with one API and two backings.
Flags read only from eager Python are backed by context variables, so
nested scopes and threads stay isolated (a new thread sees the defaults).
Flags that are read or written *inside torch.compile-traced model code*
(``_GRAPH_VISIBLE``) are backed by plain dict slots instead: dynamo
cannot trace ``ContextVar.get``/``set``, while plain reads it guards on
— the storage form these flags had before joining the tier. Their
writers and readers are single-threaded per process (TBO interleaves
ubatches on one thread; attention-TP input scattering excludes TBO), so
context isolation is not needed for correctness.
``scoped(**kw)`` — the one regular write path — restores on exit for
both backings. ``set()`` exists for the legacy unscoped setters' shims.
"""
_DEFAULTS = {
"multi_stream": False,
"moe_output_buffer": None,
# Attention-TP input-scattering (set per forward by
# AttnTpContext.maybe_input_scattered / set_attn_inputs).
"attn_input_scattered": False,
"attn_inputs": None,
# Sticky across forwards: every ForwardBatch construction writes it;
# graph runners force False around capture.
"is_extend_in_batch": False,
# Per-layer MLP collective control (set by decoder via scoped()
# around the MLP / MoE / hybrid mixer call).
# fuse_mlp_allreduce: next residual+LN absorbs the post-MLP all-reduce.
# mlp_reduce_scatter: postprocess will reduce-scatter (skip MLP AR).
# flashinfer_trtllm_bypass: deepseek dual-stream graph topk bypass.
"fuse_mlp_allreduce": False,
"mlp_reduce_scatter": False,
"flashinfer_trtllm_bypass": False,
}
# Read/written inside compiled graphs (vocab embedding, communicator,
# EP dispatch, DP gather/scatter, MLP/MoE skip-AR): plain-slot backed.
# Before moving a flag out of this set, prove no read/write site sits
# under torch.compile.
_GRAPH_VISIBLE = frozenset(
{
"attn_input_scattered",
"attn_inputs",
"is_extend_in_batch",
"fuse_mlp_allreduce",
"mlp_reduce_scatter",
"flashinfer_trtllm_bypass",
}
)
__slots__ = ("_vars", "_plain")
def __init__(self):
import contextvars
object.__setattr__(
self,
"_plain",
{
name: default
for name, default in self._DEFAULTS.items()
if name in self._GRAPH_VISIBLE
},
)
object.__setattr__(
self,
"_vars",
{
name: contextvars.ContextVar(f"forward.{name}", default=default)
for name, default in self._DEFAULTS.items()
if name not in self._GRAPH_VISIBLE
},
)
def __getattr__(self, name: str) -> Any:
plain = self._plain
if name in plain:
return plain[name]
try:
return self._vars[name].get()
except KeyError:
raise AttributeError(
f"ForwardFlags has no flag '{name}' (flags are declared in "
"ForwardFlags._DEFAULTS; check for typos)"
) from None
def __setattr__(self, name: str, value: Any) -> None:
raise AttributeError(
"ForwardFlags is written through scoped(**kw) (or the legacy "
"set() shim), never by attribute assignment"
)
def set(self, name: str, value: Any) -> None:
"""Unscoped write for legacy setter shims; persists until the next
write (current context only, for contextvar-backed flags)."""
if name in self._plain:
self._plain[name] = value
else:
self._vars[name].set(value)
@contextmanager
def scoped(self, **kwargs):
"""Set flags for the current scope, restoring on exit. Transactional
(keys validated before any write) and exception-safe."""
unknown = set(kwargs) - set(self._DEFAULTS)
if unknown:
raise ValueError(f"unknown forward flag(s): {sorted(unknown)}")
plain_saved = [
(name, self._plain[name]) for name in kwargs if name in self._plain
]
tokens = []
for name, value in kwargs.items():
if name in self._plain:
self._plain[name] = value
else:
tokens.append((self._vars[name], self._vars[name].set(value)))
try:
yield self
finally:
for var, token in reversed(tokens):
var.reset(token)
for name, value in reversed(plain_saved):
self._plain[name] = value
class _ConfigBag:
"""A resolved-config namespace bag.
Values are snapshotted from ``server_args`` at ``publish`` and this bag is
the **single source of truth** for its fields thereafter. Read is plain
attribute access; the bag is read-only by bare assignment. The sanctioned
writers are ``get_context().override(source, ...)`` (permanent) and
the scoped ``.override(**kw)`` context manager (tests). Sub-namespaces
(e.g. ``exec.moe``) are nested ``_ConfigBag`` instances reached by attribute.
Leaves and sub-bags are stored as **real instance attributes** (in
``__dict__``), so ``bag.leaf`` / ``bag.sub`` is a plain attribute load that
``torch.compile`` / dynamo can trace — config reads inside a compiled model
forward (e.g. ``get_exec().comm.enable_symm_mem`` in the embedding layer)
must not graph-break. ``_fields`` / ``_subs`` keep the authoritative
name→value maps used for override routing, membership, and scoped restore;
``__getattr__`` is only a fallback for genuinely absent names. (Deliberately
no ``__slots__``: leaves are dynamic, and the ``__dict__`` is what makes the
reads traceable.)
"""
def __init__(self, path: str):
object.__setattr__(self, "_path", path)
object.__setattr__(self, "_fields", {}) # {leaf: value}
object.__setattr__(self, "_subs", {}) # {subname: _ConfigBag}
def __getattr__(self, name: str) -> Any:
# Fallback only: real leaves/sub-bags resolve via __dict__ before this
# runs. Uses object.__getattribute__ (not self._fields) to stay safe if
# invoked before __init__ populates the bookkeeping dicts.
fields = object.__getattribute__(self, "_fields")
if name in fields:
return fields[name]
subs = object.__getattribute__(self, "_subs")
if name in subs:
return subs[name]
path = object.__getattribute__(self, "_path")
raise AttributeError(f"config namespace {path!r} has no leaf/subgroup {name!r}")
def __setattr__(self, name: str, value: Any) -> None:
raise AttributeError(
f"config namespace {self._path!r} is read-only; write via "
"get_context().override(source, ...) or the scoped .override(**kw)"
)
def _set(self, name: str, value: Any) -> None:
"""Internal write (publish + override) that bypasses the read-only guard.
Updates both the bookkeeping map and the real attribute (traceable read)."""
object.__getattribute__(self, "_fields")[name] = value
object.__setattr__(self, name, value)
def _set_sub(self, name: str, sub: _ConfigBag) -> None:
"""Register a nested bag as both a bookkeeping entry and a real
attribute (so ``bag.sub`` is a plain, traceable attribute load)."""
object.__getattribute__(self, "_subs")[name] = sub
object.__setattr__(self, name, sub)
def __contains__(self, name: str) -> bool:
return name in object.__getattribute__(self, "_fields")
@contextmanager
def override(self, **kwargs):
"""Scoped, transactional override of this bag's own leaves (keys
validated before any write; restored on exit).
For a window where one runner's value differs from the process's — a
draft model loading under ``--speculative-draft-load-format`` while the
target keeps ``--load-format`` — and for tests forcing a code path.
A permanent change goes through ``get_context().override``."""
fields = object.__getattribute__(self, "_fields")
unknown = set(kwargs) - set(fields)
if unknown:
path = object.__getattribute__(self, "_path")
raise ValueError(f"unknown config leaf for {path!r}: {sorted(unknown)}")
saved = {name: fields[name] for name in kwargs}
for name, value in kwargs.items():
self._set(name, value)
try:
yield self
finally:
for name, value in saved.items():
self._set(name, value)
def _build_config_bags(server_args: Any) -> dict:
"""Snapshot resolved ``server_args`` into the namespace bag tree, driven by
the ``NS(...)`` metadata on the dataclass fields. Returns
``{top_level_name: _ConfigBag}``, arbitrarily nested (``exec.moe.eplb.…``).
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
subgroup at the same level is a hard error — no silent shadowing."""
from sglang.srt.arg_groups.arg_utils import namespace_of
_MISSING = object()
tops: dict = {}
for field, path in namespace_of(type(server_args)).items():
value = getattr(server_args, field, _MISSING)
if value is _MISSING:
# Every NS-declared field is a dataclass field, so a resolved config
# always carries it; a miss means a malformed/partial config object
# was published. Fail loud here rather than silently omitting the
# leaf (which surfaces later as a confusing "not a published leaf").
raise AttributeError(
f"config field {field!r} is declared NS({path!r}) but absent from "
f"the published {type(server_args).__name__}; cannot project its bag leaf"
)
parts = path.split(".")
bag = tops.get(parts[0])
if bag is None:
bag = tops[parts[0]] = _ConfigBag(parts[0])
for depth in range(1, len(parts)):
name = parts[depth]
if name in object.__getattribute__(bag, "_fields"):
raise ValueError(
f"config namespace collision: {'.'.join(parts[: depth + 1])!r} "
"is declared as both a leaf and a subgroup"
)
subs = object.__getattribute__(bag, "_subs")
child = subs.get(name)
if child is None:
child = _ConfigBag(".".join(parts[: depth + 1]))
bag._set_sub(name, child)
bag = child
if field in object.__getattribute__(bag, "_subs"):
raise ValueError(
f"config namespace collision: leaf {field!r} under {path!r} "
"clashes with a subgroup of the same name"
)
bag._set(field, value)
return tops
def _snapshot_bag_values(bags: dict | None) -> dict | None:
"""Per-leaf value snapshot of a config-bag tree (bags are mutated in
place by ``override``, so reference snapshots alias live state)."""
if bags is None:
return None
snap: dict = {}
def walk(prefix: str, bag) -> None:
snap[prefix] = dict(object.__getattribute__(bag, "_fields"))
for name, sub in object.__getattribute__(bag, "_subs").items():
walk(f"{prefix}.{name}", sub)
for name, bag in bags.items():
walk(name, bag)
return snap
def _restore_bag_values(bags: dict, snap: dict) -> None:
def walk(prefix: str, bag) -> None:
for key, value in snap[prefix].items():
bag._set(key, value)
for name, sub in object.__getattribute__(bag, "_subs").items():
walk(f"{prefix}.{name}", sub)
for name, bag in bags.items():
walk(name, bag)
class RuntimeContext:
"""Container for the structured runtime accessors; exposes ``parallel``,
``server_args``, the resolved config namespace bags, ``flags``,
``resources``, and ``forward``."""
__slots__ = (
"parallel",
"_server_args",
"_config_bags",
"_overrides_log",
"_publish_role",
"flags",
"resources",
"forward",
)
def __init__(self, parallel: ParallelContext):
self.parallel = parallel
self._server_args: ServerArgs | None = None
self._config_bags: dict | None = None
self._overrides_log: list = []
self._publish_role: str | None = None
self.flags = Flags()
self.resources = Resources()
self.forward = ForwardFlags()
def get_stream(self, name: str) -> Any:
"""Named process-level CUDA side stream: get-or-create, shared by
name (the keyed-lazy pattern of the persistent buffers). Creation is
a driver call that must stay outside cuda-graph capture — call sites
lease their stream at init/warmup time."""
stream = self.resources.streams.get(name)
if stream is None:
import torch
stream = torch.cuda.Stream()
self.resources.streams[name] = stream
return stream
def set_stream(self, name: str, stream: Any) -> Any:
"""Install (or replace) the named stream — explicit injection for
tests and backends that bring their own stream."""
self.resources.streams[name] = stream
return stream
def get_buffer(self, name: str, factory: Any) -> Any:
"""Named process-level persistent buffer: get-or-create via
``factory()``, shared by name (the keyed-lazy pattern of the
persistent buffers / named streams)."""
buf = self.resources.buffers.get(name)
if buf is None:
buf = factory()
self.resources.buffers[name] = buf
return buf
@property
def server_args(self) -> ServerArgs:
"""The process-wide ``ServerArgs`` (context-owned slot)."""
server_args = self._server_args
if server_args is None:
# Verbatim legacy message: tests and user scripts may match on it.
raise ValueError("Global server args is not set yet!")
return server_args
def set_server_args(self, server_args: ServerArgs) -> None:
"""Publish the process-wide ``ServerArgs`` into the context-owned slot.
Overwrite-allowed: a re-publish replaces the slot (test kits re-publish
per test; production ordering discipline lives at the call-sites, e.g.
the draft-worker guard in ``ModelRunner.__init__``). The published
object already carries the resolved configuration (declarations
materialize at the end of ``__post_init__``).
"""
# Seed the capture tier for the new lifecycle (defaults for sentinel
# and mock publishes, which carry no config).
self.flags.capture.enable_torch_compile = getattr(
server_args, "enable_torch_compile", False
)
self._server_args = server_args
# The adaptive draft-token bound memoizes on the config *path*, so a new
# publication that reuses the path must not inherit the bound computed
# from the file's previous contents.
_adaptive_draft_token_bound.cache_clear()
# Snapshot resolved config into the namespace bags (the single source of
# truth for config reads). Driven by NS(...) metadata; a mock/partial
# config with no NS markers yields an empty tree (no bags projected).
self._config_bags = _build_config_bags(server_args)
# Wire the parallel config leaves onto the live wrapper (config-only
# leaves like pp_max_micro_batch_size are served via ParallelContext
# __getattr__; live topology properties still win by name).
self.parallel._config = self._config_bags.get("parallel")
# A direct install is roleless; ``publish`` assigns the role afterwards.
self._overrides_log = []
self._publish_role = None
def config_bag(self, name: str) -> _ConfigBag:
"""Return the top-level config namespace bag (``device`` / ``model`` /
``exec`` / ``schedule`` / ``memory`` / ``spec`` / ``lora`` / ``mm`` /
``disagg`` / ``serving`` / ``observability``). Fails closed until
``publish`` / ``set_server_args`` has projected it."""
bags = self._config_bags
if not bags or name not in bags:
raise ValueError(f"config namespace {name!r} not published")
if _ROLE_NS_MODE != "off":
self._check_role_namespace(name)
return bags[name]
def _check_role_namespace(self, name: str) -> None:
# Out of line so the mode gate above stays one dead-branch-prunable
# check under dynamo in the default "off" mode (config_bag runs inside
# compiled model forwards).
role = self._publish_role
if _ROLE_NS_MODE == "record":
if not _is_compiling():
_record_namespace_read(role, name)
elif _ROLE_NS_MODE == "enforce" and role is not None:
if role not in ROLE_NAMESPACE_SETS:
raise ValueError(
f"publish role {role!r} has no ROLE_NAMESPACE_SETS entry; "
"declare its namespace set (None for the full tree)."
)
allowed = ROLE_NAMESPACE_SETS[role]
if allowed is not None and name not in allowed:
raise ValueError(
f"config namespace {name!r} is outside the declared set "
f"for publish role {role!r} ({sorted(allowed)}). If this "
"read is legitimate for the process type, extend "
"ROLE_NAMESPACE_SETS; if not, the read belongs in a "
"different process or behind a per-instance boundary."
)
def override(self, source: str, **fields) -> None:
"""The business mutation entry: write resolved config
leaves onto the namespace bags — the single source of truth. It does
**not** touch ``server_args`` (the pristine startup record) and there is
no write-through, so the old "wrote one store, read another" desync class
cannot occur.
Each flat field name is routed to its bag by the ``NS`` metadata (flat
names are unique across namespaces). Validation is all-or-nothing: an
unknown / unprojected field aborts before any write. ``source`` is
recorded for provenance / reproduction.
"""
if not fields:
return
bags = self._config_bags
if bags is None:
raise ValueError("config not published; cannot override")
from sglang.srt.arg_groups.arg_utils import namespace_of
nsmap = namespace_of(type(self._server_args))
targets = [] # (bag, leaf, value) — resolved before any write
for name, value in fields.items():
path = nsmap.get(name)
if path is None:
raise ValueError(
f"override: unknown config field {name!r} (no NS namespace) — "
"not a resolved config leaf"
)
parts = path.split(".")
bag = bags.get(parts[0])
if bag is None:
raise ValueError(f"override: namespace {parts[0]!r} not published")
for seg in parts[1:]:
bag = object.__getattribute__(bag, "_subs").get(seg)
if bag is None:
raise ValueError(
f"override: subgroup {seg!r} missing under {path!r}"
)
if name not in bag:
raise ValueError(f"override: field {name!r} not projected on {path!r}")
targets.append((bag, name, value))
for bag, name, value in targets:
bag._set(name, value)
self._overrides_log.append((source, dict(fields)))
def overrides_log(self) -> list:
"""Provenance of post-publish ``override`` calls: ``[(source, {field: value})]``.
Returns deep-ish copies (source, dict(fields)) so callers inspecting the
log cannot mutate the recorded provenance in place."""
return [(source, dict(fields)) for source, fields in self._overrides_log]
def resolved_server_args_dict(self, base: dict | None = None) -> dict:
"""Serialize the *resolved* config: the pristine ``server_args`` fields
with every post-publish ``override`` overlaid.
``get_internal_state`` reports this, and ``/server_info`` carries it in
the ``internal_states`` block, so scheduler-side runtime changes show up
in a readback: HiCache attach/detach, the generated forward-pass-metrics
endpoint, tunables set via ``/set_internal_state``.
``base`` defaults to ``dict(vars(server_args))`` (matching the legacy
``vars`` dump); pass ``dataclasses.asdict(server_args)`` when nested
dataclass fields must be expanded first. Override leaves are flat
``ServerArgs`` field names, so overlaying them onto the top level of
either base is exact.
This covers the process-global bags only. Per-engine control-plane
changes (weight version, model path, the tokenizer's HiCache mirror)
live on the tokenizer manager — several ``Engine``s can share one
process — and ``TokenizerManager.resolved_config_dict`` overlays those
for the top-level ``/server_info`` body. The two are separate logs, not
one merged dict.
"""
d = dict(vars(self.server_args)) if base is None else dict(base)
for _source, fields in self._overrides_log:
d.update(fields)
return d
def override_server_args(self, **fields) -> _ServerArgsOverride:
"""Test-only scoped override for the config tier — the sibling of
``get_parallel().override()`` and the flag groups' ``override()``:
tests force execution paths by overriding the context instead of
hand-building config objects.
``install()`` (or entering it as a context manager) publishes a fresh
dummy-boundary ``ServerArgs`` carrying ``fields`` and returns it;
``restore()`` (or exiting) reinstates whatever the slot held before.
This is the sanctioned way for a test to get a published context, and
it stays. The transitional reason it was introduced for — production
code branching on raw ``server_args`` fields at runtime — is gone (the
read ratchet pins business reads at zero), but a test that exercises
bag readers still needs bags, and the bag tree is projected *from an
instance*: something has to publish one. Prefer the finer-grained
scoped overrides (``get_exec().override(...)``, the flag groups'
``override``) on top of a published context when a test only needs to
force one leaf.
"""
return _ServerArgsOverride(self, fields)
@contextmanager
def preserve_config(self):
"""Snapshot the full config lifecycle and reinstate it verbatim on exit.
For nested construction steps that publish a private ``ServerArgs``
copy (e.g. a draft-worker build) and must leave the enclosing
lifecycle — including its post-publish overrides — untouched.
"""
prev_server_args = self._server_args
prev_bags = self._config_bags
prev_bag_values = _snapshot_bag_values(prev_bags)
prev_overrides_log = list(self._overrides_log)
prev_publish_role = self._publish_role
prev_parallel_config = self.parallel._config
prev_capture = self.flags.capture.enable_torch_compile
try:
yield
finally:
self._server_args = prev_server_args
self._config_bags = prev_bags
if prev_bags is not None:
_restore_bag_values(prev_bags, prev_bag_values)
self._overrides_log = prev_overrides_log
self._publish_role = prev_publish_role
self.parallel._config = prev_parallel_config
self.flags.capture.enable_torch_compile = prev_capture
class _ServerArgsOverride:
"""Scoped config override (see ``RuntimeContext.override_server_args``).
Deliberately a plain class rather than a generator context manager:
fixtures that live for a whole test case install the override without a
``with`` block, and a suspended generator would run its restore whenever
the garbage collector closes it — un-publishing the active config at a
nondeterministic point.
"""
__slots__ = (
"_context",
"_fields",
"_prev_server_args",
"_prev_bags",
"_prev_overrides_log",
"_prev_publish_role",
"_prev_parallel_config",
"_prev_capture",
"_installed",
)
def __init__(self, context: RuntimeContext, fields: dict):
self._context = context
self._fields = fields
self._installed = False
def install(self) -> ServerArgs:
"""Publish a fresh dummy-boundary ``ServerArgs`` carrying the
overrides; returns the published instance."""
from sglang.srt.server_args import ServerArgs
assert not self._installed, "override_server_args already installed"
ctx = self._context
self._prev_server_args = ctx._server_args
self._prev_bags = ctx._config_bags
self._prev_overrides_log = ctx._overrides_log
self._prev_publish_role = ctx._publish_role
self._prev_parallel_config = ctx.parallel._config
self._prev_capture = ctx.flags.capture.enable_torch_compile
from sglang.srt.arg_groups.overrides import _apply_fields
server_args = ServerArgs(model_path="dummy")
# Underscore names seed private property caches (the strict guard
# exempts them); everything else must be a real config field.
unknown = {name for name in self._fields if not name.startswith("_")} - set(
type(server_args).__dataclass_fields__
)
if unknown:
raise ValueError(
f"override_server_args: unknown ServerArgs field(s): {sorted(unknown)}"
)
_apply_fields(server_args, self._fields)
# The dummy boundary skips materialization, which would leave the
# strict mutation guard unarmed on the published object — mark it
# materialized so bare post-publish writes raise like they do on a
# fully resolved config.
object.__setattr__(server_args, "_declarations_materialized", True)
ctx.set_server_args(server_args)
self._installed = True
return server_args
def restore(self) -> None:
"""Reinstate the exact pre-install lifecycle state (or the empty slot)."""
if not self._installed:
return
self._installed = False
ctx = self._context
ctx._server_args = self._prev_server_args
ctx._config_bags = self._prev_bags
ctx._overrides_log = self._prev_overrides_log
ctx._publish_role = self._prev_publish_role
ctx.parallel._config = self._prev_parallel_config
ctx.flags.capture.enable_torch_compile = self._prev_capture
self._prev_server_args = None
self._prev_bags = None
self._prev_overrides_log = None
self._prev_parallel_config = None
def __enter__(self) -> ServerArgs:
return self.install()
def __exit__(self, *exc) -> None:
self.restore()
_PARALLEL = ParallelContext()
_CONTEXT = RuntimeContext(parallel=_PARALLEL)
def get_context() -> RuntimeContext:
return _CONTEXT
def get_parallel() -> ParallelContext:
return _PARALLEL
def get_server_args() -> ServerArgs:
return _CONTEXT.server_args
def get_flags() -> Flags:
return _CONTEXT.flags
def get_resources() -> Resources:
return _CONTEXT.resources
def get_forward() -> ForwardFlags:
return _CONTEXT.forward
# --- Resolved config namespaces -------------------------
# Each returns the top-level snapshot bag; reads are `get_exec().moe.field` etc.
# All fail with ValueError("... not published") until publish has projected them.
# ``parallel`` config leaves are served by ``get_parallel()`` (live wrapper);
# their config-bag wiring is a scoped follow-up.
def get_device() -> _ConfigBag:
return _CONTEXT.config_bag("device")
def get_model() -> _ConfigBag:
return _CONTEXT.config_bag("model")
def get_exec() -> _ConfigBag:
return _CONTEXT.config_bag("exec")
def get_schedule() -> _ConfigBag:
return _CONTEXT.config_bag("schedule")
def get_memory() -> _ConfigBag:
return _CONTEXT.config_bag("memory")
def get_spec() -> _ConfigBag:
return _CONTEXT.config_bag("spec")
def get_lora() -> _ConfigBag:
return _CONTEXT.config_bag("lora")
def get_mm() -> _ConfigBag:
return _CONTEXT.config_bag("mm")
def get_disagg() -> _ConfigBag:
return _CONTEXT.config_bag("disagg")
def get_serving() -> _ConfigBag:
return _CONTEXT.config_bag("serving")
def get_observability() -> _ConfigBag:
return _CONTEXT.config_bag("observability")
# --- Per-role namespace sets (2c) -------------------------------------------
#
# ``publish(role=...)`` records which process type installed the config; this
# table declares which top-level config namespaces each role reads. ``None``
# means the full tree — either the role genuinely needs everything (scheduler)
# or its deployment shape has not been audited yet (restrict only what smoke
# coverage can verify). ``parallel`` is served by ``get_parallel()`` and every
# process legitimately reads topology config, so it is not part of this table.
#
# ``SGLANG_ROLE_NAMESPACES`` selects the mode (read once at import):
# off (default) no bookkeeping, zero overhead;
# record audit mode — collect (role, namespace) reads per process and dump
# them at exit (the data that seeds this table). Reads made inside
# torch.compile-traced code are NOT observed (recording is pruned
# under tracing to keep capture legal) — run audits with
# compilation disabled before restricting a role.
# enforce fail closed — a bag read outside the role's declared set raises.
ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
# Reads (almost) everything by design — the model-executing process.
"scheduler": None,
"launcher": None,
"test": None,
# Audited (record-mode smokes, plain + DP-attention): the DP controller
# reads only the elastic-EP gate; its module's static read set agrees.
"dp_controller": frozenset({"exec"}),
# Record-mode audit (2026-08-06, text model, /generate + /get_server_info +
# /v1/models): reads exactly {"serving"} — the per-instance managers read
# self.server_args by design. Still declared full, because that run did not
# exercise the multimodal processors, LoRA/score endpoints, the disagg
# roles, or the gRPC bridge; narrowing needs those shapes audited too, and
# a wrong set fails a request rather than a test.
"tokenizer": None,
# Deployment shapes not exercised locally; audit before restricting.
"encoder": None,
"expert_backup": None,
"weight_cache_daemon": None,
}
def _validated_role_ns_mode(value: str) -> str:
mode = value.strip().lower()
if mode not in ("off", "record", "enforce"):
raise ValueError(
f"SGLANG_ROLE_NAMESPACES={value!r} is not one of off / record / "
"enforce — refusing to guess (a typo here would silently disable "
"enforcement)."
)
return mode
def _role_ns_mode_from_env() -> str:
# Resolved once at import so the config_bag gate stays a dynamo-prunable
# constant; validated fail-loud here (EnvField's warn-and-default parse
# would silently turn a typo into "off").
from sglang.srt.environ import envs
return _validated_role_ns_mode(envs.SGLANG_ROLE_NAMESPACES.get())
_ROLE_NS_MODE = _role_ns_mode_from_env()
_RECORDED_NS_READS: set[tuple[str | None, str]] = set()
_RECORD_DUMP_REGISTERED = False
def _is_compiling() -> bool:
# Recording has Python side effects (set mutation, file I/O, atexit) that
# must never run under tracing; torch.compiler.is_compiling() is dynamo's
# sanctioned probe. The lazy lookup keeps this module import-light.
torch = sys.modules.get("torch")
return torch is not None and torch.compiler.is_compiling()
def _ensure_record_dump_registered() -> None:
global _RECORD_DUMP_REGISTERED
if not _RECORD_DUMP_REGISTERED:
_RECORD_DUMP_REGISTERED = True
import atexit
atexit.register(_dump_recorded_namespace_reads)
def _append_role_ns_out(role: str | None, name: str) -> None:
# Persist immediately: worker processes are routinely torn down with
# signals that skip atexit, and the audit must survive that.
from sglang.srt.environ import envs
out = envs.SGLANG_ROLE_NAMESPACES_OUT.get()
if not out:
return
try:
with open(out, "a") as f:
f.write(f"{role} {name}\n")
except OSError as e:
# The entry stays in the in-memory set; the exit summary still covers it.
print(
f"[role-namespaces] pid={os.getpid()} failed to append "
f"({role}, {name}) to {out!r}: {e}",
file=sys.stderr,
flush=True,
)
def _record_namespace_read(role: str | None, name: str) -> None:
if (role, name) in _RECORDED_NS_READS:
return
_RECORDED_NS_READS.add((role, name))
_append_role_ns_out(role, name)
_ensure_record_dump_registered()
def _dump_recorded_namespace_reads() -> None:
"""Emit the record-mode audit: one line per role with the namespaces its
process actually read (multi-process runs dump once per process). The
process's own publish role is always included, so a zero-read role emits
an (empty) line rather than being indistinguishable from a process where
recording never ran."""
by_role: dict = {}
own_role = _CONTEXT._publish_role
if own_role is not None:
by_role.setdefault(own_role, set())
for role, name in _RECORDED_NS_READS:
if name == "-": # publish-time marker, not a namespace read
by_role.setdefault(role, set())
continue
by_role.setdefault(role, set()).add(name)
for role in sorted(by_role, key=str):
print(
f"[role-namespaces] pid={os.getpid()} role={role} "
f"read={','.join(sorted(by_role[role]))}",
file=sys.stderr,
flush=True,
)
def publish(server_args, *, role: str, hf_config: Any = None) -> RuntimeContext:
"""Install process-wide config for this OS process.
Records the process ``role`` (``tokenizer`` / ``scheduler`` /
``dp_controller`` / ``encoder`` / ``expert_backup`` /
``weight_cache_daemon`` / ``launcher`` / ``test``) and
projects the config bags. Draft workers skip publish (they must not clobber
the target). ``role`` is provenance, and — when ``SGLANG_ROLE_NAMESPACES``
is ``enforce`` — the key into ``ROLE_NAMESPACE_SETS`` for fail-closed
namespace-read enforcement (``record`` audits the reads instead).
``hf_config`` is accepted for forward-compat and currently unused.
Normally one call per process, but re-publish is allowed and is
**last-publish-wins** (bags re-projected, provenance reset, role
overwritten). Two sanctioned multi-publish shapes exist: the in-process
Engine builds its ``TokenizerManager`` inside the launcher process (the
process ends up with the tokenizer publish), and multiple Engines in one
process publish in sequence — which is exactly why per-instance managers
must read ``self.server_args`` for anything engine-specific rather than
the process-global bags.
"""
if _ROLE_NS_MODE == "enforce" and role not in ROLE_NAMESPACE_SETS:
# Fail closed at publish time, not at the first stray read.
raise ValueError(
f"publish role {role!r} has no ROLE_NAMESPACE_SETS entry; declare "
"its namespace set (None for the full tree)."
)
_CONTEXT.set_server_args(server_args)
_CONTEXT._publish_role = role
if _ROLE_NS_MODE == "record":
# The '-' marker distinguishes a zero-read role from a process where
# recording never ran (signal teardown skips atexit).
_record_namespace_read(role, "-")
print(
f"[role-namespaces] pid={os.getpid()} role={role} recording; note: "
"reads inside torch.compile-traced code are not observed — audit "
"with compilation disabled before restricting a role.",
file=sys.stderr,
flush=True,
)
return _CONTEXT
def publish_role() -> str | None:
"""The role recorded by the last ``publish`` (None for a legacy set)."""
return _CONTEXT._publish_role
def get_stream(name: str) -> Any:
return _CONTEXT.get_stream(name)
def set_stream(name: str, stream: Any) -> Any:
return _CONTEXT.set_stream(name, stream)
def get_buffer(name: str, factory: Any) -> Any:
return _CONTEXT.get_buffer(name, factory)
_GLOBAL_DWDP_MANAGER: Any = None
def get_global_dwdp_manager() -> Any:
return _GLOBAL_DWDP_MANAGER
def set_global_dwdp_manager(manager: Any) -> None:
global _GLOBAL_DWDP_MANAGER
_GLOBAL_DWDP_MANAGER = manager
def reset_context() -> None:
"""Clear the context-owned store (unit-test teardown): drop the published
``server_args`` and install fresh ``Flags`` and ``Resources``.
Wrapper subsystems (``parallel``) hold no state and are unaffected.
"""
_CONTEXT._server_args = None
_CONTEXT._config_bags = None
_adaptive_draft_token_bound.cache_clear()
_CONTEXT._overrides_log = []
_CONTEXT._publish_role = None
_CONTEXT.parallel._config = None
_CONTEXT.flags = Flags()
_CONTEXT.resources = Resources()
_CONTEXT.forward = ForwardFlags()
set_global_dwdp_manager(None)
def mamba_extra_buffer_enabled() -> bool:
"""Whether the mamba radix cache keeps its extra state buffer.
A predicate over two published leaves (``memory.disable_radix_cache`` and
``exec.mamba.mamba_radix_cache_strategy``), so it reads the bags rather
than the startup record — the ``ServerArgs`` member of the same name is the
pre-publish equivalent used inside the resolution pipeline.
"""
return (
get_memory().disable_radix_cache is False
and get_exec().mamba.mamba_radix_cache_strategy
in ("extra_buffer", "extra_buffer_lazy")
)
def mamba_extra_buffer_lazy_enabled() -> bool:
"""The lazy variant of :func:`mamba_extra_buffer_enabled`."""
return (
get_memory().disable_radix_cache is False
and get_exec().mamba.mamba_radix_cache_strategy == "extra_buffer_lazy"
)
# --- Derived config accessors ------------------------------------------------
#
# A few values are computed from several config fields plus the HF config, so
# they are ``ServerArgs`` members rather than namespace leaves. Business code
# must not reach for the startup record to get them: these accessors are the
# named home, and this module — which owns the slot — is the only place that
# reads it. Each one keeps the member's exact semantics, including which model
# config it derives from (always the process's, i.e. the target's).
def mamba_cache_chunk_size() -> int:
"""The caching point granularity for mamba state: ``max(the model's mamba
chunk size, page_size)``. Cached on the config after the first call."""
return get_server_args().mamba_cache_chunk_size
def max_speculative_num_draft_tokens() -> int | None:
"""The largest draft-token count speculative decoding may use.
All three inputs are ``spec`` leaves, so this derives from the bags and
follows a post-publish override; ``ServerArgs.max_speculative_num_draft_tokens``
is the pre-publish equivalent. Adaptive spec resolves the count from its
candidate-step table instead of the flat field.
"""
spec = get_spec()
if spec.speculative_num_draft_tokens is None:
return None
if not spec.speculative_adaptive:
return spec.speculative_num_draft_tokens
# The adaptive branch parses a JSON config, and this is called per decode
# batch (`spec_prepare_for_decode`), so memoize on the inputs -- keyed, not
# cached once, so a post-publish override still recomputes.
return _adaptive_draft_token_bound(spec.speculative_adaptive_config)
@functools.lru_cache(maxsize=8)
def _adaptive_draft_token_bound(cfg_path: str | None) -> int:
from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config,
)
candidate_steps = resolve_candidate_steps_from_config(cfg_path=cfg_path)
# Adaptive spec requires topk=1 today, so each runtime state needs
# steps + 1 draft-token slots (mirrors the ServerArgs member).
return max(candidate_steps) + 1
def uses_mla_backend() -> bool:
"""Whether this process's model runs the MLA attention path."""
return get_server_args().use_mla_backend()
def attention_backends() -> tuple:
"""The configured ``(prefill, decode)`` backend pair, split fields falling
back to ``attention_backend``.
All three inputs are ``exec.kernel`` leaves, so this derives from the bags
and follows a post-publish override; ``ServerArgs.get_attention_backends``
is the pre-publish equivalent the resolution pipeline uses. A built runner
stamps its own resolved pair (``ModelRunner.prefill_attention_backend_str``);
read that when there is a runner in hand.
"""
from sglang.srt.arg_groups.overrides import attention_backends_of
# All three leaves live in the same bag, so the resolution pipeline's own
# helper applies directly -- one definition of the fallback rule.
return attention_backends_of(get_exec().kernel)
def process_model_config():
"""The process's ``ModelConfig`` (built once from the published config)."""
return get_server_args().get_model_config()
def cutedsl_moe_max_num_tokens() -> int:
"""The CuteDSL A2A per-rank token budget.
Every input is a published leaf (``spec``, ``schedule``, ``exec.graph``), so
this derives from the bags and follows a post-publish override;
``ServerArgs.cutedsl_moe_max_num_tokens`` is the pre-publish equivalent the
resolution pipeline uses. Max over the prefill bound, the piecewise-prefill
capture, and the decode/verify bound.
"""
from sglang.srt.model_executor.cuda_graph_config import Backend
spec = get_spec()
num_tokens_per_req = (
(spec.speculative_num_draft_tokens or 1) if spec.speculative_algorithm else 1
)
prefill_tokens = get_schedule().max_prefill_tokens
cg_config = get_exec().graph.cuda_graph_config
if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE:
prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0)
decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0
return max(prefill_tokens, decode_max_bs * num_tokens_per_req)
# --- Configured (not live) parallel sizes ------------------------------------
#
# ``get_parallel()`` shadows these names with the LIVE topology, which is the
# right answer almost everywhere. A handful of call sites need what was
# *configured* instead — before the groups exist, in a process that has none,
# or where the live value is deliberately aliased to another dimension. Each
# accessor below names that intent so no business call site has to reach for
# the startup record; the per-site reasons live in the read ratchet.
#
# They read the published leaf rather than the record: the bag is what
# ``override`` writes, and once the instance holds only the user's raw input
# the record would answer with what was *typed* instead of what resolution
# produced. Going through the bag directly is what gets past the live property
# that shadows these four names on ``get_parallel()``.
def _configured_parallel(name: str):
# The bag itself, not ParallelContext, whose live property shadows these
# four names. Read through the parallel slot the way the leaf accessor
# does — ``parallel`` is deliberately outside the per-role namespace table
# (every process reads topology config), so this must not route through
# ``config_bag()``'s role check, which would record or reject the read.
config = _CONTEXT.parallel._config
if config is None:
raise ValueError("config namespace 'parallel' not published")
return getattr(config, name)
def configured_tp_size() -> int:
return _configured_parallel("tp_size")
def configured_pp_size() -> int:
return _configured_parallel("pp_size")
def configured_moe_dp_size() -> int:
return _configured_parallel("moe_dp_size")
def configured_attn_cp_size() -> int:
return _configured_parallel("attn_cp_size")
def is_ep_joiner() -> bool:
"""True in a process launched as an elastic-EP joiner (scale or recover).
A predicate over the published ``exec.moe.ep_join_mode`` leaf, so it follows
a post-publish override; the same-named ``ServerArgs`` property is the
pre-publish equivalent.
"""
return get_exec().moe.ep_join_mode in ("scale", "recover")
def is_ep_scale_joiner() -> bool:
"""True in a process launched as an elastic-EP scale-up joiner."""
return get_exec().moe.ep_join_mode == "scale"