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