config: retire the alias-form process-global config reads

`sa = get_server_args()` followed by `sa.field` reads the same startup record as
the direct form; the read ratchet added in the previous slice pinned twelve of
them as the remaining surface. Eleven now read the accessor for what they
actually want:

- `is_enable_moe_cp_allgather` compares the attention-CP and MoE-DP sizes to
  decide whether a forward needs an allgather, so it reads the live topology
  through `get_parallel()` — the same source `get_moe_cp_size()` right above it
  already uses. Both groups exist once model-parallel init has run, which is
  before any forward.
- The DeepSeek MLA decode-backend gate and Inkling's attention paths read
  `get_exec().kernel`; Inkling's KV-dtype checks read `get_model()`. These are
  per-runner fields, and the value they get is the config published for the
  runner being built — unchanged from what the alias returned.
- The int8 mamba checkpoint pool reads `get_exec().mamba`. It keeps its guard
  for callers that construct the pool with no published config; that guard now
  catches the namespace accessor instead of the slot.

`model_loader`'s `moe_dp_size` stays on the instance and is exempt: the dict it
belongs to already reports the live size under `"dp"`, so that entry is the
configured intent, and `get_parallel()` shadows the name with the live value.

Alias-form baseline 12 -> 0. What remains on `get_server_args()` in the package
is the derived API (properties and methods computed from several fields plus the
HF config) and four config-intent reads of live-shadowed sizes, each exempt by
name with its reason.
This commit is contained in:
Cheng Wan
2026-08-05 19:29:20 -07:00
committed by GitHub
parent 1d47952c7c
commit bebebb8f6c
7 changed files with 49 additions and 34 deletions
+9 -5
View File
@@ -27,7 +27,7 @@ from sglang.srt.distributed import (
from sglang.srt.distributed.device_communicators.pynccl_allocator import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
from sglang.srt.runtime_context import get_flags from sglang.srt.runtime_context import get_flags, get_server_args
from sglang.srt.utils import get_bool_env_var, is_hip from sglang.srt.utils import get_bool_env_var, is_hip
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -979,11 +979,15 @@ def get_moe_cp_size() -> int:
def is_enable_moe_cp_allgather() -> bool: def is_enable_moe_cp_allgather() -> bool:
"""True when moe_dp_size < attn_cp_size, requiring allgather across CP ranks before MoE.""" """True when moe_dp_size < attn_cp_size, requiring allgather across CP ranks before MoE.
from sglang.srt.runtime_context import get_server_args
sa = get_server_args() Reads the configured sizes, not the live groups: that very configuration makes
return sa.attn_cp_size > sa.moe_dp_size ``initialize_model_parallel`` alias ``_MOE_DP`` to ``_ATTN_CP``
(``parallel_state.py``), so the live sizes are equal and the comparison would
always be false.
"""
server_args = get_server_args()
return server_args.attn_cp_size > server_args.moe_dp_size
def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor): def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
+2 -3
View File
@@ -64,11 +64,10 @@ def _init_state() -> Optional[_State]:
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
CustomAllReduceV2, CustomAllReduceV2,
) )
from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.runtime_context import get_exec, get_parallel
from sglang.srt.utils.common import get_device_sm from sglang.srt.utils.common import get_device_sm
server_args = get_server_args() a2a = get_exec().moe.moe_a2a_backend
a2a = server_args.moe_a2a_backend
group = get_parallel().attn_tp_group group = get_parallel().attn_tp_group
comm = group.ca_comm comm = group.ca_comm
if ( if (
@@ -313,21 +313,21 @@ def maybe_init_int8_mamba_checkpoint_pool(
allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable allocating, so an oversized ``--int8-mamba-ckpt-size`` fails with an actionable
message instead of a cryptic mid-allocation CUDA OOM. message instead of a cryptic mid-allocation CUDA OOM.
""" """
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_exec
try: try:
_sa = get_server_args() mamba = get_exec().mamba
except ValueError: except ValueError:
# Some unit-test / mock runners construct HybridReqToTokenPool directly # Some unit-test / mock runners construct HybridReqToTokenPool directly
# without a global server-args context. The int8 checkpoint pool is opt-in # without a global server-args context. The int8 checkpoint pool is opt-in
# via a CLI flag, so an unset context unambiguously means it is off. # via a CLI flag, so an unset context unambiguously means it is off.
_sa = None mamba = None
if not getattr(_sa, "enable_int8_mamba_checkpoint", False): if mamba is None or not mamba.enable_int8_mamba_checkpoint:
return None return None
GB = 1 << 30 GB = 1 << 30
H, d_v, d_k = cache_params.shape.temporal H, d_v, d_k = cache_params.shape.temporal
ckpt_size = _sa.int8_mamba_ckpt_size or (2 * mamba_size) ckpt_size = mamba.int8_mamba_ckpt_size or (2 * mamba_size)
kwargs = dict( kwargs = dict(
num_layers=len(mamba_layer_ids), num_layers=len(mamba_layer_ids),
num_slots=ckpt_size, num_slots=ckpt_size,
+12 -11
View File
@@ -29,7 +29,11 @@ from sglang.srt.models.inkling_common.kernels.comm import (
from sglang.srt.models.inkling_common.norm import RMSNorm from sglang.srt.models.inkling_common.norm import RMSNorm
from sglang.srt.models.inkling_common.sconv import SconvType, ShortConvolution from sglang.srt.models.inkling_common.sconv import SconvType, ShortConvolution
from sglang.srt.models.utils import apply_qk_norm from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.runtime_context import (
get_exec,
get_model,
get_parallel,
)
from sglang.srt.utils import add_prefix, get_current_device_stream_fast from sglang.srt.utils import add_prefix, get_current_device_stream_fast
try: try:
@@ -410,8 +414,7 @@ class InklingAttention(nn.Module):
) )
sfk = sfv = None sfk = sfv = None
do_mxfp8_store = False do_mxfp8_store = False
server_args = get_server_args() if get_model().kv_cache_dtype == "mxfp8" and hasattr(
if server_args.kv_cache_dtype == "mxfp8" and hasattr(
pool, "get_kv_scale_buffer" pool, "get_kv_scale_buffer"
): ):
sfk, sfv = pool.get_kv_scale_buffer(self.layer_id) sfk, sfv = pool.get_kv_scale_buffer(self.layer_id)
@@ -525,8 +528,7 @@ class InklingAttention(nn.Module):
) )
sfk = sfv = None sfk = sfv = None
do_mxfp8_store = False do_mxfp8_store = False
server_args = get_server_args() if get_model().kv_cache_dtype == "mxfp8" and hasattr(
if server_args.kv_cache_dtype == "mxfp8" and hasattr(
pool, "get_kv_scale_buffer" pool, "get_kv_scale_buffer"
): ):
sfk, sfv = pool.get_kv_scale_buffer(self.layer_id) sfk, sfv = pool.get_kv_scale_buffer(self.layer_id)
@@ -645,8 +647,7 @@ class InklingAttention(nn.Module):
) )
sfk = sfv = None sfk = sfv = None
do_mxfp8_store = False do_mxfp8_store = False
server_args = get_server_args() if get_model().kv_cache_dtype == "mxfp8" and hasattr(
if server_args.kv_cache_dtype == "mxfp8" and hasattr(
pool, "get_kv_scale_buffer" pool, "get_kv_scale_buffer"
): ):
sfk, sfv = pool.get_kv_scale_buffer(self.layer_id) sfk, sfv = pool.get_kv_scale_buffer(self.layer_id)
@@ -726,12 +727,12 @@ class InklingAttention(nn.Module):
apply_log_scaling = log_scaling_tau is not None and not self.is_local apply_log_scaling = log_scaling_tau is not None and not self.is_local
server_args = get_server_args() attention_backend = get_exec().kernel.attention_backend
assert server_args.attention_backend in ("fa4", "triton") assert attention_backend in ("fa4", "triton")
# The overlap threads a CUDA event into the FA4 sheared-bias kernel, so it # The overlap threads a CUDA event into the FA4 sheared-bias kernel, so it
# is FA4-only for now. # is FA4-only for now.
# TODO(triton): plumb rel_bias_event through the triton attn path too. # TODO(triton): plumb rel_bias_event through the triton attn path too.
fa4 = server_args.attention_backend == "fa4" fa4 = attention_backend == "fa4"
rel_event = None rel_event = None
prologue_did_store = False prologue_did_store = False
@@ -870,7 +871,7 @@ class InklingAttention(nn.Module):
) )
extra_attn_kwargs = {} extra_attn_kwargs = {}
if server_args.kv_cache_dtype == "mxfp8": if get_model().kv_cache_dtype == "mxfp8":
# Must run AFTER v is joined above (wait_event(v_event)): v (and k) # Must run AFTER v is joined above (wait_event(v_event)): v (and k)
# may be produced by sconv on the alt stream, and quantizing them on # may be produced by sconv on the alt stream, and quantizing them on
# the main stream before the join reads half-written buffers under # the main stream before the join reads half-written buffers under
+3 -3
View File
@@ -34,7 +34,7 @@ from sglang.srt.layers.attention.vision import (
) )
from sglang.srt.models.kimi_vl_moonvit import tpool_patch_merger from sglang.srt.models.kimi_vl_moonvit import tpool_patch_merger
from sglang.srt.multimodal.mm_utils import concat_or_single from sglang.srt.multimodal.mm_utils import concat_or_single
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_mm
from sglang.srt.utils import get_bool_env_var, is_hip, print_info_once from sglang.srt.utils import get_bool_env_var, is_hip, print_info_once
_is_hip = is_hip() _is_hip = is_hip()
@@ -59,10 +59,10 @@ def _resolve_grid_thw_list(
def _get_mm_attention_backend() -> str: def _get_mm_attention_backend() -> str:
try: try:
server_args = get_server_args() return get_mm().mm_attention_backend or "auto"
except ValueError: except ValueError:
# config not published yet (import-time probes)
return "auto" return "auto"
return server_args.mm_attention_backend or "auto"
def _is_fa4_available() -> bool: def _is_fa4_available() -> bool:
@@ -21,7 +21,7 @@ from sglang.srt.multimodal.kimi_k3_vit_cuda_graph_runner import (
KimiK3ViTCudaGraphRunner, KimiK3ViTCudaGraphRunner,
) )
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_context, get_parallel
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu") register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -357,7 +357,7 @@ def test_kimi_k3_position_interpolation_uses_contiguous_chw(monkeypatch):
assert torch.equal(actual, expected) assert torch.equal(actual, expected)
def test_kimi_k3_prepares_shared_attention_metadata_once(monkeypatch): def test_kimi_k3_prepares_shared_attention_metadata_once(monkeypatch, request):
metadata_ids = [] metadata_ids = []
values_are_contiguous = [] values_are_contiguous = []
@@ -370,11 +370,13 @@ def test_kimi_k3_prepares_shared_attention_metadata_once(monkeypatch):
values_are_contiguous.append(v.is_contiguous()) values_are_contiguous.append(v.is_contiguous())
return q return q
monkeypatch.setattr( # Force the backend through the context, not by patching an import binding:
kimi_k3_vl, # production reads the published config bag (get_mm().mm_attention_backend).
"get_server_args", override = get_context().override_server_args(
lambda: SimpleNamespace(mm_attention_backend="flashinfer_cudnn"), mm_attention_backend="flashinfer_cudnn"
) )
override.install()
request.addfinalizer(override.restore)
monkeypatch.setitem(kimi_k3_vl.QKV_BACKEND_IMPL, "flashinfer_cudnn", FakeAttention) monkeypatch.setitem(kimi_k3_vl.QKV_BACKEND_IMPL, "flashinfer_cudnn", FakeAttention)
encoder = MoonViT3dEncoder( encoder = MoonViT3dEncoder(
@@ -31,6 +31,12 @@ What legitimately remains:
reads ``get_dcp_group()``, and that group is only installed when DCP is on. reads ``get_dcp_group()``, and that group is only installed when DCP is on.
- ``cuda_ipc_transport_utils.tp_size`` runs in the tokenizer process, which has - ``cuda_ipc_transport_utils.tp_size`` runs in the tokenizer process, which has
no groups at all (the call site already guards for "not published yet"). no groups at all (the call site already guards for "not published yet").
- ``dp_attention.attn_cp_size`` / ``moe_dp_size``: the configuration the
predicate detects (``attn_cp_size > moe_dp_size``) is the one where
``initialize_model_parallel`` aliases ``_MOE_DP`` to ``_ATTN_CP``, so the live
sizes are equal there and a live comparison is always false.
- ``model_loader/loader.py`` reports both: the same dict carries the live
``moe_dp_size`` under ``"dp"``, so this entry is the configured intent.
- The alias-form baseline is not zero yet. Lowering it is the next slice; the - The alias-form baseline is not zero yet. Lowering it is the next slice; the
failure message lists the sites whenever the count moves. failure message lists the sites whenever the count moves.
""" """
@@ -67,13 +73,16 @@ _DERIVED_MEMBERS = frozenset(
_CONFIG_INTENT_SIZES = frozenset( _CONFIG_INTENT_SIZES = frozenset(
{ {
("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"), ("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"),
("srt/layers/dp_attention.py", "attn_cp_size"),
("srt/layers/dp_attention.py", "moe_dp_size"),
("srt/mem_cache/allocation.py", "dcp_size"), ("srt/mem_cache/allocation.py", "dcp_size"),
("srt/model_loader/loader.py", "moe_dp_size"),
("srt/utils/cuda_ipc_transport_utils.py", "tp_size"), ("srt/utils/cuda_ipc_transport_utils.py", "tp_size"),
} }
) )
_DIRECT_BASELINE = 0 _DIRECT_BASELINE = 0
_ALIAS_BASELINE = 12 _ALIAS_BASELINE = 0
def _is_global_call(node) -> bool: def _is_global_call(node) -> bool: