[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
|
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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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":
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user