[refactor] Read resolved config from server_args fields; retire the flags mirror tier (#30346)

This commit is contained in:
Cheng Wan
2026-07-07 21:28:34 -07:00
committed by GitHub
parent b14f7b4f75
commit be32c57598
60 changed files with 238 additions and 795 deletions
+2 -2
View File
@@ -77,8 +77,8 @@ class Arg:
no_cli: bool = False
# When True, this field may be written by config resolution (model
# overrides and post-process passes): it is part of the whitelist accepted
# by the apply_model_overrides gate, and its resolved value lives on the
# flags tier (the server_args field itself stays the pristine user input).
# by the declaration stash, and its resolved value materializes onto the
# field at the end of __post_init__.
resolvable: bool = False
+15 -112
View File
@@ -14,8 +14,9 @@
"""Declarative model-override registry.
Model-identity adjustments to the server configuration are DECLARED here and
resolved into the flags tier through the ``apply_model_overrides`` gate —
model code never mutates ``ServerArgs``, which stays the pristine user input.
materialized onto ``server_args`` at the end of ``__post_init__`` (gate
order, last writer wins) — model code never mutates ``ServerArgs`` fields
imperatively.
Two declaration forms, keyed on ``hf_config.architectures[0]``:
@@ -30,11 +31,10 @@ from __future__ import annotations
import dataclasses
import logging
from typing import Any, Callable, Dict, Iterable, List, Optional, Sequence, Tuple
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
from sglang.srt.arg_groups.arg_utils import resolvable_fields
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.runtime_context import resolve_flag_leaf
from sglang.srt.utils.common import (
cpu_has_amx_support,
get_device_capability,
@@ -259,16 +259,19 @@ def mamba_extra_buffer_of(cfg: Any) -> bool:
def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
"""Declare a load-time resolved field (model-file config overrides,
weight-resolved dtypes): apply it onto the published ``server_args`` —
resolution has already materialized, so post-init declarations write
through — and record it into the flags tier through the runtime gate."""
weight-resolved dtypes) on the published ``server_args``: resolution has
already materialized, so the declaration writes through, joining the
declaration stash for provenance and republish consistency."""
from sglang.srt.runtime_context import get_context
ctx = get_context()
entry = (source, dict(declared))
validate_declarations(ctx.server_args, [entry])
_apply_fields(ctx.server_args, declared)
ctx.record_runtime_overrides([entry])
server_args = get_context().server_args
validate_declarations(server_args, [(source, dict(declared))])
override = getattr(server_args, "override", None)
if override is not None:
override(source, **declared)
else:
# Config-shaped fixtures without the mutation entry point.
_apply_fields(server_args, declared)
def collect_model_override_declarations(
@@ -2054,84 +2057,6 @@ def _dllm_page_size(view: Any) -> dict:
return {}
@dataclasses.dataclass(frozen=True)
class OverrideRecord:
"""Provenance of one resolved write: ``base`` is the value before this
declaration applied (the pristine value for the first writer)."""
source: str
field: str
base: Any
resolved: Any
def apply_model_overrides(
flags: Any,
server_args: Any,
declarations: Sequence[Tuple[str, Dict[str, Any]]],
*,
terminal: Sequence[Tuple[str, Dict[str, Any]]] = (),
whitelist: Optional[Iterable[str]] = None,
leaf_map: Optional[Dict[str, str]] = None,
) -> List[OverrideRecord]:
"""Resolve model-override declarations into the flags tier.
- **Transactional**: every declaration (``terminal`` included) is
validated against the whitelist and the flag-leaf layout BEFORE any
write; on error nothing is applied.
- **Ordering**: ``declarations`` apply in order (last writer wins), then
``terminal`` (the enforce-disable pass) applies after everything.
- **Materialization**: every whitelisted field becomes a flag leaf —
declared fields carry the resolved value, undeclared ones the pristine
``server_args`` value — so readers only ever read flags, never a
"flag or fallback to config" combination.
- ``server_args`` is read-only here: resolution output lives on flags.
Returns the provenance log, one record per declared write.
"""
if whitelist is None:
whitelist = resolvable_fields(type(server_args))
whitelist = frozenset(whitelist)
ordered = list(declarations) + list(terminal)
problems = [
f"{source}: {sorted(set(decl) - whitelist)} not model-overridable"
for source, decl in ordered
if set(decl) - whitelist
]
if problems:
raise ValueError(
"model override validation failed (nothing was applied): "
+ "; ".join(problems)
)
for field in sorted(whitelist):
owner, leaf = resolve_flag_leaf(flags, field, leaf_map=leaf_map)
if leaf not in type(owner).__dataclass_fields__:
raise ValueError(
f"flag leaf for '{field}' is not declared on "
f"{type(owner).__name__} (declare the dataclass field and map "
"it in FLAG_LEAF_MAP); nothing was applied"
)
if getattr(owner, "_frozen", False):
raise RuntimeError(
f"cannot resolve '{field}': {type(owner).__name__} is frozen; "
"nothing was applied"
)
resolved = {field: getattr(server_args, field) for field in whitelist}
records: List[OverrideRecord] = []
for source, decl in ordered:
for field, value in decl.items():
records.append(OverrideRecord(source, field, resolved[field], value))
resolved[field] = value
for field, value in resolved.items():
owner, leaf = resolve_flag_leaf(flags, field, leaf_map=leaf_map)
setattr(owner, leaf, value)
return records
def validate_declarations(
server_args: Any,
declarations: Sequence[Tuple[str, Dict[str, Any]]],
@@ -2168,25 +2093,3 @@ def _hrm_text_attention_force(view: Any) -> dict:
"attention."
)
return {"attention_backend": "triton"}
def assert_flag_parity(
flags: Any,
server_args: Any,
fields: Iterable[str],
*,
leaf_map: Optional[Dict[str, str]] = None,
) -> None:
"""Drift guard: each declared field's flag leaf must equal the
(materialized) ``server_args`` value."""
mismatches = []
for field in fields:
owner, leaf = resolve_flag_leaf(flags, field, leaf_map=leaf_map)
flag_value = getattr(owner, leaf)
args_value = getattr(server_args, field)
if flag_value != args_value:
mismatches.append(
f"{field}: flags={flag_value!r} server_args={args_value!r}"
)
if mismatches:
raise AssertionError("flag/server_args parity broken: " + "; ".join(mismatches))
@@ -39,7 +39,7 @@ from sglang.srt.model_executor.forward_batch_info import (
compute_position,
)
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
@@ -634,7 +634,7 @@ class TboForwardBatchPreparer:
sum_field=None,
)
_, child_b.extend_start_loc = compute_position(
get_flags().attn.backend,
get_server_args().attention_backend,
child_b.extend_prefix_lens,
child_b.extend_seq_lens,
child_b.extend_num_tokens,
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
ForwardMode,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils.common import (
is_cpu,
@@ -335,7 +335,7 @@ class LogitsProcessor(nn.Module):
self.config = config
self.vocab_size = config.vocab_size
self.logit_scale = logit_scale
self.use_attn_tp_group = get_flags().enable_dp_lm_head
self.use_attn_tp_group = get_server_args().enable_dp_lm_head
self.use_fp32_lm_head = get_global_server_args().enable_fp32_lm_head
if self.use_attn_tp_group:
self.attn_tp_size = get_parallel().attn_tp_size
@@ -18,7 +18,6 @@ from sglang.srt.layers.rotary_embedding.yarn import (
yarn_get_mscale_simple,
yarn_linear_ramp_mask,
)
from sglang.srt.runtime_context import get_flags
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
cpu_has_amx_support,
@@ -45,6 +44,8 @@ if _is_xpu:
import triton
import triton.language as tl
from sglang.srt.runtime_context import get_server_args
@triton.jit
def apply_interleaved_rope_kernel(
@@ -227,7 +228,7 @@ class MRotaryEmbedding(RotaryEmbedding):
last_dim = cos_sin.size()[-1]
cos, sin = cos_sin.chunk(2, dim=-1)
if self.mrope_interleaved:
if support_triton(get_flags().attn.backend):
if support_triton(get_server_args().attention_backend):
cos = apply_interleaved_rope_triton(cos, self.mrope_section)
sin = apply_interleaved_rope_triton(sin, self.mrope_section)
else:
+3 -3
View File
@@ -13,7 +13,7 @@ from sglang.srt.layers.dp_attention import (
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.utils.hash import murmur_hash32
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
from sglang.srt.runtime_context import get_flags
from sglang.srt.runtime_context import get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import get_global_server_args
@@ -80,7 +80,7 @@ class Sampler(nn.Module):
)
# In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer.
self.use_log_softmax_logprob = self.rl_on_policy_target is not None
self.use_ascend_backend = get_flags().sampling_backend == "ascend"
self.use_ascend_backend = get_server_args().sampling_backend == "ascend"
def _preprocess_logits(
self, logits: torch.Tensor, sampling_info: SamplingBatchInfo
@@ -231,7 +231,7 @@ class Sampler(nn.Module):
positions=positions,
)
else:
backend = get_flags().sampling_backend
backend = get_server_args().sampling_backend
if backend == "flashinfer":
assert (
sampling_info.sampling_seed is None
-8
View File
@@ -554,14 +554,6 @@ class Scheduler(
self.init_batch_result_processor()
# The config-resolution lifecycle of this scheduler process ends
# here: every load-time stage has run (target and draft model init,
# weight-resolved kv-cache dtype), so lock the static flag groups.
# flags.capture stays writable; late resolution writes now raise.
from sglang.srt.runtime_context import get_context
get_context().freeze_flags()
self.is_initializing = False
def init_zbal_on_npu(self):
+3 -3
View File
@@ -25,7 +25,7 @@ from sglang.srt.mem_cache.triton_ops.common import (
get_last_loc_triton_safe,
write_req_to_token_pool_triton,
)
from sglang.srt.runtime_context import get_flags
from sglang.srt.runtime_context import get_server_args
from sglang.srt.server_args import ServerArgs, get_global_server_args
from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
@@ -134,7 +134,7 @@ def write_cache_indices(
prefix_tensors: list[torch.Tensor],
req_to_token_pool: ReqToTokenPool,
):
if support_triton(get_flags().attn.backend):
if support_triton(get_server_args().attention_backend):
prefix_pointers = torch.tensor(
[t.data_ptr() for t in prefix_tensors],
dtype=torch.uint64,
@@ -175,7 +175,7 @@ def get_last_loc(
req_pool_indices_tensor: torch.Tensor,
prefix_lens_tensor: torch.Tensor,
) -> torch.Tensor:
attn_backend = get_flags().attn.backend
attn_backend = get_server_args().attention_backend
uses_triton_dispatch = attn_backend not in ("ascend", "torch_native")
if _is_hip and uses_triton_dispatch:
@@ -182,7 +182,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
from sglang.srt.model_loader.utils import set_default_torch_dtype
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import get_flags
from sglang.srt.runtime_context import get_flags, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ( # noqa: F401 (re-export)
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
@@ -541,7 +541,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.init_threads_binding()
# Set float32 matmul precision
if get_flags().enable_tf32_matmul:
if get_server_args().enable_tf32_matmul:
torch.set_float32_matmul_precision("high")
# Get available memory before model loading.
+2 -2
View File
@@ -52,7 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader,
maybe_remap_kv_scale_name,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__)
@@ -446,7 +446,7 @@ class ApertusForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
+2 -2
View File
@@ -46,7 +46,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader,
maybe_remap_kv_scale_name,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__)
@@ -405,7 +405,7 @@ class ArceeForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
+2 -2
View File
@@ -77,7 +77,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
@@ -832,7 +832,7 @@ class BailingMoEForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
@@ -58,7 +58,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
BumpAllocator,
@@ -1089,7 +1089,7 @@ class BailingMoELinearForCausalLM(nn.Module):
config.hidden_size,
params_dtype=torch.float32,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
)
self.logits_processor = LogitsProcessor(config)
@@ -42,7 +42,7 @@ from sglang.srt.models.bailing_moe_linear import (
BailingMoeV2_5ForCausalLM,
)
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import BumpAllocator, add_prefix
LoraConfig = None
@@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid":
+3 -3
View File
@@ -57,7 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
@@ -172,7 +172,7 @@ class DeepseekModelNextN(nn.Module):
if (
_is_npu
and self.quant_config is None
and get_flags().quantization is not None
and get_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
@@ -330,7 +330,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+4 -4
View File
@@ -176,7 +176,7 @@ from sglang.srt.models.deepseek_common.utils import (
_use_aiter_bpreshuffle_gfx95,
_use_aiter_gfx95,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import (
@@ -541,7 +541,7 @@ class DeepseekV2MoE(nn.Module):
n_shared_experts = (
0 if config.n_shared_experts is None else int(config.n_shared_experts)
)
_fusion_disabled = get_flags().disable_shared_experts_fusion
_fusion_disabled = get_server_args().disable_shared_experts_fusion
# num_fused_shared_experts drives weight remapping in deepseek_weight_loader:
# mlp.shared_experts → mlp.experts.256 when > 0.
@@ -2703,7 +2703,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
else:
# ranks other than the last rank will have a placeholder layer
@@ -2742,7 +2742,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self.num_fused_shared_experts = 0
server_args = get_global_server_args()
if get_flags().disable_shared_experts_fusion:
if get_server_args().disable_shared_experts_fusion:
return
disable_reason = None
+3 -3
View File
@@ -125,7 +125,7 @@ from sglang.srt.models.deepseek_v2 import (
_is_npu,
_is_xpu,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
if not _is_hip:
from sglang.srt.layers.utils.cp_utils import (
@@ -2171,7 +2171,7 @@ class DeepseekV4ForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
else:
self.lm_head = PPMissingLayer()
@@ -2209,7 +2209,7 @@ class DeepseekV4ForCausalLM(nn.Module):
def determine_num_fused_shared_experts(self):
self.num_fused_shared_experts = 0
if get_flags().disable_shared_experts_fusion:
if get_server_args().disable_shared_experts_fusion:
return
disable_reason = None
@@ -38,7 +38,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -233,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -31,7 +31,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader,
maybe_remap_kv_scale_name,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, make_layers
from sglang.utils import get_exception_traceback, logger
@@ -443,7 +443,7 @@ class Exaone4ForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -62,7 +62,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -652,7 +652,7 @@ class ExaoneMoEForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
# For EAGLE3 support
+2 -2
View File
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -63,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers
logger = logging.getLogger(__name__)
@@ -477,7 +477,7 @@ class FalconH1ForCausalLM(nn.Module):
quant_config=quant_config,
org_num_embeddings=config.vocab_size,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.lm_head = self.lm_head.float()
self.lm_head_multiplier = config.lm_head_multiplier
+6 -4
View File
@@ -82,7 +82,7 @@ from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -404,7 +404,9 @@ class Glm4MoeSparseMoeBlock(nn.Module):
self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = (
0 if get_flags().disable_shared_experts_fusion else config.n_shared_experts
0
if get_server_args().disable_shared_experts_fusion
else config.n_shared_experts
)
self.config = config
@@ -1184,7 +1186,7 @@ class Glm4MoeForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
@@ -1192,7 +1194,7 @@ class Glm4MoeForCausalLM(nn.Module):
self.capture_aux_hidden_states = False
def determine_num_fused_shared_experts(self):
if get_flags().disable_shared_experts_fusion:
if get_server_args().disable_shared_experts_fusion:
return
disable_reason = None
+6 -4
View File
@@ -74,7 +74,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
)
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
BumpAllocator,
@@ -188,7 +188,9 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts
self.num_fused_shared_experts = (
0 if get_flags().disable_shared_experts_fusion else config.n_shared_experts
0
if get_server_args().disable_shared_experts_fusion
else config.n_shared_experts
)
self.config = config
self.layer_id = layer_id
@@ -918,7 +920,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
@@ -939,7 +941,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self, architecture: str = "Glm4MoeLiteForCausalLM"
):
self.num_fused_shared_experts = 0
if get_flags().disable_shared_experts_fusion:
if get_server_args().disable_shared_experts_fusion:
return
disable_reason = None
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
Glm4MoeLiteDecoderLayer,
Glm4MoeLiteForCausalLM,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
@@ -155,12 +155,12 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = (
0 if get_flags().disable_shared_experts_fusion else 1
0 if get_server_args().disable_shared_experts_fusion else 1
)
@torch.no_grad()
+3 -3
View File
@@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu
@@ -141,12 +141,12 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = (
0 if get_flags().disable_shared_experts_fusion else 1
0 if get_server_args().disable_shared_experts_fusion else 1
)
@torch.no_grad()
+3 -3
View File
@@ -18,7 +18,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4_moe import Glm4MoeModel
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -70,7 +70,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
else:
# ranks other than the last rank will have a placeholder layer
@@ -84,7 +84,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
self.capture_aux_hidden_states = False
def determine_num_fused_shared_experts(self):
if get_flags().disable_shared_experts_fusion:
if get_server_args().disable_shared_experts_fusion:
return
disable_reason = None
+3 -3
View File
@@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4 import Glm4DecoderLayer
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -134,12 +134,12 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = (
0 if get_flags().disable_shared_experts_fusion else 1
0 if get_server_args().disable_shared_experts_fusion else 1
)
@torch.no_grad()
+3 -3
View File
@@ -68,7 +68,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
LazyValue,
@@ -390,7 +390,7 @@ class GptOssAttention(nn.Module):
# Choose dtype of sinks based on attention backend: trtllm_mha requires float32,
# others can use bfloat16
attn_backend = get_flags().attn.backend
attn_backend = get_server_args().attention_backend
sinks_dtype = torch.float32 if attn_backend == "trtllm_mha" else torch.bfloat16
self.sinks = nn.Parameter(
torch.empty(self.num_heads, dtype=sinks_dtype), requires_grad=False
@@ -745,7 +745,7 @@ class GptOssForCausalLM(nn.Module):
config.hidden_size,
# quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False
+2 -2
View File
@@ -53,7 +53,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, make_layers
@@ -649,7 +649,7 @@ class LagunaForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
else:
self.lm_head = PPMissingLayer()
+2 -2
View File
@@ -76,7 +76,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -796,7 +796,7 @@ class LLaDA2MoeModelLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config, return_full_logits=True)
+2 -2
View File
@@ -52,7 +52,7 @@ from sglang.srt.model_loader.weight_utils import (
kv_cache_scales_loader,
maybe_remap_kv_scale_name,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers
from sglang.utils import get_exception_traceback
@@ -501,7 +501,7 @@ class LlamaForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True)
+2 -2
View File
@@ -86,7 +86,7 @@ from sglang.srt.model_loader.utils import (
)
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import (
BumpAllocator,
add_prefix,
@@ -644,7 +644,7 @@ class LongcatFlashForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False
+2 -2
View File
@@ -76,7 +76,7 @@ from sglang.srt.model_loader.weight_utils import (
)
from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
LazyValue,
@@ -1041,7 +1041,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
else:
self.lm_head = PPMissingLayer()
+2 -2
View File
@@ -44,7 +44,7 @@ from sglang.srt.models.mimo_v2 import (
MiMoV2MLP,
load_mimo_v2_qkv_proj_weight,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
MiMoV2Config = None
@@ -259,7 +259,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -87,7 +87,7 @@ from sglang.srt.models.nemotron_h_utils import (
pad_to_original_num_tokens,
)
from sglang.srt.models.utils import WeightsMapper
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -919,7 +919,7 @@ class NemotronHForCausalLM(nn.Module):
else lora_config.lora_vocab_padding_size
),
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+2 -2
View File
@@ -39,7 +39,7 @@ from sglang.srt.models.nemotron_h import (
NemotronHMoEDecoderLayer,
)
from sglang.srt.models.nemotron_h_utils import is_attn_layer
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
@@ -339,7 +339,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM):
self.config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+3 -3
View File
@@ -93,7 +93,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -148,7 +148,7 @@ def can_fuse_shared_expert(
Caller must still gate on the model/backend support flag.
"""
if (
get_flags().disable_shared_experts_fusion is True
get_server_args().disable_shared_experts_fusion is True
or getattr(config, "shared_expert_intermediate_size", 0) <= 0
or config.shared_expert_intermediate_size != config.moe_intermediate_size
or get_moe_a2a_backend().is_deepep()
@@ -1003,7 +1003,7 @@ class Qwen2MoeForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
# For EAGLE3 support
+2 -2
View File
@@ -33,7 +33,7 @@ from sglang.srt.model_loader.weight_utils import (
from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP
from sglang.srt.models.qwen2 import Qwen2Model
from sglang.srt.models.utils import apply_qk_norm
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu
@@ -493,7 +493,7 @@ class Qwen3ForCausalLM(nn.Module):
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+3 -3
View File
@@ -91,7 +91,7 @@ from sglang.srt.models.utils import (
fused_qk_gemma_rmsnorm,
fused_qk_gemma_rmsnorm_with_gate,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
# Utils
from sglang.srt.utils import (
@@ -133,7 +133,7 @@ cached_get_processor = lru_cache(get_processor)
def _disable_shared_experts_fusion() -> bool:
# Resolved lazily: the global server args is not set at module import time
# (e.g. when this module is imported by unit tests).
return get_flags().disable_shared_experts_fusion
return get_server_args().disable_shared_experts_fusion
if _is_cuda:
@@ -1177,7 +1177,7 @@ class Qwen3_5ForCausalLM(nn.Module):
# so the model still gets the #25885 multi-streaming path. ROCm-only.
if (
config.model_type == "qwen3_5_moe_text"
and not get_flags().disable_shared_experts_fusion
and not get_server_args().disable_shared_experts_fusion
and not can_fuse_shared_expert(config, quant_config)
):
from sglang.srt.arg_groups.overrides import declare_load_time_override
+2 -2
View File
@@ -34,7 +34,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu
@@ -157,7 +157,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
if (
is_npu()
and self.quant_config is None
and get_flags().quantization is not None
and get_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
+2 -2
View File
@@ -72,7 +72,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
LazyValue,
@@ -960,7 +960,7 @@ class Qwen3MoeForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False
+2 -2
View File
@@ -30,7 +30,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__)
@@ -63,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_loader.weight_utils import (
sharded_weight_loader,
)
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.utils import (
LazyValue,
add_prefix,
@@ -1027,7 +1027,7 @@ class Qwen3NextForCausalLM(nn.Module):
quant_config=quant_config,
org_num_embeddings=config.vocab_size,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
# For EAGLE3 support
+3 -3
View File
@@ -32,7 +32,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu
@@ -84,7 +84,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
# Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so
@@ -114,7 +114,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
if (
is_npu()
and self.quant_config is None
and get_flags().quantization is not None
and get_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
+2 -2
View File
@@ -68,7 +68,7 @@ from sglang.srt.models.utils import (
)
from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
@@ -1278,7 +1278,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
self.config.vocab_size,
self.config.hidden_size,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+2 -2
View File
@@ -60,7 +60,7 @@ from sglang.srt.models.bailing_moe import BailingMoEForCausalLM
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
DeepseekMHAForwardMixin,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
BumpAllocator,
@@ -1241,7 +1241,7 @@ class SarvamMLAForCausalLM(nn.Module):
config.hidden_size,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
)
self.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -41,7 +41,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers
@@ -471,7 +471,7 @@ class SDARForCausalLM(nn.Module):
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+2 -2
View File
@@ -57,7 +57,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -566,7 +566,7 @@ class SDARMoeForCausalLM(nn.Module):
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+2 -2
View File
@@ -46,7 +46,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers
@@ -826,7 +826,7 @@ class Step3p5ForCausalLM(nn.Module):
config.vocab_size,
config.hidden_size,
quant_config=quant_config,
use_attn_tp_group=get_flags().enable_dp_lm_head,
use_attn_tp_group=get_server_args().enable_dp_lm_head,
prefix=add_prefix("lm_head", prefix),
)
else:
+20 -229
View File
@@ -28,13 +28,13 @@ tier). The context owns the storage: publishing goes through
``server_args.py`` are thin shims over this slot), and the object is returned
by reference — the same live instance everywhere, never a copy.
``get_flags()`` returns the resolved-flags tier: what the system *resolved*
the configuration to (``server_args`` stays the pristine user input). Flags
live in typed dataclass groups (``flags.attn`` / ``flags.moe`` / flat generic
leaves on ``flags`` itself); reads and writes are plain attribute access.
Static groups are writable during resolution and locked by ``freeze()``;
``flags.capture`` stays writable (capture-time state). Each group offers a
transactional, test-only ``override(**kw)`` that also works on frozen groups.
``get_flags()`` returns the runtime-flags tier. Resolved configuration lives
on ``server_args`` fields (declarations materialize at the end of
``__post_init__``), so this tier only carries genuine runtime state that is
not a function of the configuration alone — today the capture lifecycle
(``flags.capture``). 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
@@ -238,18 +238,13 @@ class _FlagGroupBase:
f"{type(self).__name__} has no flag '{name}' (leaves are "
"declared as dataclass fields; check for typos)"
)
if getattr(self, "_frozen", False):
raise RuntimeError(
f"{type(self).__name__} is frozen; cannot write '{name}'. "
"Test-scoped changes go through override()."
)
object.__setattr__(self, name, value)
@contextmanager
def override(self, **kwargs):
"""Temporarily force flag values, restoring on exit. Transactional
(keys validated before any write) and usable on frozen groups — this
is the test-only injection primitive."""
(keys validated before any write) — the test-only injection
primitive."""
fields = type(self).__dataclass_fields__
unknown = set(kwargs) - set(fields)
if unknown:
@@ -266,37 +261,6 @@ class _FlagGroupBase:
object.__setattr__(self, name, value)
class _StaticFlags(_FlagGroupBase):
"""Static flag-group: writable during resolution, locked by ``freeze()``."""
def freeze(self) -> None:
object.__setattr__(self, "_frozen", True)
@property
def frozen(self) -> bool:
return getattr(self, "_frozen", False)
@dataclasses.dataclass
class AttnFlags(_StaticFlags):
"""Attention-family resolved flags (leaves arrive with the V3 sweeps)."""
# Resolved attention backend; the pristine user request stays on
# server_args.attention_backend.
backend: str | None = None
prefill_backend: str | None = None
decode_backend: str | None = None
@dataclasses.dataclass
class MoeFlags(_StaticFlags):
"""MoE-family resolved flags (leaves arrive with the V3 sweeps)."""
# Resolved MoE runner backend; the pristine user request stays on
# server_args.moe_runner_backend.
runner_backend: str = "auto"
@dataclasses.dataclass
class CaptureFlags(_FlagGroupBase):
"""Capture-time flags; never frozen (written during cuda-graph capture)."""
@@ -307,96 +271,28 @@ class CaptureFlags(_FlagGroupBase):
@dataclasses.dataclass
class Flags(_StaticFlags):
"""Root of the resolved-flags tier.
class Flags(_FlagGroupBase):
"""Root of the runtime-flags tier.
Family groups hang off it (``flags.attn`` / ``flags.moe`` / ``flags.capture``);
single generic flags live flat on this container, declared as fields here.
``freeze()`` locks the container and every static sub-group; ``capture``
stays writable.
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 (today: ``capture``).
"""
attn: AttnFlags = dataclasses.field(default_factory=AttnFlags)
moe: MoeFlags = dataclasses.field(default_factory=MoeFlags)
capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags)
# -- resolved config leaves (flat; materialized at publish) --------------
# Pristine user requests stay on the matching server_args fields; these
# leaves carry the model-resolved values.
dtype: str = "auto"
enable_tf32_matmul: bool = False
enable_multi_layer_eagle: bool = False
swa_full_tokens_ratio: float = 0.8
disable_hybrid_swa_memory: bool = False
sampling_backend: str | None = None
page_size: int | None = None
quantization: str | None = None
fp8_gemm_runner_backend: str = "auto"
disable_overlap_schedule: bool = False
uses_mamba_radix_cache: bool = False
mamba_radix_cache_strategy: str = "auto"
speculative_moe_runner_backend: str | None = None
speculative_moe_a2a_backend: str | None = None
disable_shared_experts_fusion: bool = False
kv_cache_dtype: str = "auto"
dsa_prefill_backend: str | None = None
dsa_decode_backend: str | None = None
flashinfer_allreduce_fusion_backend: str | None = None
# Parallel-request fields: flat transitional home, to be re-homed by the
# Parallel Parameters Clarification module.
enable_dp_attention: bool = False
enable_dp_lm_head: bool = False
moe_a2a_backend: str = "none"
ep_size: int = 1
moe_dense_tp_size: int | None = None
attn_cp_size: int = 1
def freeze(self) -> None:
for field in dataclasses.fields(self):
value = getattr(self, field.name)
if isinstance(value, _StaticFlags):
value.freeze()
super().freeze()
# Resolved-config field name → dotted flag-leaf path (e.g. a V3 sweep adds
# "use_mla_backend": "attn.use_mla_backend"). Fields not listed default to a
# flat leaf of the same name on the Flags container. Populated per field
# family as readers migrate.
FLAG_LEAF_MAP: dict[str, str] = {
"attention_backend": "attn.backend",
"prefill_attention_backend": "attn.prefill_backend",
"decode_attention_backend": "attn.decode_backend",
"moe_runner_backend": "moe.runner_backend",
}
def resolve_flag_leaf(
flags: Flags, field: str, *, leaf_map: dict[str, str] | None = None
) -> tuple[Any, str]:
"""Return ``(owning group, leaf attribute name)`` for a resolved-config field."""
path = (FLAG_LEAF_MAP if leaf_map is None else leaf_map).get(field, field)
owner: Any = flags
*groups, leaf = path.split(".")
for part in groups:
owner = getattr(owner, part)
return owner, leaf
class RuntimeContext:
"""Container for the structured runtime accessors; exposes ``parallel``,
``server_args``, and ``flags``."""
__slots__ = ("parallel", "_server_args", "flags", "_runtime_overrides")
__slots__ = ("parallel", "_server_args", "flags")
def __init__(self, parallel: ParallelContext):
self.parallel = parallel
self._server_args: ServerArgs | None = None
self.flags = Flags()
# Post-publish resolution declarations (runner- and load-time
# resolved fields), replayed after the publish-time stash on
# every re-resolve. Cleared on (re-)publish and reset.
self._runtime_overrides: list[tuple[str, dict]] = []
@property
def server_args(self) -> ServerArgs:
@@ -412,29 +308,10 @@ class RuntimeContext:
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__``).
Publishing also resolves the stashed model-override declarations
into the flags tier (skipped for objects without the stash — dummy /
"none" fixture ServerArgs and test-kit mocks never compute it).
Resolution runs first: if it fails, the previous publish stays intact.
A publish after ``freeze_flags()`` is an ordering violation and raises.
the draft-worker guard in ``ModelRunner.__init__``). The published
object already carries the resolved configuration (declarations
materialize at the end of ``__post_init__``).
"""
if self.flags.frozen:
raise RuntimeError(
"set_server_args() after freeze_flags(): the flags tier is "
"frozen for this process; use reset_context() in tests."
)
# A (re-)publish starts a fresh resolution lifecycle; a failed
# resolve keeps the previous lifecycle (including its recorded
# runtime overrides) intact.
saved_runtime_overrides = self._runtime_overrides
self._runtime_overrides = []
try:
self._resolve_flags(server_args)
except BaseException:
self._runtime_overrides = saved_runtime_overrides
raise
# 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(
@@ -442,91 +319,6 @@ class RuntimeContext:
)
self._server_args = server_args
def record_runtime_overrides(
self, entries: list[tuple[str, dict]]
) -> list[tuple[str, dict]]:
"""Append post-publish resolution declarations (the runner- and
load-time stages) and
atomically re-resolve the flags tier.
Target-worker only, and only before ``freeze_flags()``. During the
dual-apply transition the call sites keep their imperative
``server_args`` writes; the recorded declarations must match them —
parity is re-asserted on every declared field. On failure the
recorded entries are rolled back and the previous flags stay
installed.
"""
server_args = self._server_args
if server_args is None:
raise ValueError("Global server args is not set yet!")
if self.flags.frozen:
raise RuntimeError(
"record_runtime_overrides() after freeze_flags(): runtime "
"resolution stages must complete before the flags tier "
"freezes."
)
entries = [(source, dict(declared)) for source, declared in entries]
self._runtime_overrides.extend(entries)
try:
self._resolve_flags(server_args)
except BaseException:
del self._runtime_overrides[len(self._runtime_overrides) - len(entries) :]
raise
return entries
def freeze_flags(self) -> None:
"""Lock every static flag group (the resolution end point: after the
load-time stages, before serving). ``flags.capture`` stays writable."""
self.flags.freeze()
def _resolve_flags(self, server_args: ServerArgs) -> None:
declarations = getattr(server_args, "_resolved_overrides", None)
if declarations is None and not self._runtime_overrides:
# Stash-less publish. For a config-shaped object (a dataclass:
# mock ServerArgs fixtures, dummy-path instances that skipped the
# monolith) still materialize the whitelist from its own fields,
# so flag reads match legacy server_args reads. Skip only for
# field-less sentinels (tests publishing object()).
if not dataclasses.is_dataclass(server_args):
return
from sglang.srt.arg_groups.arg_utils import resolvable_fields
if any(
field not in vars(server_args)
for field in resolvable_fields(type(server_args))
):
# Bare object.__new__ fixtures: dataclass defaults live on
# the class, not the instance — nothing was populated, so
# treat it as a sentinel (hasattr would see the class
# defaults and materialize them, clobbering resolved flags).
return
declarations = ()
declarations = list(declarations or ()) + self._runtime_overrides
from sglang.srt.arg_groups.overrides import (
apply_model_overrides,
assert_flag_parity,
)
# Resolve into a fresh container and only install it once everything
# passed: a failed resolution (gate validation or the parity assert)
# must not leave the process-global flags half-written for callers
# that catch the error or republish (same install-fresh semantics as
# reset_context()).
flags = Flags()
apply_model_overrides(flags, server_args, declarations)
# Transition-period drift guard: dual-apply keeps the declared fields
# on server_args byte-identical to the resolved flag leaves.
assert_flag_parity(
flags,
server_args,
{field for _source, decl in declarations for field in decl},
)
# The capture tier is not part of the static resolution: carry it
# across re-resolves so runtime-stage recording cannot clobber a
# capture-time write (set_server_args re-seeds it per lifecycle).
flags.capture = self.flags.capture
self.flags = flags
_PARALLEL = ParallelContext()
_CONTEXT = RuntimeContext(parallel=_PARALLEL)
@@ -550,10 +342,9 @@ def get_flags() -> Flags:
def reset_context() -> None:
"""Clear the context-owned store (unit-test teardown): drop the published
``server_args`` and install a fresh, unfrozen ``Flags``.
``server_args`` and install a fresh ``Flags``.
Wrapper subsystems (``parallel``) hold no state and are unaffected.
"""
_CONTEXT._server_args = None
_CONTEXT.flags = Flags()
_CONTEXT._runtime_overrides = []