[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 no_cli: bool = False
# When True, this field may be written by config resolution (model # When True, this field may be written by config resolution (model
# overrides and post-process passes): it is part of the whitelist accepted # 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 # by the declaration stash, and its resolved value materializes onto the
# flags tier (the server_args field itself stays the pristine user input). # field at the end of __post_init__.
resolvable: bool = False resolvable: bool = False
+15 -112
View File
@@ -14,8 +14,9 @@
"""Declarative model-override registry. """Declarative model-override registry.
Model-identity adjustments to the server configuration are DECLARED here and Model-identity adjustments to the server configuration are DECLARED here and
resolved into the flags tier through the ``apply_model_overrides`` gate — materialized onto ``server_args`` at the end of ``__post_init__`` (gate
model code never mutates ``ServerArgs``, which stays the pristine user input. order, last writer wins) — model code never mutates ``ServerArgs`` fields
imperatively.
Two declaration forms, keyed on ``hf_config.architectures[0]``: Two declaration forms, keyed on ``hf_config.architectures[0]``:
@@ -30,11 +31,10 @@ from __future__ import annotations
import dataclasses import dataclasses
import logging 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.arg_groups.arg_utils import resolvable_fields
from sglang.srt.model_executor.cuda_graph_config import Backend 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 ( from sglang.srt.utils.common import (
cpu_has_amx_support, cpu_has_amx_support,
get_device_capability, 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: def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
"""Declare a load-time resolved field (model-file config overrides, """Declare a load-time resolved field (model-file config overrides,
weight-resolved dtypes): apply it onto the published ``server_args`` — weight-resolved dtypes) on the published ``server_args``: resolution has
resolution has already materialized, so post-init declarations write already materialized, so the declaration writes through, joining the
through — and record it into the flags tier through the runtime gate.""" declaration stash for provenance and republish consistency."""
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context
ctx = get_context() server_args = get_context().server_args
entry = (source, dict(declared)) validate_declarations(server_args, [(source, dict(declared))])
validate_declarations(ctx.server_args, [entry]) override = getattr(server_args, "override", None)
_apply_fields(ctx.server_args, declared) if override is not None:
ctx.record_runtime_overrides([entry]) override(source, **declared)
else:
# Config-shaped fixtures without the mutation entry point.
_apply_fields(server_args, declared)
def collect_model_override_declarations( def collect_model_override_declarations(
@@ -2054,84 +2057,6 @@ def _dllm_page_size(view: Any) -> dict:
return {} 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( def validate_declarations(
server_args: Any, server_args: Any,
declarations: Sequence[Tuple[str, Dict[str, Any]]], declarations: Sequence[Tuple[str, Dict[str, Any]]],
@@ -2168,25 +2093,3 @@ def _hrm_text_attention_force(view: Any) -> dict:
"attention." "attention."
) )
return {"attention_backend": "triton"} 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, compute_position,
) )
from sglang.srt.model_executor.forward_context import get_attn_backend 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.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
@@ -634,7 +634,7 @@ class TboForwardBatchPreparer:
sum_field=None, sum_field=None,
) )
_, child_b.extend_start_loc = compute_position( _, child_b.extend_start_loc = compute_position(
get_flags().attn.backend, get_server_args().attention_backend,
child_b.extend_prefix_lens, child_b.extend_prefix_lens,
child_b.extend_seq_lens, child_b.extend_seq_lens,
child_b.extend_num_tokens, child_b.extend_num_tokens,
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch, ForwardBatch,
ForwardMode, 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.server_args import get_global_server_args
from sglang.srt.utils.common import ( from sglang.srt.utils.common import (
is_cpu, is_cpu,
@@ -335,7 +335,7 @@ class LogitsProcessor(nn.Module):
self.config = config self.config = config
self.vocab_size = config.vocab_size self.vocab_size = config.vocab_size
self.logit_scale = logit_scale 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 self.use_fp32_lm_head = get_global_server_args().enable_fp32_lm_head
if self.use_attn_tp_group: if self.use_attn_tp_group:
self.attn_tp_size = get_parallel().attn_tp_size 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_get_mscale_simple,
yarn_linear_ramp_mask, 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
@@ -45,6 +44,8 @@ if _is_xpu:
import triton import triton
import triton.language as tl import triton.language as tl
from sglang.srt.runtime_context import get_server_args
@triton.jit @triton.jit
def apply_interleaved_rope_kernel( def apply_interleaved_rope_kernel(
@@ -227,7 +228,7 @@ class MRotaryEmbedding(RotaryEmbedding):
last_dim = cos_sin.size()[-1] last_dim = cos_sin.size()[-1]
cos, sin = cos_sin.chunk(2, dim=-1) cos, sin = cos_sin.chunk(2, dim=-1)
if self.mrope_interleaved: 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) cos = apply_interleaved_rope_triton(cos, self.mrope_section)
sin = apply_interleaved_rope_triton(sin, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section)
else: 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.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.utils.hash import murmur_hash32 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.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_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import get_global_server_args 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. # 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_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( def _preprocess_logits(
self, logits: torch.Tensor, sampling_info: SamplingBatchInfo self, logits: torch.Tensor, sampling_info: SamplingBatchInfo
@@ -231,7 +231,7 @@ class Sampler(nn.Module):
positions=positions, positions=positions,
) )
else: else:
backend = get_flags().sampling_backend backend = get_server_args().sampling_backend
if backend == "flashinfer": if backend == "flashinfer":
assert ( assert (
sampling_info.sampling_seed is None sampling_info.sampling_seed is None
-8
View File
@@ -554,14 +554,6 @@ class Scheduler(
self.init_batch_result_processor() 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 self.is_initializing = False
def init_zbal_on_npu(self): 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, get_last_loc_triton_safe,
write_req_to_token_pool_triton, 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.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 import is_cuda, is_hip, is_npu, support_triton
from sglang.srt.utils.common import ceil_align, is_pin_memory_available 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], prefix_tensors: list[torch.Tensor],
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
): ):
if support_triton(get_flags().attn.backend): if support_triton(get_server_args().attention_backend):
prefix_pointers = torch.tensor( prefix_pointers = torch.tensor(
[t.data_ptr() for t in prefix_tensors], [t.data_ptr() for t in prefix_tensors],
dtype=torch.uint64, dtype=torch.uint64,
@@ -175,7 +175,7 @@ def get_last_loc(
req_pool_indices_tensor: torch.Tensor, req_pool_indices_tensor: torch.Tensor,
prefix_lens_tensor: torch.Tensor, prefix_lens_tensor: torch.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") uses_triton_dispatch = attn_backend not in ("ascend", "torch_native")
if _is_hip and uses_triton_dispatch: 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.utils import set_default_torch_dtype
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.platforms import current_platform 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.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.server_args import ( # noqa: F401 (re-export) from sglang.srt.server_args import ( # noqa: F401 (re-export)
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS, CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
@@ -541,7 +541,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.init_threads_binding() self.init_threads_binding()
# Set float32 matmul precision # Set float32 matmul precision
if get_flags().enable_tf32_matmul: if get_server_args().enable_tf32_matmul:
torch.set_float32_matmul_precision("high") torch.set_float32_matmul_precision("high")
# Get available memory before model loading. # 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, kv_cache_scales_loader,
maybe_remap_kv_scale_name, 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.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -446,7 +446,7 @@ class ApertusForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) 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, kv_cache_scales_loader,
maybe_remap_kv_scale_name, 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.srt.utils import add_prefix, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -405,7 +405,7 @@ class ArceeForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) 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, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers 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, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
from sglang.srt.models.utils import WeightsMapper 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -1089,7 +1089,7 @@ class BailingMoELinearForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
params_dtype=torch.float32, params_dtype=torch.float32,
quant_config=quant_config, 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) self.logits_processor = LogitsProcessor(config)
@@ -42,7 +42,7 @@ from sglang.srt.models.bailing_moe_linear import (
BailingMoeV2_5ForCausalLM, BailingMoeV2_5ForCausalLM,
) )
from sglang.srt.models.utils import WeightsMapper 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 from sglang.srt.utils import BumpAllocator, add_prefix
LoraConfig = None LoraConfig = None
@@ -208,7 +208,7 @@ class BailingMoeForCausalLMNextN(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
if hasattr(self.config, "model_type") and config.model_type == "bailing_hybrid": 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_common.utils import enable_nextn_moe_bf16_cast_to_fp8
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
from sglang.srt.models.utils import WeightsMapper 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.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
@@ -172,7 +172,7 @@ class DeepseekModelNextN(nn.Module):
if ( if (
_is_npu _is_npu
and self.quant_config is None and self.quant_config is None
and get_flags().quantization is not None and get_server_args().quantization is not None
): ):
# ascend mtp unquant # ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
@@ -330,7 +330,7 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.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_bpreshuffle_gfx95,
_use_aiter_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.server_args import get_global_server_args
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -541,7 +541,7 @@ class DeepseekV2MoE(nn.Module):
n_shared_experts = ( n_shared_experts = (
0 if config.n_shared_experts is None else int(config.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: # num_fused_shared_experts drives weight remapping in deepseek_weight_loader:
# mlp.shared_experts → mlp.experts.256 when > 0. # mlp.shared_experts → mlp.experts.256 when > 0.
@@ -2703,7 +2703,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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: else:
# ranks other than the last rank will have a placeholder layer # 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 self.num_fused_shared_experts = 0
server_args = get_global_server_args() server_args = get_global_server_args()
if get_flags().disable_shared_experts_fusion: if get_server_args().disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
+3 -3
View File
@@ -125,7 +125,7 @@ from sglang.srt.models.deepseek_v2 import (
_is_npu, _is_npu,
_is_xpu, _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: if not _is_hip:
from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.layers.utils.cp_utils import (
@@ -2171,7 +2171,7 @@ class DeepseekV4ForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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: else:
self.lm_head = PPMissingLayer() self.lm_head = PPMissingLayer()
@@ -2209,7 +2209,7 @@ class DeepseekV4ForCausalLM(nn.Module):
def determine_num_fused_shared_experts(self): def determine_num_fused_shared_experts(self):
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
if get_flags().disable_shared_experts_fusion: if get_server_args().disable_shared_experts_fusion:
return return
disable_reason = None 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_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend from sglang.srt.model_executor.forward_context import get_attn_backend
from sglang.srt.models.deepseek_v4 import DeepseekV4DecoderLayer, DeepseekV4ForCausalLM 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 from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -233,7 +233,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -31,7 +31,7 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader, default_weight_loader,
maybe_remap_kv_scale_name, 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.srt.utils import add_prefix, make_layers
from sglang.utils import get_exception_traceback, logger from sglang.utils import get_exception_traceback, logger
@@ -443,7 +443,7 @@ class Exaone4ForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.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.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode 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.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.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -652,7 +652,7 @@ class ExaoneMoEForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
# For EAGLE3 support # 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.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.exaone_moe import ExaoneMoEForCausalLM, ExaoneMoEModel 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 from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -63,7 +63,7 @@ class ExaoneMoEForCausalLMMTP(ExaoneMoEForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.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_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import get_attn_backend 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.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 from sglang.srt.utils import add_prefix, is_cuda, make_layers
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -477,7 +477,7 @@ class FalconH1ForCausalLM(nn.Module):
quant_config=quant_config, quant_config=quant_config,
org_num_embeddings=config.vocab_size, org_num_embeddings=config.vocab_size,
prefix=add_prefix("lm_head", prefix), 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 = self.lm_head.float()
self.lm_head_multiplier = config.lm_head_multiplier 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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
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_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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -404,7 +404,9 @@ class Glm4MoeSparseMoeBlock(nn.Module):
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_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.config = config
@@ -1184,7 +1186,7 @@ class Glm4MoeForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
@@ -1192,7 +1194,7 @@ class Glm4MoeForCausalLM(nn.Module):
self.capture_aux_hidden_states = False self.capture_aux_hidden_states = False
def determine_num_fused_shared_experts(self): def determine_num_fused_shared_experts(self):
if get_flags().disable_shared_experts_fusion: if get_server_args().disable_shared_experts_fusion:
return return
disable_reason = None 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_common.utils import _is_cuda, _use_aiter
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -188,7 +188,9 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
self.routed_scaling_factor = config.routed_scaling_factor self.routed_scaling_factor = config.routed_scaling_factor
self.n_shared_experts = config.n_shared_experts self.n_shared_experts = config.n_shared_experts
self.num_fused_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.config = config
self.layer_id = layer_id self.layer_id = layer_id
@@ -918,7 +920,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
@@ -939,7 +941,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
self, architecture: str = "Glm4MoeLiteForCausalLM" self, architecture: str = "Glm4MoeLiteForCausalLM"
): ):
self.num_fused_shared_experts = 0 self.num_fused_shared_experts = 0
if get_flags().disable_shared_experts_fusion: if get_server_args().disable_shared_experts_fusion:
return return
disable_reason = None disable_reason = None
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
Glm4MoeLiteDecoderLayer, Glm4MoeLiteDecoderLayer,
Glm4MoeLiteForCausalLM, 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.server_args import get_global_server_args
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
@@ -155,12 +155,12 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = ( 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() @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.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
@@ -141,12 +141,12 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = ( 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() @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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.glm4_moe import Glm4MoeModel from sglang.srt.models.glm4_moe import Glm4MoeModel
from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel 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.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 import add_prefix, get_device_sm, is_cuda, log_info_on_rank0
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
@@ -70,7 +70,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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: else:
# ranks other than the last rank will have a placeholder layer # ranks other than the last rank will have a placeholder layer
@@ -84,7 +84,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration):
self.capture_aux_hidden_states = False self.capture_aux_hidden_states = False
def determine_num_fused_shared_experts(self): def determine_num_fused_shared_experts(self):
if get_flags().disable_shared_experts_fusion: if get_server_args().disable_shared_experts_fusion:
return return
disable_reason = None 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.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.glm4 import Glm4DecoderLayer from sglang.srt.models.glm4 import Glm4DecoderLayer
from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration 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 from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -134,12 +134,12 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
self.num_fused_shared_experts = ( 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() @torch.no_grad()
+3 -3
View File
@@ -68,7 +68,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
@@ -390,7 +390,7 @@ class GptOssAttention(nn.Module):
# Choose dtype of sinks based on attention backend: trtllm_mha requires float32, # Choose dtype of sinks based on attention backend: trtllm_mha requires float32,
# others can use bfloat16 # 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 sinks_dtype = torch.float32 if attn_backend == "trtllm_mha" else torch.bfloat16
self.sinks = nn.Parameter( self.sinks = nn.Parameter(
torch.empty(self.num_heads, dtype=sinks_dtype), requires_grad=False torch.empty(self.num_heads, dtype=sinks_dtype), requires_grad=False
@@ -745,7 +745,7 @@ class GptOssForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
# quant_config=quant_config, # quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False 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_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
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_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.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, make_layers from sglang.srt.utils import LazyValue, add_prefix, make_layers
@@ -649,7 +649,7 @@ class LagunaForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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: else:
self.lm_head = PPMissingLayer() self.lm_head = PPMissingLayer()
+2 -2
View File
@@ -76,7 +76,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -796,7 +796,7 @@ class LLaDA2MoeModelLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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) 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, kv_cache_scales_loader,
maybe_remap_kv_scale_name, 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.srt.utils import add_prefix, is_cuda, is_npu, is_xpu, make_layers
from sglang.utils import get_exception_traceback from sglang.utils import get_exception_traceback
@@ -501,7 +501,7 @@ class LlamaForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) 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.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA 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 ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
add_prefix, add_prefix,
@@ -644,7 +644,7 @@ class LongcatFlashForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False 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_audio import AudioEncoderMixin, MiMoAudioEncoderConfig
from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
@@ -1041,7 +1041,7 @@ class MiMoV2ForCausalLM(nn.Module, AudioEncoderMixin):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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: else:
self.lm_head = PPMissingLayer() self.lm_head = PPMissingLayer()
+2 -2
View File
@@ -44,7 +44,7 @@ from sglang.srt.models.mimo_v2 import (
MiMoV2MLP, MiMoV2MLP,
load_mimo_v2_qkv_proj_weight, 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 from sglang.srt.utils import add_prefix
MiMoV2Config = None MiMoV2Config = None
@@ -259,7 +259,7 @@ class MiMoV2MTP(MiMoV2ForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.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, pad_to_original_num_tokens,
) )
from sglang.srt.models.utils import WeightsMapper 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -919,7 +919,7 @@ class NemotronHForCausalLM(nn.Module):
else lora_config.lora_vocab_padding_size else lora_config.lora_vocab_padding_size
), ),
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: else:
+2 -2
View File
@@ -39,7 +39,7 @@ from sglang.srt.models.nemotron_h import (
NemotronHMoEDecoderLayer, NemotronHMoEDecoderLayer,
) )
from sglang.srt.models.nemotron_h_utils import is_attn_layer 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 from sglang.srt.utils import add_prefix
@@ -339,7 +339,7 @@ class NemotronHForCausalLMMTP(NemotronHForCausalLM):
self.config.hidden_size, self.config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.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.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.runner import get_is_capture_mode 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.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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -148,7 +148,7 @@ def can_fuse_shared_expert(
Caller must still gate on the model/backend support flag. Caller must still gate on the model/backend support flag.
""" """
if ( 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 getattr(config, "shared_expert_intermediate_size", 0) <= 0
or config.shared_expert_intermediate_size != config.moe_intermediate_size or config.shared_expert_intermediate_size != config.moe_intermediate_size
or get_moe_a2a_backend().is_deepep() or get_moe_a2a_backend().is_deepep()
@@ -1003,7 +1003,7 @@ class Qwen2MoeForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
# For EAGLE3 support # 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 Qwen2MLP as Qwen3MLP
from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.qwen2 import Qwen2Model
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_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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu 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.vocab_size,
config.hidden_size, config.hidden_size,
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: else:
+3 -3
View File
@@ -91,7 +91,7 @@ from sglang.srt.models.utils import (
fused_qk_gemma_rmsnorm, fused_qk_gemma_rmsnorm,
fused_qk_gemma_rmsnorm_with_gate, 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 # Utils
from sglang.srt.utils import ( from sglang.srt.utils import (
@@ -133,7 +133,7 @@ cached_get_processor = lru_cache(get_processor)
def _disable_shared_experts_fusion() -> bool: def _disable_shared_experts_fusion() -> bool:
# Resolved lazily: the global server args is not set at module import time # Resolved lazily: the global server args is not set at module import time
# (e.g. when this module is imported by unit tests). # (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: if _is_cuda:
@@ -1177,7 +1177,7 @@ class Qwen3_5ForCausalLM(nn.Module):
# so the model still gets the #25885 multi-streaming path. ROCm-only. # so the model still gets the #25885 multi-streaming path. ROCm-only.
if ( if (
config.model_type == "qwen3_5_moe_text" 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) and not can_fuse_shared_expert(config, quant_config)
): ):
from sglang.srt.arg_groups.overrides import declare_load_time_override 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_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
@@ -157,7 +157,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
if ( if (
is_npu() is_npu()
and self.quant_config is None and self.quant_config is None
and get_flags().quantization is not None and get_server_args().quantization is not None
): ):
# ascend mtp unquant # ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) 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, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
LazyValue, LazyValue,
@@ -960,7 +960,7 @@ class Qwen3MoeForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
self.capture_aux_hidden_states = False 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.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.qwen3_moe import Qwen3MoeForCausalLM, Qwen3MoeModel 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 from sglang.srt.utils import add_prefix
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -63,7 +63,7 @@ class Qwen3MoeForCausalLMMTP(Qwen3MoeForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -47,7 +47,7 @@ from sglang.srt.model_loader.weight_utils import (
sharded_weight_loader, sharded_weight_loader,
) )
from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock 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 ( from sglang.srt.utils import (
LazyValue, LazyValue,
add_prefix, add_prefix,
@@ -1027,7 +1027,7 @@ class Qwen3NextForCausalLM(nn.Module):
quant_config=quant_config, quant_config=quant_config,
org_num_embeddings=config.vocab_size, org_num_embeddings=config.vocab_size,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
# For EAGLE3 support # 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.layers.vocab_parallel_embedding import ParallelLMHead
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils import add_prefix, is_npu
@@ -84,7 +84,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("model.shared_head.head", prefix), 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.logits_processor = LogitsProcessor(config)
# Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so # Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so
@@ -114,7 +114,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
if ( if (
is_npu() is_npu()
and self.quant_config is None and self.quant_config is None
and get_flags().quantization is not None and get_server_args().quantization is not None
): ):
# ascend mtp unquant # ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) 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.mm_utils import run_dp_sharded_mrope_vision_model
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
add_prefix, add_prefix,
@@ -1278,7 +1278,7 @@ class Qwen3VLForConditionalGeneration(nn.Module):
self.config.vocab_size, self.config.vocab_size,
self.config.hidden_size, self.config.hidden_size,
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: 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 ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
DeepseekMHAForwardMixin, 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.server_args import get_global_server_args
from sglang.srt.utils import ( from sglang.srt.utils import (
BumpAllocator, BumpAllocator,
@@ -1241,7 +1241,7 @@ class SarvamMLAForCausalLM(nn.Module):
config.hidden_size, config.hidden_size,
quant_config=quant_config, quant_config=quant_config,
prefix=add_prefix("lm_head", prefix), 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.logits_processor = LogitsProcessor(config)
+2 -2
View File
@@ -41,7 +41,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, make_layers from sglang.srt.utils import add_prefix, is_cuda, make_layers
@@ -471,7 +471,7 @@ class SDARForCausalLM(nn.Module):
config.vocab_size, config.vocab_size,
config.hidden_size, config.hidden_size,
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: else:
+2 -2
View File
@@ -57,7 +57,7 @@ from sglang.srt.models.utils import (
create_fused_set_kv_buffer_arg, create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer, 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.server_args import get_global_server_args
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
@@ -566,7 +566,7 @@ class SDARMoeForCausalLM(nn.Module):
config.vocab_size, config.vocab_size,
config.hidden_size, config.hidden_size,
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: 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_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader 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.server_args import get_global_server_args
from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, make_layers 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.vocab_size,
config.hidden_size, config.hidden_size,
quant_config=quant_config, 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), prefix=add_prefix("lm_head", prefix),
) )
else: 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 ``server_args.py`` are thin shims over this slot), and the object is returned
by reference — the same live instance everywhere, never a copy. by reference — the same live instance everywhere, never a copy.
``get_flags()`` returns the resolved-flags tier: what the system *resolved* ``get_flags()`` returns the runtime-flags tier. Resolved configuration lives
the configuration to (``server_args`` stays the pristine user input). Flags on ``server_args`` fields (declarations materialize at the end of
live in typed dataclass groups (``flags.attn`` / ``flags.moe`` / flat generic ``__post_init__``), so this tier only carries genuine runtime state that is
leaves on ``flags`` itself); reads and writes are plain attribute access. not a function of the configuration alone — today the capture lifecycle
Static groups are writable during resolution and locked by ``freeze()``; (``flags.capture``). Flags live in typed dataclass groups; reads and writes
``flags.capture`` stays writable (capture-time state). Each group offers a are plain attribute access, and each group offers a transactional, test-only
transactional, test-only ``override(**kw)`` that also works on frozen groups. ``override(**kw)``.
""" """
from __future__ import annotations from __future__ import annotations
@@ -238,18 +238,13 @@ class _FlagGroupBase:
f"{type(self).__name__} has no flag '{name}' (leaves are " f"{type(self).__name__} has no flag '{name}' (leaves are "
"declared as dataclass fields; check for typos)" "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) object.__setattr__(self, name, value)
@contextmanager @contextmanager
def override(self, **kwargs): def override(self, **kwargs):
"""Temporarily force flag values, restoring on exit. Transactional """Temporarily force flag values, restoring on exit. Transactional
(keys validated before any write) and usable on frozen groups — this (keys validated before any write) — the test-only injection
is the test-only injection primitive.""" primitive."""
fields = type(self).__dataclass_fields__ fields = type(self).__dataclass_fields__
unknown = set(kwargs) - set(fields) unknown = set(kwargs) - set(fields)
if unknown: if unknown:
@@ -266,37 +261,6 @@ class _FlagGroupBase:
object.__setattr__(self, name, value) 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 @dataclasses.dataclass
class CaptureFlags(_FlagGroupBase): class CaptureFlags(_FlagGroupBase):
"""Capture-time flags; never frozen (written during cuda-graph capture).""" """Capture-time flags; never frozen (written during cuda-graph capture)."""
@@ -307,96 +271,28 @@ class CaptureFlags(_FlagGroupBase):
@dataclasses.dataclass @dataclasses.dataclass
class Flags(_StaticFlags): class Flags(_FlagGroupBase):
"""Root of the resolved-flags tier. """Root of the runtime-flags tier.
Family groups hang off it (``flags.attn`` / ``flags.moe`` / ``flags.capture``); Resolved configuration lives on ``server_args`` fields (materialized at
single generic flags live flat on this container, declared as fields here. the end of ``__post_init__``) — this tier only carries genuine runtime
``freeze()`` locks the container and every static sub-group; ``capture`` state whose value is not a function of the configuration alone, grouped
stays writable. 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) 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: class RuntimeContext:
"""Container for the structured runtime accessors; exposes ``parallel``, """Container for the structured runtime accessors; exposes ``parallel``,
``server_args``, and ``flags``.""" ``server_args``, and ``flags``."""
__slots__ = ("parallel", "_server_args", "flags", "_runtime_overrides") __slots__ = ("parallel", "_server_args", "flags")
def __init__(self, parallel: ParallelContext): def __init__(self, parallel: ParallelContext):
self.parallel = parallel self.parallel = parallel
self._server_args: ServerArgs | None = None self._server_args: ServerArgs | None = None
self.flags = Flags() 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 @property
def server_args(self) -> ServerArgs: def server_args(self) -> ServerArgs:
@@ -412,29 +308,10 @@ class RuntimeContext:
Overwrite-allowed: a re-publish replaces the slot (test kits re-publish Overwrite-allowed: a re-publish replaces the slot (test kits re-publish
per test; production ordering discipline lives at the call-sites, e.g. per test; production ordering discipline lives at the call-sites, e.g.
the draft-worker guard in ``ModelRunner.__init__``). the draft-worker guard in ``ModelRunner.__init__``). The published
object already carries the resolved configuration (declarations
Publishing also resolves the stashed model-override declarations materialize at the end of ``__post_init__``).
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.
""" """
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 # Seed the capture tier for the new lifecycle (defaults for sentinel
# and mock publishes, which carry no config). # and mock publishes, which carry no config).
self.flags.capture.enable_torch_compile = getattr( self.flags.capture.enable_torch_compile = getattr(
@@ -442,91 +319,6 @@ class RuntimeContext:
) )
self._server_args = server_args 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() _PARALLEL = ParallelContext()
_CONTEXT = RuntimeContext(parallel=_PARALLEL) _CONTEXT = RuntimeContext(parallel=_PARALLEL)
@@ -550,10 +342,9 @@ def get_flags() -> Flags:
def reset_context() -> None: def reset_context() -> None:
"""Clear the context-owned store (unit-test teardown): drop the published """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. Wrapper subsystems (``parallel``) hold no state and are unaffected.
""" """
_CONTEXT._server_args = None _CONTEXT._server_args = None
_CONTEXT.flags = Flags() _CONTEXT.flags = Flags()
_CONTEXT._runtime_overrides = []
@@ -17,10 +17,6 @@ register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd")
def _make_server_args(*, sampling_backend: str) -> SimpleNamespace: def _make_server_args(*, sampling_backend: str) -> SimpleNamespace:
# The install gate reads the resolved backend from the flags tier.
from sglang.srt.runtime_context import get_flags
get_flags().sampling_backend = sampling_backend
return SimpleNamespace(sampling_backend=sampling_backend) return SimpleNamespace(sampling_backend=sampling_backend)
@@ -23,15 +23,16 @@ register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd")
def _mock_global_server_args(backend="pytorch"): def _mock_global_server_args(backend="pytorch"):
from sglang.srt.layers import sampler as sampler_mod from sglang.srt.layers import sampler as sampler_mod
from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import (
from sglang.srt.server_args import ServerArgs ServerArgs,
set_global_server_args_for_scheduler,
)
sampler_mod.get_global_server_args = lambda: ServerArgs( # Publish for real: the sampler reads the context slot through
model_path="dummy", # get_server_args(), which a module-attribute rebinding cannot intercept.
sampling_backend=backend, set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", sampling_backend=backend)
) )
# The sampler reads the resolved backend from the flags tier.
get_flags().sampling_backend = backend
class _DummyTPGroup: class _DummyTPGroup:
device_group = None device_group = None
-2
View File
@@ -7,7 +7,6 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.runtime_context import get_flags
from sglang.srt.server_args import ( from sglang.srt.server_args import (
ServerArgs, ServerArgs,
get_global_server_args, get_global_server_args,
@@ -45,7 +44,6 @@ class TestLMHeadFP32(unittest.TestCase):
def _make_logprocessor(self, vocab_size, enable_fp32): def _make_logprocessor(self, vocab_size, enable_fp32):
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
get_flags().enable_dp_lm_head = False
get_global_server_args().enable_fp32_lm_head = enable_fp32 get_global_server_args().enable_fp32_lm_head = enable_fp32
cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None) cfg = SimpleNamespace(vocab_size=vocab_size, final_logit_softcapping=None)
return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None) return LogitsProcessor(cfg, skip_all_gather=True, logit_scale=None)
@@ -38,11 +38,9 @@ def _make_target_verify_batch(bs: int) -> ForwardBatch:
def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch: def _filter(batch: ForwardBatch, *, lo: int, hi: int) -> ForwardBatch:
fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3") fake_args = SimpleNamespace(moe_dense_tp_size=None, attention_backend="fa3")
from sglang.srt.runtime_context import get_flags
with get_parallel().override(attn_tp_size=1), patch.object( with get_parallel().override(attn_tp_size=1), patch.object(
tbo, "get_global_server_args", lambda: fake_args tbo, "get_global_server_args", lambda: fake_args
), get_flags().attn.override(backend="fa3"): ):
return TboForwardBatchPreparer.filter_batch( return TboForwardBatchPreparer.filter_batch(
batch, batch,
start_token_index=lo, start_token_index=lo,
@@ -102,10 +102,6 @@ def _make_model_runner(
sa = SimpleNamespace() sa = SimpleNamespace()
sa.swa_full_tokens_ratio = swa_full_tokens_ratio sa.swa_full_tokens_ratio = swa_full_tokens_ratio
# The configurator reads the resolved ratio from the flags tier.
from sglang.srt.runtime_context import get_flags
get_flags().swa_full_tokens_ratio = swa_full_tokens_ratio
sa.page_size = page_size sa.page_size = page_size
sa.disable_radix_cache = disable_radix_cache sa.disable_radix_cache = disable_radix_cache
sa.chunked_prefill_size = chunked_prefill_size sa.chunked_prefill_size = chunked_prefill_size
@@ -2,7 +2,7 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM
from sglang.srt.runtime_context import get_context, get_flags, reset_context from sglang.srt.runtime_context import get_context, reset_context
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -10,9 +10,8 @@ register_cpu_ci(est_time=4, suite="base-a-test-cpu")
class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase): class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
"""The disable decision is a load-time resolution: it lands on the flags """The disable decision is a load-time resolution: it writes through to
tier through declare_load_time_override (dual-applied onto the published the published config via declare_load_time_override."""
config during the transition)."""
def setUp(self): def setUp(self):
self._saved_server_args = get_context()._server_args self._saved_server_args = get_context()._server_args
@@ -41,7 +40,6 @@ class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
self.assertEqual(model.num_fused_shared_experts, 0) self.assertEqual(model.num_fused_shared_experts, 0)
self.assertTrue(get_flags().disable_shared_experts_fusion)
# post-init declaration writes through to the published config # post-init declaration writes through to the published config
self.assertTrue(server_args.disable_shared_experts_fusion) self.assertTrue(server_args.disable_shared_experts_fusion)
@@ -52,7 +50,6 @@ class TestDeepseekV4SharedExpertFusionPolicy(unittest.TestCase):
DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model) DeepseekV4ForCausalLM.determine_num_fused_shared_experts(model)
self.assertEqual(model.num_fused_shared_experts, 1) self.assertEqual(model.num_fused_shared_experts, 1)
self.assertFalse(get_flags().disable_shared_experts_fusion)
self.assertFalse(server_args.disable_shared_experts_fusion) self.assertFalse(server_args.disable_shared_experts_fusion)
+19 -117
View File
@@ -18,14 +18,11 @@ from unittest.mock import patch
from sglang.srt.arg_groups import overrides as overrides_module from sglang.srt.arg_groups import overrides as overrides_module
from sglang.srt.arg_groups.arg_utils import A, Arg, resolvable_fields from sglang.srt.arg_groups.arg_utils import A, Arg, resolvable_fields
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
OverrideRecord,
apply_model_overrides,
collect_model_override_declarations, collect_model_override_declarations,
register_model_override, register_model_override,
validate_declarations, validate_declarations,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
_StaticFlags,
get_context, get_context,
get_server_args, get_server_args,
reset_context, reset_context,
@@ -233,89 +230,6 @@ class TestResolvedViewAndPasses(CustomTestCase):
) )
@dataclasses.dataclass
class _FakeAttnGroup(_StaticFlags):
backend: str = "unset"
@dataclasses.dataclass
class _FakeFlags(_StaticFlags):
attn: _FakeAttnGroup = dataclasses.field(default_factory=_FakeAttnGroup)
resolved_by_model: str = "unset"
also_resolved: Optional[int] = None
class TestApplyModelOverridesGate(CustomTestCase):
def _fresh(self):
return _FakeFlags(), _FakeArgs()
def test_materializes_declared_and_pristine_leaves(self):
flags, args = self._fresh()
records = apply_model_overrides(
flags, args, [("src", {"resolved_by_model": "dsv4"})]
)
self.assertEqual(flags.resolved_by_model, "dsv4") # declared
self.assertIsNone(flags.also_resolved) # undeclared -> pristine value
self.assertEqual(args.resolved_by_model, "auto") # server_args untouched
self.assertEqual(
records, [OverrideRecord("src", "resolved_by_model", "auto", "dsv4")]
)
def test_last_writer_wins_then_terminal_wins_last(self):
flags, args = self._fresh()
records = apply_model_overrides(
flags,
args,
[
("first", {"resolved_by_model": "a"}),
("second", {"resolved_by_model": "b"}),
],
terminal=[("enforce_disable", {"resolved_by_model": "off"})],
)
self.assertEqual(flags.resolved_by_model, "off")
self.assertEqual([r.resolved for r in records], ["a", "b", "off"])
self.assertEqual(records[1].base, "a") # provenance chains the writers
def test_non_whitelisted_field_rejected_before_any_write(self):
flags, args = self._fresh()
with self.assertRaises(ValueError):
apply_model_overrides(
flags,
args,
[("ok", {"resolved_by_model": "x"}), ("bad", {"plain": 1})],
)
self.assertEqual(flags.resolved_by_model, "unset") # transactional
def test_missing_leaf_rejected_before_any_write(self):
flags, args = self._fresh()
with self.assertRaises(ValueError):
apply_model_overrides(
flags,
args,
[("src", {"resolved_by_model": "x"})],
whitelist={"resolved_by_model", "field_without_leaf"},
)
self.assertEqual(flags.resolved_by_model, "unset")
def test_frozen_flags_rejected(self):
flags, args = self._fresh()
flags.freeze()
with self.assertRaises(RuntimeError):
apply_model_overrides(flags, args, [("src", {"resolved_by_model": "x"})])
def test_leaf_map_routes_to_group_leaf(self):
flags, args = self._fresh()
apply_model_overrides(
flags,
args,
[("src", {"resolved_by_model": "fa3"})],
whitelist={"resolved_by_model"},
leaf_map={"resolved_by_model": "attn.backend"},
)
self.assertEqual(flags.attn.backend, "fa3")
self.assertEqual(flags.resolved_by_model, "unset") # flat leaf untouched
class _IsolatedPublish(CustomTestCase): class _IsolatedPublish(CustomTestCase):
"""Publishing writes the process-global context; save/restore around it.""" """Publishing writes the process-global context; save/restore around it."""
@@ -335,9 +249,9 @@ class _NoOverridableArgs:
x: int = 1 x: int = 1
class TestPublishResolvesFlags(_IsolatedPublish): class TestPublishInstallsSlot(_IsolatedPublish):
"""Publish wiring: stash-carrying publishes resolve into flags via the """Publish wiring: set_server_args installs the already-resolved object
gate; publishes without the stash skip resolution.""" into the context-owned slot (no transformation at publish time)."""
def test_dummy_fixture_has_empty_stash_and_publishes_cleanly(self): def test_dummy_fixture_has_empty_stash_and_publishes_cleanly(self):
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -357,23 +271,12 @@ class TestPublishResolvesFlags(_IsolatedPublish):
get_context().set_server_args(sa) get_context().set_server_args(sa)
self.assertIs(get_server_args(), sa) self.assertIs(get_server_args(), sa)
def test_non_whitelisted_declaration_fails_at_publish(self):
from sglang.srt.runtime_context import get_flags
flags_before = get_flags()
sa = _NoOverridableArgs()
sa._resolved_overrides = [("rogue", {"x": 2})]
with self.assertRaises(ValueError):
get_context().set_server_args(sa)
# a failed publish must leave BOTH the slot and the flags untouched
self.assertIs(get_flags(), flags_before)
class TestGoldenModelOverrides(_IsolatedPublish): class TestGoldenModelOverrides(_IsolatedPublish):
"""Per-arch golden diff for migrated families: the declarative path must """Per-arch golden diff for migrated families: the declarative path must
reproduce the legacy imperative writes byte-identically on server_args reproduce the legacy imperative writes byte-identically on the
(dual-apply) and materialize the same values on the flags tier at materialized server_args fields; the publish round-trip returns the same
publish.""" object."""
_MINI_CONFIG = { _MINI_CONFIG = {
"hidden_size": 64, "hidden_size": 64,
@@ -408,11 +311,13 @@ class TestGoldenModelOverrides(_IsolatedPublish):
return ServerArgs(model_path=config_dir, **server_kwargs) return ServerArgs(model_path=config_dir, **server_kwargs)
def _publish(self, server_args): def _publish(self, server_args):
from sglang.srt.runtime_context import get_flags from sglang.srt.server_args import (
from sglang.srt.server_args import set_global_server_args_for_scheduler get_global_server_args,
set_global_server_args_for_scheduler,
)
set_global_server_args_for_scheduler(server_args) set_global_server_args_for_scheduler(server_args)
return get_flags() return get_global_server_args()
def test_mistral_large3_forces_bfloat16(self): def test_mistral_large3_forces_bfloat16(self):
sa = self._construct("MistralLarge3ForCausalLM", "mistral") sa = self._construct("MistralLarge3ForCausalLM", "mistral")
@@ -601,7 +506,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
] ]
self.assertEqual(len(deterministic_fills), 1) self.assertEqual(len(deterministic_fills), 1)
self.assertEqual(sa.attention_backend, deterministic_fills[0]) self.assertEqual(sa.attention_backend, deterministic_fills[0])
self.assertEqual(flags.attn.backend, deterministic_fills[0]) self.assertEqual(flags.attention_backend, deterministic_fills[0])
def test_deterministic_incompatible_backend_raises(self): def test_deterministic_incompatible_backend_raises(self):
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
@@ -645,8 +550,8 @@ class TestGoldenModelOverrides(_IsolatedPublish):
("_dllm_attention_backend", {"attention_backend": "flashinfer"}), ("_dllm_attention_backend", {"attention_backend": "flashinfer"}),
sa._resolved_overrides, sa._resolved_overrides,
) )
# first MAPPED leaf: attention_backend routes to flags.attn.backend # the deterministic fill lands on the attention_backend field
self.assertEqual(self._publish(sa).attn.backend, "flashinfer") self.assertEqual(self._publish(sa).attention_backend, "flashinfer")
def test_attention_backend_leaf_materializes_end_state(self): def test_attention_backend_leaf_materializes_end_state(self):
# The default-fill pass declares the platform-selected backend; the # The default-fill pass declares the platform-selected backend; the
@@ -660,7 +565,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
] ]
self.assertTrue(declared_values) # default fill declared self.assertTrue(declared_values) # default fill declared
self.assertEqual(sa.attention_backend, declared_values[-1]) # materialized self.assertEqual(sa.attention_backend, declared_values[-1]) # materialized
self.assertEqual(self._publish(sa).attn.backend, declared_values[-1]) self.assertEqual(self._publish(sa).attention_backend, declared_values[-1])
def test_post_materialize_pass_writes_through(self): def test_post_materialize_pass_writes_through(self):
from sglang.srt.arg_groups.overrides import run_post_process_pass from sglang.srt.arg_groups.overrides import run_post_process_pass
@@ -679,12 +584,12 @@ class TestGoldenModelOverrides(_IsolatedPublish):
run_post_process_pass(sa, _force_triton) run_post_process_pass(sa, _force_triton)
if resolved_before != "triton": if resolved_before != "triton":
self.assertEqual(sa.attention_backend, "triton") self.assertEqual(sa.attention_backend, "triton")
self.assertEqual(self._publish(sa).attn.backend, sa.attention_backend) self.assertEqual(self._publish(sa).attention_backend, sa.attention_backend)
def test_attention_backend_user_choice_declares_nothing_extra(self): def test_attention_backend_user_choice_declares_nothing_extra(self):
sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton") sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton")
self.assertEqual(sa.attention_backend, "triton") self.assertEqual(sa.attention_backend, "triton")
self.assertEqual(self._publish(sa).attn.backend, "triton") self.assertEqual(self._publish(sa).attention_backend, "triton")
def test_compatibility_passes_at_callable_level(self): def test_compatibility_passes_at_callable_level(self):
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
@@ -2158,13 +2063,10 @@ class TestGoldenModelOverrides(_IsolatedPublish):
class TestDeclarationValidation(CustomTestCase): class TestDeclarationValidation(CustomTestCase):
def test_declarations_never_mutate_server_args(self): def test_declarations_never_mutate_server_args(self):
flags, args = _FakeFlags(), _FakeArgs() args = _FakeArgs()
declarations = [("src", {"resolved_by_model": "dsv4", "also_resolved": 7})] declarations = [("src", {"resolved_by_model": "dsv4", "also_resolved": 7})]
apply_model_overrides(flags, args, declarations)
validate_declarations(args, declarations) validate_declarations(args, declarations)
# the leaves carry the declared values; the fields stay pristine # validation is a pure whitelist check: the fields stay untouched
self.assertEqual(flags.resolved_by_model, "dsv4")
self.assertEqual(flags.also_resolved, 7)
self.assertEqual(args.resolved_by_model, _FakeArgs.resolved_by_model) self.assertEqual(args.resolved_by_model, _FakeArgs.resolved_by_model)
self.assertEqual(args.also_resolved, _FakeArgs.also_resolved) self.assertEqual(args.also_resolved, _FakeArgs.also_resolved)
+51 -187
View File
@@ -15,13 +15,11 @@ from sglang.srt.runtime_context import (
ParallelContext, ParallelContext,
RuntimeContext, RuntimeContext,
_FlagGroupBase, _FlagGroupBase,
_StaticFlags,
get_context, get_context,
get_flags, get_flags,
get_parallel, get_parallel,
get_server_args, get_server_args,
reset_context, reset_context,
resolve_flag_leaf,
) )
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -215,91 +213,47 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
self.assertFalse(hasattr(server_args_module, "_global_server_args")) self.assertFalse(hasattr(server_args_module, "_global_server_args"))
@dataclasses.dataclass
class _FakeStaticGroup(_StaticFlags):
alpha: int = 1
beta: str = "b"
@dataclasses.dataclass @dataclasses.dataclass
class _FakeCaptureGroup(_FlagGroupBase): class _FakeCaptureGroup(_FlagGroupBase):
gamma: int = 0 gamma: int = 0
class TestFlagsTier(_IsolatedServerArgs): class TestFlagsTier(_IsolatedServerArgs):
"""V3a skeleton: typed dataclass groups, freeze guard, override primitive.""" """Runtime-flags tier: typed groups, typo-safe writes, override primitive.
Resolved configuration lives on server_args fields (materialized at the
end of __post_init__); the flags tier only carries runtime state
(today: the capture lifecycle)."""
def test_wiring_and_groups(self): def test_wiring_and_groups(self):
flags = get_flags() flags = get_flags()
self.assertIs(flags, get_context().flags) self.assertIs(flags, get_context().flags)
self.assertIsInstance(flags, Flags) self.assertIsInstance(flags, Flags)
for group in ("attn", "moe", "capture"): self.assertTrue(hasattr(flags, "capture"))
self.assertTrue(hasattr(flags, group))
self.assertFalse(flags.frozen)
def test_typo_safety(self): def test_typo_safety(self):
group = _FakeStaticGroup() group = _FakeCaptureGroup()
with self.assertRaises(AttributeError): with self.assertRaises(AttributeError):
group.alpha_misspelled = 2 # undeclared leaf group.gamma_misspelled = 2 # undeclared leaf
with self.assertRaises(AttributeError): with self.assertRaises(AttributeError):
get_flags().not_a_flag = 1 get_flags().not_a_flag = 1
def test_static_group_writable_until_freeze(self): def test_override_is_transactional(self):
group = _FakeStaticGroup()
group.alpha = 5
self.assertEqual(group.alpha, 5)
group.freeze()
with self.assertRaises(RuntimeError):
group.alpha = 6
self.assertEqual(group.alpha, 5)
def test_override_is_transactional_and_works_on_frozen(self):
group = _FakeStaticGroup()
group.freeze()
with group.override(alpha=99, beta="x"):
self.assertEqual(group.alpha, 99)
self.assertEqual(group.beta, "x")
self.assertEqual(group.alpha, 1)
self.assertEqual(group.beta, "b")
with self.assertRaises(ValueError):
with group.override(alpha=2, gamma=3): # gamma undeclared
pass
self.assertEqual(group.alpha, 1) # validated before any write
def test_non_static_group_has_no_freeze(self):
group = _FakeCaptureGroup() group = _FakeCaptureGroup()
group.gamma = 42 with group.override(gamma=99):
self.assertEqual(group.gamma, 42) self.assertEqual(group.gamma, 99)
self.assertFalse(hasattr(group, "freeze")) self.assertEqual(group.gamma, 0)
with self.assertRaises(ValueError):
with group.override(gamma=2, delta=3): # delta undeclared
pass
self.assertEqual(group.gamma, 0) # validated before any write
def test_container_freeze_cascades_except_capture(self): def test_reset_context_installs_fresh_flags(self):
flags = Flags() # fresh container, not the process singleton old = get_flags()
flags.freeze() old.capture.enable_torch_compile = True
self.assertTrue(flags.frozen) reset_context()
self.assertTrue(flags.attn.frozen) self.assertIsNot(get_flags(), old)
self.assertTrue(flags.moe.frozen) self.assertFalse(get_flags().capture.enable_torch_compile)
with self.assertRaises(RuntimeError):
flags.attn = flags.attn # container leaves lock too
self.assertFalse(getattr(flags.capture, "_frozen", False))
def test_resolve_flag_leaf_flat_default_and_mapped(self):
flags = Flags()
owner, leaf = resolve_flag_leaf(flags, "some_field")
self.assertIs(owner, flags)
self.assertEqual(leaf, "some_field")
owner, leaf = resolve_flag_leaf(flags, "x", leaf_map={"x": "attn.x"})
self.assertIs(owner, flags.attn)
self.assertEqual(leaf, "x")
def test_reset_context_installs_fresh_unfrozen_flags(self):
try:
old = get_flags()
old.freeze()
reset_context()
self.assertIsNot(get_flags(), old)
self.assertFalse(get_flags().frozen)
finally:
reset_context() # never leave the singleton frozen for other tests
@dataclasses.dataclass @dataclasses.dataclass
@@ -311,84 +265,15 @@ class _FakeResolvedArgs:
_resolved_overrides: list = dataclasses.field(default_factory=list) _resolved_overrides: list = dataclasses.field(default_factory=list)
class TestRuntimeResolutionStages(_IsolatedServerArgs): class TestPublishLifecycle(_IsolatedServerArgs):
"""Runtime stages: post-publish declarations re-resolve the flags tier """Publish installs the resolved server_args and seeds the capture tier."""
atomically; freeze_flags() ends the resolution lifecycle."""
def _publish(self, **kw): def _publish(self, **kw):
args = _FakeResolvedArgs(**kw) args = _FakeResolvedArgs(**kw)
get_context().set_server_args(args) get_context().set_server_args(args)
return args return args
def test_record_before_publish_raises(self): def test_capture_tier_seeded_at_publish(self):
reset_context()
with self.assertRaises(ValueError):
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
def test_record_updates_leaves_and_accumulates_stages(self):
args = self._publish(page_size=1, sampling_backend="flashinfer")
self.assertEqual(get_flags().page_size, 1) # publish-time materialize
# dual-apply transition: the call site keeps its imperative write
args.page_size = 64
get_context().record_runtime_overrides([("stage.runner", {"page_size": 64})])
self.assertEqual(get_flags().page_size, 64)
args.sampling_backend = "pytorch"
get_context().record_runtime_overrides(
[("stage.load", {"sampling_backend": "pytorch"})]
)
self.assertEqual(get_flags().sampling_backend, "pytorch")
self.assertEqual(get_flags().page_size, 64) # earlier stage survives
def test_record_whitelist_violation_rolls_back(self):
self._publish()
with self.assertRaises(ValueError):
get_context().record_runtime_overrides([("bad", {"nope": 1})])
self.assertEqual(get_context()._runtime_overrides, [])
def test_freeze_ends_the_resolution_lifecycle(self):
args = self._publish(page_size=1)
try:
get_context().freeze_flags()
self.assertTrue(get_flags().frozen)
with self.assertRaises(RuntimeError):
get_context().record_runtime_overrides([("late", {"page_size": 64})])
with self.assertRaises(RuntimeError):
get_context().set_server_args(args)
finally:
reset_context()
def test_declare_load_time_override_applies_and_records(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
args = self._publish(page_size=1)
declare_load_time_override("model.load_time", {"page_size": 64})
# post-init declaration: written through to the field and resolved
# into the leaf
self.assertEqual(args.page_size, 64)
self.assertEqual(get_flags().page_size, 64)
self.assertEqual(
get_context()._runtime_overrides,
[("model.load_time", {"page_size": 64})],
)
def test_failed_republish_keeps_previous_lifecycle(self):
args = self._publish(page_size=1)
args.page_size = 64
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
flags_before = get_flags()
bad = _FakeResolvedArgs(page_size=1)
bad._resolved_overrides = [("bad", {"nope": 1})] # gate rejects
with self.assertRaises(ValueError):
get_context().set_server_args(bad)
# previous publish fully intact: slot, flags, and the recorded stages
self.assertIs(get_context()._server_args, args)
self.assertIs(get_flags(), flags_before)
self.assertEqual(
get_context()._runtime_overrides, [("stage", {"page_size": 64})]
)
def test_capture_tier_seeded_at_publish_and_survives_stages(self):
# seeded from the published config
args = self._publish(page_size=1) args = self._publish(page_size=1)
args.enable_torch_compile = True args.enable_torch_compile = True
get_context().set_server_args(args) # re-publish picks up the value get_context().set_server_args(args) # re-publish picks up the value
@@ -396,59 +281,38 @@ class TestRuntimeResolutionStages(_IsolatedServerArgs):
# capture-time write (B4) targets the capture leaf # capture-time write (B4) targets the capture leaf
get_flags().capture.enable_torch_compile = False get_flags().capture.enable_torch_compile = False
self.assertFalse(get_flags().capture.enable_torch_compile) self.assertFalse(get_flags().capture.enable_torch_compile)
# a runtime-stage re-resolve must not clobber the capture write
args.page_size = 64
get_context().record_runtime_overrides([("stage", {"page_size": 64})])
self.assertFalse(get_flags().capture.enable_torch_compile)
# capture stays writable after freeze
try:
get_context().freeze_flags()
get_flags().capture.enable_torch_compile = True
self.assertTrue(get_flags().capture.enable_torch_compile)
finally:
reset_context()
def test_declared_leaf_wins_over_stale_field(self):
# A stash entry always drives the leaf at publish, even if the field
# value diverged (e.g. a fixture that skipped materialization).
@dataclasses.dataclass
class _Args:
enable_dp_lm_head: A[bool, Arg(help="d", resolvable=True)] = True
_resolved_overrides: list = dataclasses.field(default_factory=list)
args = _Args()
args._resolved_overrides = [("dp", {"enable_dp_lm_head": False})]
args._declarations_materialized = True
args.enable_dp_lm_head = False
get_context().set_server_args(args)
self.assertFalse(get_flags().enable_dp_lm_head)
def test_bare_dataclass_publish_skips_materialization(self):
# object.__new__(ServerArgs) fixtures (no __init__, no field values)
# must publish without touching the flags tier — dataclass defaults
# live on the class, so materializing from them would clobber
# previously resolved flags with defaults.
from sglang.srt.server_args import ServerArgs
self._publish(page_size=64)
self.assertEqual(get_flags().page_size, 64)
bare = object.__new__(ServerArgs)
get_context().set_server_args(bare)
self.assertIs(get_server_args(), bare)
self.assertEqual(get_flags().page_size, 64) # not clobbered
def test_capture_tier_defaults_for_sentinel_publish(self): def test_capture_tier_defaults_for_sentinel_publish(self):
get_context().set_server_args(object()) get_context().set_server_args(object())
self.assertFalse(get_flags().capture.enable_torch_compile) self.assertFalse(get_flags().capture.enable_torch_compile)
def test_republish_clears_runtime_overrides(self): def test_declare_load_time_override_writes_through(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
args = self._publish(page_size=1) args = self._publish(page_size=1)
args.page_size = 64 declare_load_time_override("model.load_time", {"page_size": 64})
get_context().record_runtime_overrides([("stage", {"page_size": 64})]) self.assertEqual(args.page_size, 64)
self.assertEqual(get_flags().page_size, 64)
self._publish(page_size=1) # fresh lifecycle def test_declare_load_time_override_validates_whitelist(self):
self.assertEqual(get_flags().page_size, 1) from sglang.srt.arg_groups.overrides import declare_load_time_override
self.assertEqual(get_context()._runtime_overrides, [])
args = self._publish(page_size=1)
with self.assertRaises(ValueError):
declare_load_time_override("bad", {"nope": 1})
self.assertEqual(args.page_size, 1)
def test_declare_load_time_override_records_provenance(self):
from sglang.srt.arg_groups.overrides import declare_load_time_override
from sglang.srt.server_args import ServerArgs
class _Args(_FakeResolvedArgs):
override = ServerArgs.override
args = _Args(page_size=1)
get_context().set_server_args(args)
declare_load_time_override("model.load_time", {"page_size": 64})
self.assertEqual(args.page_size, 64)
self.assertIn(("model.load_time", {"page_size": 64}), args._resolved_overrides)
if __name__ == "__main__": if __name__ == "__main__":
@@ -68,8 +68,8 @@ class TestServerArgsMutationRatchet(CustomTestCase):
f"server_args mutations outside the resolution pipeline grew: " f"server_args mutations outside the resolution pipeline grew: "
f"{count} > baseline {_BASELINE}. Configuration is resolved in " f"{count} > baseline {_BASELINE}. Configuration is resolved in "
"ServerArgs.__post_init__; declare through the pipeline " "ServerArgs.__post_init__; declare through the pipeline "
"(passes / declare_load_time_override / " "(passes / declare_load_time_override) or go through "
"record_runtime_overrides) instead of assigning fields." "ServerArgs.override(source, ...) instead of assigning fields."
) )
if count < _BASELINE: if count < _BASELINE:
self.fail( self.fail(