config: the resolution callbacks into the record go to zero (#36972)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-29 04:18:05 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 48b88e1256
commit b65e677e48
50 changed files with 1087 additions and 979 deletions
+15 -6
View File
@@ -29,8 +29,14 @@ logger = logging.getLogger(__name__)
def handle_attention_backend_compatibility(server_args: Any):
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
model_config_of,
use_mla_backend,
)
cfg = resolving_view(server_args)
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
# The attention_backend write clusters of this handler moved to the
# resolution pipeline (arg_groups/overrides.py), each invoked below at
@@ -127,7 +133,7 @@ def handle_attention_backend_compatibility(server_args: Any):
run_post_process_pass(server_args, _cutedsl_prefill_backend_fill)
prefill_backend, decode_backend = server_args._resolved_attention_backends()
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
if "trtllm_mha" in (prefill_backend, decode_backend):
if prefill_backend == "trtllm_mha" and not (
is_sm90_supported() or is_sm100_supported() or is_sm120_supported()
@@ -184,8 +190,8 @@ def handle_attention_backend_compatibility(server_args: Any):
# Other platforms backends
run_post_process_pass(server_args, _attention_backend_platform_fallbacks)
prefill_backend, decode_backend = server_args._resolved_attention_backends()
if server_args.use_mla_backend() and prefill_backend == "intel_xpu":
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
if use_mla_backend(server_args) and prefill_backend == "intel_xpu":
raise ValueError(
"intel_xpu backend is only supported on decode for MLA models, please set --decode-attention-backend to intel_xpu and do not set --attention-backend or --prefill-attention-backend to intel_xpu for prefill instead use triton."
)
@@ -451,6 +457,8 @@ def handle_multi_item_scoring(server_args: Any):
changing it silently could surprise users who intentionally picked
a non-flashinfer backend.
"""
from sglang.srt.arg_groups.overrides import attention_backends_of
cfg = resolving_view(server_args)
if not cfg.enable_mis:
return
@@ -488,7 +496,7 @@ def handle_multi_item_scoring(server_args: Any):
chunked_prefill_size=-1,
)
prefill_backend, decode_backend = server_args._resolved_attention_backends()
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
assert prefill_backend == "flashinfer" and decode_backend == "flashinfer", (
"Multi-item scoring requires flashinfer attention backend for custom attention mask support. "
f"Please set --attention-backend flashinfer when using --enable-mis. "
@@ -497,6 +505,7 @@ def handle_multi_item_scoring(server_args: Any):
def handle_deterministic_inference(server_args: Any):
from sglang.srt.arg_groups.overrides import model_config_of
from sglang.srt.server_args import (
RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND,
)
@@ -549,7 +558,7 @@ def handle_deterministic_inference(server_args: Any):
is_deepseek_model = False
if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE:
try:
hf_config = server_args.get_model_config().hf_config
hf_config = model_config_of(server_args).hf_config
model_arch = hf_config.architectures[0]
is_deepseek_model = model_arch in [
"DeepseekV2ForCausalLM",
+132 -15
View File
@@ -110,6 +110,8 @@ def apply_cuda_graph_compatibility(server_args: Any):
prefill backend (this folds in the old
--enforce-piecewise-cuda-graph contract).
"""
from sglang.srt.arg_groups.overrides import attention_backends_of, model_config_of
cfg = resolving_view(server_args)
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
return
@@ -120,11 +122,13 @@ def apply_cuda_graph_compatibility(server_args: Any):
# this runs first, so piecewise would otherwise silently win.
if (
cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE
and server_args.get_model_config().is_multimodal_piecewise_cuda_graph_supported
and not server_args.get_model_config().is_multimodal_breakable_cuda_graph_supported
and model_config_of(server_args).is_multimodal_piecewise_cuda_graph_supported
and not model_config_of(
server_args
).is_multimodal_breakable_cuda_graph_supported
# Keep trtllm_mla on the preferred breakable path, which now serves
# MLA by falling back to the flashinfer MLA impl for extend.
and server_args._resolved_attention_backends()[0] != "trtllm_mla"
and attention_backends_of(resolved_view(server_args))[0] != "trtllm_mla"
):
logger.info(
"Using tc_piecewise CUDA graph for validated multimodal " "decoder prefill."
@@ -149,12 +153,14 @@ def disable_tc_piecewise_cudagraph_if_incompatible(server_args: Any):
"""TcPiecewise (torch.compile + piecewise) is incompatible with
these configurations. Most are torch.compile / dynamo limitations.
"""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
rules = [
(
"model-arch blacklist",
lambda: server_args.get_model_config().is_piecewise_cuda_graph_disabled_model,
lambda: model_config_of(server_args).is_piecewise_cuda_graph_disabled_model,
),
("DP attention", lambda: resolved_view(server_args).enable_dp_attention),
("full torch.compile mode", lambda: cfg.enable_torch_compile),
@@ -177,8 +183,10 @@ def disable_tc_piecewise_cudagraph_if_incompatible(server_args: Any):
("LoRA", lambda: bool(cfg.lora_paths) or cfg.enable_lora),
(
"multimodal model",
lambda: server_args.get_model_config().is_multimodal
and not server_args.get_model_config().is_multimodal_piecewise_cuda_graph_supported,
lambda: model_config_of(server_args).is_multimodal
and not model_config_of(
server_args
).is_multimodal_piecewise_cuda_graph_supported,
),
(
"GGUF quantization",
@@ -236,6 +244,8 @@ def disable_breakable_cudagraph_if_incompatible(server_args: Any):
memory-saver rejection in its own __init__; config-time rules can be
added here as they're discovered.
"""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import is_deepseek_v4
from sglang.srt.layers.cp.bcg import supports_prefill_cp_bcg
@@ -245,7 +255,7 @@ def disable_breakable_cudagraph_if_incompatible(server_args: Any):
# c4 indexer scratch is pinned in the capture pool and OOMs. Disable.
(
"DeepSeek-V4 (heavy capture-pool memory pressure)",
lambda: is_deepseek_v4(server_args.get_model_config().hf_config),
lambda: is_deepseek_v4(model_config_of(server_args).hf_config),
),
# CP all_gather replay size mismatch under BCG.
(
@@ -271,8 +281,10 @@ def disable_breakable_cudagraph_if_incompatible(server_args: Any):
# Multimodal prefill replay faults under BCG; allowlisted archs opt back in.
(
"multimodal model",
lambda: server_args.get_model_config().is_multimodal
and not server_args.get_model_config().is_multimodal_breakable_cuda_graph_supported,
lambda: model_config_of(server_args).is_multimodal
and not model_config_of(
server_args
).is_multimodal_breakable_cuda_graph_supported,
),
]
for name, predicate in rules:
@@ -319,6 +331,8 @@ def disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args: Any):
breakable) trtllm_mla falls back to FlashAttention for prefill and regresses
performance, so disable whichever prefill graph backend is in effect.
"""
from sglang.srt.arg_groups.overrides import attention_backends_of, model_config_of
cfg = resolving_view(server_args)
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
@@ -327,10 +341,10 @@ def disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args: Any):
return
if (
"DeepseekV3ForCausalLM"
not in server_args.get_model_config().hf_config.architectures
not in model_config_of(server_args).hf_config.architectures
):
return
prefill_attention_backend, _ = server_args._resolved_attention_backends()
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
if prefill_attention_backend != "trtllm_mla":
return
logger.warning(
@@ -362,7 +376,7 @@ def apply_deepep_adjustments(server_args: Any):
if bs is None:
# 2048 = documented prefill default; max_bs unresolved here.
max_bs = cfg.cuda_graph_config.prefill.max_bs or 2048
bs = server_args._generate_prefill_cuda_graph_batch_sizes(max_bs)
bs = generate_prefill_cuda_graph_batch_sizes(server_args, max_bs)
aligned = sorted({((b + 7) // 8) * 8 for b in bs})
if aligned != sorted(bs):
logger.info(
@@ -389,6 +403,8 @@ def apply_inkling_prefill_cuda_graph_default(server_args: Any):
auto-disabled for this multimodal arch, and declarative model overrides
materialize too late to steer cuda-graph resolution. Honors an explicit
--cuda-graph-backend-prefill / --disable-prefill-cuda-graph."""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if (
cfg.cuda_graph_backend_prefill is not None
@@ -396,7 +412,7 @@ def apply_inkling_prefill_cuda_graph_default(server_args: Any):
or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE
):
return
arch = server_args.get_model_config().hf_config.architectures[0]
arch = model_config_of(server_args).hf_config.architectures[0]
if arch in (
"InklingForConditionalGeneration",
"InklingForConditionalGenerationMTP",
@@ -409,13 +425,15 @@ def apply_inkling_prefill_cuda_graph_default(server_args: Any):
def apply_muse_glimmer_prefill_cuda_graph_max_bs_default(server_args: Any):
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if (
cfg.cuda_graph_max_bs_prefill is not None
or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE
):
return
arch = server_args.get_model_config().hf_config.architectures[0]
arch = model_config_of(server_args).hf_config.architectures[0]
if arch in ("MuseGlimmerForCausalLM", "MuseGlimmerForConditionalGeneration"):
declare_resolution(
server_args,
@@ -430,7 +448,7 @@ def handle_cuda_graph_config(server_args: Any):
parse_cuda_graph_config(server_args)
apply_cuda_graph_compatibility(server_args)
apply_deepep_adjustments(server_args)
server_args._apply_cuda_graph_disaggregation_roles()
apply_cuda_graph_disaggregation_roles(server_args)
validate_cuda_graph_config(server_args)
# Warn on the final resolved config (not inside the compat cascade —
# that path is skipped when the user explicitly sets the backend,
@@ -453,3 +471,102 @@ def validate_cuda_graph_config(server_args: Any):
f"--cuda-graph-config[{phase}].backend={backend!r} not allowed; "
f"allowed: {ALLOWED_BACKENDS_PER_PHASE[phase]}"
)
def generate_prefill_cuda_graph_batch_sizes(server_args: Any, max_bs: int):
"""
Generate the list of batch sizes for prefill CUDA graph capture
based on max_bs. For tc_piecewise prefill, bs carries the
captured token count (one shape knob per phase).
"""
capture_sizes = (
list(range(4, 33, 4))
+ list(range(48, 257, 16))
+ list(range(288, 513, 32))
+ list(range(576, 1024 + 1, 64))
+ list(range(1280, 4096 + 1, 256))
+ list(range(4608, max_bs + 1, 512))
)
capture_sizes = [s for s in capture_sizes if s <= max_bs]
return capture_sizes
def generate_decode_cuda_graph_batch_sizes(server_args: Any, max_bs: int):
"""
Generate the list of batch sizes for CUDA graph capture based on max_bs.
This integrates the logic from cuda_graph_runner.py.
"""
cfg = resolving_view(server_args)
# Handle disable_cuda_graph_padding as the first condition for both spec and non-spec
if cfg.disable_cuda_graph_padding:
capture_bs = list(range(1, max_bs + 1))
elif cfg.speculative_algorithm is None:
# Normal case:
capture_bs = (
[1, 2, 4, 8, 12]
+ list(range(16, 257, 8))
+ list(range(272, 512, 16))
+ list(range(512, max_bs + 1, 32))
)
else:
# Spec decoding case: less padding for smaller batch sizes
capture_bs = (
list(range(1, 9, 1))
+ list(range(10, 33, 2))
+ list(range(40, 65, 4))
+ list(range(72, 257, 8))
+ list(range(272, max_bs + 1, 16))
)
capture_bs = [bs for bs in capture_bs if bs <= max_bs]
if max_bs not in capture_bs:
capture_bs.append(max_bs)
return capture_bs
def generate_cpu_graph_batch_sizes(server_args: Any):
"""
Generate the list of batch sizes for CPU graph capture based on torch_compile_max_bs.
"""
cfg = resolving_view(server_args)
if cfg.disable_cuda_graph_padding:
capture_bs = list(range(1, cfg.torch_compile_max_bs + 1))
else:
capture_bs = sorted(
set().union(
range(1, 17),
range(18, 31, 2),
range(32, 81, 4),
range(84, cfg.torch_compile_max_bs + 1, 8),
{cfg.torch_compile_max_bs},
)
)
capture_bs = [bs for bs in capture_bs if bs <= cfg.torch_compile_max_bs]
return capture_bs
def apply_cuda_graph_disaggregation_roles(server_args: Any):
cfg = resolving_view(server_args)
if cfg.disaggregation_mode == "prefill":
if (Phase.DECODE, "backend") not in server_args._cuda_graph_config_locked:
declare_resolution(
server_args,
"_apply_cuda_graph_disaggregation_roles",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
elif cfg.disaggregation_mode == "decode":
if (Phase.PREFILL, "backend") not in server_args._cuda_graph_config_locked:
declare_resolution(
server_args,
"_apply_cuda_graph_disaggregation_roles",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
@@ -27,6 +27,8 @@ logger = logging.getLogger(__name__)
def handle_expert_pack(server_args: Any) -> None:
"""Normalize expert-pack settings and report all startup errors together."""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if cfg.load_format != "expert_pack":
return
@@ -146,7 +148,7 @@ def handle_expert_pack(server_args: Any) -> None:
)
else:
try:
hf_config = server_args.get_model_config().hf_config
hf_config = model_config_of(server_args).hf_config
except Exception as exc:
errors.append(f"failed to load expert_pack model config: {exc}")
else:
+3 -1
View File
@@ -68,6 +68,8 @@ def handle_hicache_ratio_default(server_args: Any):
def resolve_hicache_dcp_compatibility(server_args: Any):
from sglang.srt.arg_groups.overrides import use_mla_backend
cfg = resolving_view(server_args)
if cfg.dcp_size <= 1 or not cfg.enable_hierarchical_cache:
return
@@ -95,7 +97,7 @@ def resolve_hicache_dcp_compatibility(server_args: Any):
"--enable-hisparse with --dcp-size > 1 is not supported: the "
"HiSparse host pool is constructed without DCP translation."
)
if not server_args.use_mla_backend():
if not use_mla_backend(server_args):
raise NotImplementedError(
"HiCache with --dcp-size > 1 is only supported for MLA models: "
"the index translation lives in MLATokenToKVPoolHost, and the "
@@ -19,7 +19,8 @@ HISPARSE_KV_CACHE_DTYPES = ("bfloat16", "fp8_e4m3")
def _is_hip() -> bool:
from sglang.srt.server_args import is_hip
"""The one place this family asks about ROCm, and the seam the tests patch."""
from sglang.srt.utils.common import is_hip
return is_hip()
@@ -81,6 +82,8 @@ def validate_hisparse_kv_cache_dtype(server_args: ServerArgs) -> None:
def validate_hisparse(server_args: ServerArgs) -> None:
"""Validate --enable-hisparse constraints (model class, radix cache, DSA backend)."""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if not cfg.enable_hisparse:
return
@@ -90,7 +93,7 @@ def validate_hisparse(server_args: ServerArgs) -> None:
is_deepseek_v4,
)
hf_config = server_args.get_model_config().hf_config
hf_config = model_config_of(server_args).hf_config
is_v4_hisparse = is_deepseek_v4(hf_config)
is_hip = _is_hip()
assert is_deepseek_dsa(hf_config) or is_v4_hisparse, (
+19 -11
View File
@@ -37,13 +37,15 @@ def handle_mxfp8_kv_cache_compatibility(server_args: Any) -> None:
def handle_kv4_compatibility(server_args: Any) -> None:
"""Check FP4 KV cache compatibility with the attention backend"""
from sglang.srt.arg_groups.overrides import attention_backends_of, use_mla_backend
cfg = resolving_view(server_args)
if cfg.kv_cache_dtype not in ("nvfp4", "fp4_mx_block16"):
return
use_mla_backend = server_args.use_mla_backend()
prefill_backend, decode_backend = server_args._resolved_attention_backends()
uses_mla = use_mla_backend(server_args)
prefill_backend, decode_backend = attention_backends_of(resolved_view(server_args))
attention_backend = resolved_view(server_args).attention_backend
if is_cuda():
@@ -64,7 +66,7 @@ def handle_kv4_compatibility(server_args: Any) -> None:
)
else:
if prefill_backend == "fa4":
if use_mla_backend: # FA4 + MLA
if uses_mla: # FA4 + MLA
KV4_FA4_MLA_BACKEND_CHOICES = [
"cutlass_mla",
"flashinfer",
@@ -85,7 +87,7 @@ def handle_kv4_compatibility(server_args: Any) -> None:
f"{KV4_FA4_MHA_BACKEND_CHOICES}, but got {decode_backend}"
)
else:
if use_mla_backend: # !FA4 + MLA
if uses_mla: # !FA4 + MLA
KV4_ATTENTION_MLA_BACKEND_CHOICES = [
"cutlass_mla",
"flashinfer",
@@ -120,6 +122,8 @@ def handle_prefill_only_disable_kv_cache(server_args: Any) -> None:
still None, backends haven't settled yet and the resolved (prefill,
decode) pair would be a stale (None, None).
"""
from sglang.srt.arg_groups.overrides import attention_backends_of
cfg = resolving_view(server_args)
if not cfg.prefill_only_disable_kv_cache:
@@ -130,7 +134,7 @@ def handle_prefill_only_disable_kv_cache(server_args: Any) -> None:
"_handle_attention_backend_compatibility() so the prefill backend is resolved."
)
prefill_backend, _ = server_args._resolved_attention_backends()
prefill_backend, _ = attention_backends_of(resolved_view(server_args))
if prefill_backend not in ("fa3", "fa4"):
raise ValueError(
"--prefill-only-disable-kv-cache currently requires the FA prefill backend "
@@ -194,6 +198,8 @@ def handle_cache_compatibility(server_args: Any) -> None:
def handle_unified_memory_pool(server_args: Any) -> None:
from sglang.srt.arg_groups.overrides import attention_backends_of
cfg = resolving_view(server_args)
if not cfg.enable_unified_memory:
return
@@ -236,7 +242,7 @@ def handle_unified_memory_pool(server_args: Any) -> None:
# Both roles: verify routes to either backend depending on
# --speculative-attention-mode.
spec_allowed = {"triton", "trtllm_mla", "cutedsl_mla", "tokenspeed_mla"}
spec_backends = set(server_args._resolved_attention_backends())
spec_backends = set(attention_backends_of(resolved_view(server_args)))
spec_backends.discard(None)
assert spec_backends <= spec_allowed, (
"--enable-unified-memory + DSPARK requires spec-verify-audited "
@@ -273,6 +279,8 @@ def handle_page_major_kv_layout(server_args: Any):
# The unified pool stores state in the page-major envelope-strided layout, so
# enabling it implies --enable-page-major-kv-layout — routing it through the
# single page-major path + stride-aware Triton asserts (set before the guard).
from sglang.srt.arg_groups.overrides import attention_backends_of, use_mla_backend
cfg = resolving_view(server_args)
if cfg.enable_unified_memory:
declare_resolution(
@@ -287,14 +295,14 @@ def handle_page_major_kv_layout(server_args: Any):
# exposes each layer as a DENSE contiguous per-layer view
# (build_dense_mla_views), which the paged MLA kernels consume directly,
# with their kv_indices / block tables remapped to dense ids. Names below
# are the RESOLVED ids from _resolved_attention_backends: "flashinfer" is
# are the RESOLVED ids from attention_backends_of: "flashinfer" is
# FlashInferMLAAttnBackend for an MLA model, "trtllm_mla" the trtllm
# decode kernel; "cutedsl_mla" and "tokenspeed_mla" subclass
# TRTLLMMLABackend and inherit its dense read/write path; "fa3" remaps its
# page_table (in-kernel for captured decode, one funnel for eager).
# flashmla / cutlass_mla share the create_flashmla block-table path and
# can be added the same way once exercised.
if cfg.enable_unified_memory and server_args.use_mla_backend():
if cfg.enable_unified_memory and use_mla_backend(server_args):
allowed_full = {
"triton",
"fa3",
@@ -305,7 +313,7 @@ def handle_page_major_kv_layout(server_args: Any):
}
else:
allowed_full = {"triton"}
backends = set(server_args._resolved_attention_backends())
backends = set(attention_backends_of(resolved_view(server_args)))
backends.discard(None)
assert backends <= allowed_full, (
"--enable-page-major-kv-layout requires the Triton attention backend "
@@ -327,7 +335,7 @@ def handle_page_major_kv_layout(server_args: Any):
# are MLA-hybrid) from GDN models (GQA-hybrid) for the KDA-only caveat.
decode_allowed = {"triton", "flashinfer"}
prefill_allowed = {"triton", "flashkda"}
if server_args.use_mla_backend():
if use_mla_backend(server_args):
decode_allowed.update({"cutedsl", "helion"})
prefill_allowed.update({"cutedsl", "helion"})
resolved_linear_decode = cfg.linear_attn_decode_backend or cfg.linear_attn_backend
@@ -403,7 +411,7 @@ def validate_prefill_only_disable_kv_cache_args(server_args: Any):
# Context-parallel prefill stages K/V through cp_allgather_and_save_kv_cache,
# which writes to the pool via set_kv_buffer. NoOpMHATokenToKVPool intentionally
# raises on writes, so the engine would boot fine but fail on the first request.
if server_args._resolved().attn_cp_size > 1:
if resolved_view(server_args).attn_cp_size > 1:
raise ValueError(
"--prefill-only-disable-kv-cache is incompatible with --attn-cp-size > 1: "
"the context-parallel attention path writes K/V to the pool via set_kv_buffer, "
+15 -8
View File
@@ -7,6 +7,7 @@ import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_late_resolution,
resolving_view,
)
from sglang.srt.environ import envs
@@ -23,7 +24,9 @@ def check_lora_server_args(server_args: Any):
# Enable LoRA if any LoRA paths are provided for backward compatibility.
if cfg.lora_paths:
if cfg.enable_lora is None:
server_args._late_resolution("check_lora_server_args", enable_lora=True)
declare_late_resolution(
server_args, "check_lora_server_args", enable_lora=True
)
logger.warning(
"--enable-lora is set to True because --lora-paths is provided."
)
@@ -34,8 +37,8 @@ def check_lora_server_args(server_args: Any):
if cfg.enable_lora:
if cfg.enable_lora_overlap_loading is None:
server_args._late_resolution(
"check_lora_server_args", enable_lora_overlap_loading=False
declare_late_resolution(
server_args, "check_lora_server_args", enable_lora_overlap_loading=False
)
if cfg.enable_lora_overlap_loading:
@@ -90,11 +93,12 @@ def check_lora_server_args(server_args: Any):
"Expected a string or a dictionary."
)
parsed_lora_paths.append(lora_ref)
server_args._late_resolution(
"check_lora_server_args", lora_paths=parsed_lora_paths
declare_late_resolution(
server_args, "check_lora_server_args", lora_paths=parsed_lora_paths
)
elif isinstance(cfg.lora_paths, dict):
server_args._late_resolution(
declare_late_resolution(
server_args,
"check_lora_server_args",
lora_paths=[
LoRARef(
@@ -107,7 +111,9 @@ def check_lora_server_args(server_args: Any):
],
)
elif cfg.lora_paths is None:
server_args._late_resolution("check_lora_server_args", lora_paths=[])
declare_late_resolution(
server_args, "check_lora_server_args", lora_paths=[]
)
else:
raise ValueError(
f"Invalid type for --lora-paths: {type(cfg.lora_paths)}. "
@@ -117,7 +123,8 @@ def check_lora_server_args(server_args: Any):
# Normalize target modules to a set; keep {"all"} as a sentinel
# that gets resolved model-awarely in lora_manager.init_lora_shapes().
if cfg.lora_target_modules:
server_args._late_resolution(
declare_late_resolution(
server_args,
"check_lora_server_args",
lora_target_modules=set(cfg.lora_target_modules),
)
+125 -15
View File
@@ -9,9 +9,11 @@ from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend
logger = logging.getLogger(__name__)
@@ -40,6 +42,17 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
The coefficient 1.5 is a heuristic value, in the future, we can do better estimation by looking at the model types, hidden sizes or even do a dummy run.
"""
from sglang.srt.arg_groups.cuda_graph_hook import (
generate_cpu_graph_batch_sizes,
generate_decode_cuda_graph_batch_sizes,
generate_prefill_cuda_graph_batch_sizes,
)
from sglang.srt.arg_groups.overrides import (
model_config_of,
post_capture_kv_sizing_planned,
use_mla_backend,
)
cfg = resolving_view(server_args)
# A copy, so an earlier declaration keeps the value it recorded.
cuda_graph_config = copy.deepcopy(cfg.cuda_graph_config)
@@ -139,10 +152,8 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
# Set cuda graph batch sizes
if cfg.device != "cpu":
if decode_cuda_graph_config.bs is None:
decode_cuda_graph_config.bs = (
server_args._generate_decode_cuda_graph_batch_sizes(
decode_cuda_graph_config.max_bs
)
decode_cuda_graph_config.bs = generate_decode_cuda_graph_batch_sizes(
server_args, decode_cuda_graph_config.max_bs
)
else:
decode_cuda_graph_config.max_bs = max(decode_cuda_graph_config.bs)
@@ -164,7 +175,7 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
torch_compile_max_bs=cfg.torch_compile_max_bs
or decode_cuda_graph_config.max_bs,
)
decode_cuda_graph_config.bs = server_args._generate_cpu_graph_batch_sizes()
decode_cuda_graph_config.bs = generate_cpu_graph_batch_sizes(server_args)
assert (
cfg.torch_compile_max_bs > 0
@@ -175,7 +186,7 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
# Refer to pr #15927, by default we set the prefill max_bs to the chunked prefill size.
# For MLA backend, the introduction of piecewise cuda graph will influence the kernel dispatch difference compared to the original mode.
# To avoid the performance regression, we set max_bs to 2048 by default.
if not server_args.use_mla_backend():
if not use_mla_backend(server_args):
prefill_cuda_graph_config.max_bs = cfg.chunked_prefill_size
else:
prefill_cuda_graph_config.max_bs = 2048
@@ -194,10 +205,8 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
)
if prefill_cuda_graph_config.bs is None:
prefill_cuda_graph_config.bs = (
server_args._generate_prefill_cuda_graph_batch_sizes(
prefill_cuda_graph_config.max_bs
)
prefill_cuda_graph_config.bs = generate_prefill_cuda_graph_batch_sizes(
server_args, prefill_cuda_graph_config.max_bs
)
if cuda_graph_config != cfg.cuda_graph_config:
@@ -208,7 +217,7 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
)
if cfg.mem_fraction_static is None:
if server_args.post_capture_kv_sizing_planned():
if post_capture_kv_sizing_planned(server_args):
# Post-capture sizing measures free memory after graph capture, so
# skip the graph/activation reserve; keep only the floor + parallel slack.
reserved_mem = 1536
@@ -230,11 +239,11 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
reserved_mem += activation_tokens * 1.5
# Some adjustments for large parallel size
reserved_mem += cfg.tp_size * cfg.pp_size / 8 * 1024
reserved_mem += server_args.reserve_for_graph_mb()
reserved_mem += reserve_for_graph_mb(server_args)
if gpu_mem is not None and gpu_mem > 60 * 1024:
reserved_mem = max(reserved_mem, 10 * 1024)
# Reserve headroom for DeepEP all-to-all buffers on top of the floor.
reserved_mem += server_args.reserve_for_deepep_a2a_mb()
reserved_mem += reserve_for_deepep_a2a_mb(server_args)
declare_resolution(
server_args,
@@ -250,14 +259,14 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
# so we adjust the mem_fraction_static accordingly. The VLM encoder
# only runs on the prefill stage, so PD decode engines do not need
# this headroom; prefill engines and normal (non-PD) engines do.
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
if (
model_config.is_multimodal
and not cfg.language_only
and not cfg.language_model_only
and cfg.disaggregation_mode != "decode"
):
server_args.adjust_mem_fraction_for_vlm(model_config)
adjust_mem_fraction_for_vlm(server_args, model_config)
# If symm mem is enabled and prealloc size is not set, set it to 4GB
if cfg.enable_symm_mem and not envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.is_set():
@@ -266,3 +275,104 @@ def handle_gpu_memory_settings(server_args: Any, gpu_mem):
"Symmetric memory is enabled, setting symmetric memory prealloc size to 4GB as default."
"Use environment variable SGLANG_SYMM_MEM_PREALLOC_GB_SIZE to change the prealloc size."
)
def reserve_for_graph_mb(server_args: Any) -> float:
from sglang.srt.arg_groups.overrides import use_mla_backend
cfg = resolving_view(server_args)
decode_cuda_graph_config = cfg.cuda_graph_config.decode
prefill_cuda_graph_config = cfg.cuda_graph_config.prefill
reserved_mem = 0.0
if (
cfg.disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED
):
reserved_mem += decode_cuda_graph_config.max_bs * 2
if (
resolved_view(server_args).enable_dp_attention
and cfg.disaggregation_mode != "prefill"
):
# DP attention needs more padding for some operations, and much more for large
# cuda graph max bs (torch allocator / implementation inefficiencies).
reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 3
if decode_cuda_graph_config.max_bs > 300:
reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 1.5
if (
cfg.disaggregation_mode != "decode"
and prefill_cuda_graph_config.backend != Backend.DISABLED
):
if not use_mla_backend(server_args):
# Only non-torch memory is counted; torch memory is reused by cuda graph capture.
reserved_mem += len(prefill_cuda_graph_config.bs) * 8
else:
# MLA backend overhead is much higher than expected with fa3.
reserved_mem += 1.5 * 1024
if (
prefill_cuda_graph_config.backend == Backend.BREAKABLE
and resolved_view(server_args).moe_a2a_backend == "deepep"
):
# Prefill-BCG DeepEP delta (bridge pool + NVL first-touch
# during capture); decode-side DeepEP is a baseline cost.
reserved_mem += 1 * 1024
return reserved_mem
def reserve_for_deepep_a2a_mb(server_args: Any) -> float:
# DeepEP all-to-all buffers captured in the decode graph are real extra
# allocations, reserved on top of the floor.
cfg = resolving_view(server_args)
decode_cuda_graph_config = cfg.cuda_graph_config.decode
if (
cfg.disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED
and resolved_view(server_args).moe_a2a_backend == "deepep"
):
return 2 * 1024
return 0.0
def adjust_mem_fraction_for_vlm(server_args: Any, model_config):
cfg = resolving_view(server_args)
vision_config = getattr(model_config.hf_config, "vision_config", None)
if vision_config is None:
return
# roughly reduce the mem_fraction_static base on params of Vit
original_server_arg_mem_fraction = cfg.mem_fraction_static
# a base mem_fraction_static factor for regular Vit
base_mem_fraction_reduction_ratio = 0.95
vit_num_layers = getattr(vision_config, "num_hidden_layers", 24)
vit_hidden_size = getattr(vision_config, "hidden_size", 1024)
# baseline ViT params (ViT-L/14)
baseline_vit_layers = 24
baseline_vit_hidden_size = 1024
# weight params count
current_complexity_score = vit_num_layers * (vit_hidden_size**2)
baseline_complexity_score = baseline_vit_layers * (baseline_vit_hidden_size**2)
complexity_ratio = (
current_complexity_score / baseline_complexity_score
if baseline_complexity_score > 0
else 1.0
)
# every time the complexity grows 100%, adjust final factor for 10%
sensitivity_scale = 0.1
dynamic_adjustment_factor = 1.0 - sensitivity_scale * (complexity_ratio - 1.0)
dynamic_adjustment_factor = max(0.8, min(1.05, dynamic_adjustment_factor))
final_overall_factor = base_mem_fraction_reduction_ratio * dynamic_adjustment_factor
declare_resolution(
server_args,
"adjust_mem_fraction_for_vlm",
mem_fraction_static=original_server_arg_mem_fraction * final_overall_factor,
)
+36 -17
View File
@@ -8,6 +8,7 @@ from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
mamba_cache_chunk_size,
resolved_view,
resolving_view,
)
@@ -34,6 +35,12 @@ logger = logging.getLogger(__name__)
def handle_model_specific_adjustments(server_args: Any):
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
model_config_of,
use_mla_backend,
)
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import (
get_mimo_v2_fused_qkv_expected_tp_size,
@@ -57,7 +64,7 @@ def handle_model_specific_adjustments(server_args: Any):
# key them on.
return
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
hf_config = model_config.hf_config
model_arch = hf_config.architectures[0]
@@ -203,10 +210,14 @@ def handle_model_specific_adjustments(server_args: Any):
import torch
major, _ = torch.cuda.get_device_capability()
server_args._set_default_dsa_kv_cache_dtype(
major, resolved_view(server_args).quantization
from sglang.srt.arg_groups.overrides import (
_dsa_kv_cache_dtype_default,
_dsa_split_backend_resolution,
run_post_process_pass,
)
server_args._set_default_dsa_backends(major)
run_post_process_pass(server_args, _dsa_kv_cache_dtype_default)
run_post_process_pass(server_args, _dsa_split_backend_resolution)
if cfg.enable_prefill_cp:
assert (
@@ -270,7 +281,7 @@ def handle_model_specific_adjustments(server_args: Any):
# MLA prefill CP auto-config: the field declarations moved to
# the override registry (arg_groups/overrides.py:
# _deepseek_family_overrides).
if cfg.enable_prefill_cp and server_args.use_mla_backend():
if cfg.enable_prefill_cp and use_mla_backend(server_args):
declare_resolution(
server_args,
"_handle_model_specific_adjustments",
@@ -304,7 +315,7 @@ def handle_model_specific_adjustments(server_args: Any):
# here for the rest of the DSA family (DeepSeek-V3.2 /
# GLM-5.x) that shares the same decode top-k path.
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
if not server_args._resolved().enable_dp_attention and cfg.nnodes == 1:
if not resolved_view(server_args).enable_dp_attention and cfg.nnodes == 1:
# TODO (Hubert): Put this back later
# server_args.enable_aiter_allreduce_fusion = True
logger.info("Enable Aiter AllReduce Fusion for DeepseekV3ForCausalLM")
@@ -372,8 +383,8 @@ def handle_model_specific_adjustments(server_args: Any):
"intel_xpu",
"aiter",
]
prefill_attn_backend, decode_attn_backend = (
server_args._resolved_attention_backends()
prefill_attn_backend, decode_attn_backend = attention_backends_of(
resolved_view(server_args)
)
assert (
prefill_attn_backend in supported_backends
@@ -387,7 +398,7 @@ def handle_model_specific_adjustments(server_args: Any):
quant_method = get_quantization_config(hf_config)
is_mxfp4_quant_format = quant_method == "mxfp4"
if (
not server_args._resolved().enable_dp_attention
not resolved_view(server_args).enable_dp_attention
and cfg.nnodes == 1
and is_hip()
):
@@ -407,13 +418,13 @@ def handle_model_specific_adjustments(server_args: Any):
if resolved_view(server_args).moe_runner_backend == "triton_kernel":
assert (
server_args._resolved().ep_size == 1
resolved_view(server_args).ep_size == 1
), "Triton kernel MoE is only supported when ep_size == 1"
elif model_arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM"):
if model_arch == "MiMoV2ForCausalLM" and not cfg.encoder_only:
expected_attn_tp_size = get_mimo_v2_fused_qkv_expected_tp_size(hf_config)
view = server_args._resolved()
view = resolved_view(server_args)
attn_dp_size = cfg.dp_size if view.enable_dp_attention else 1
effective_attn_tp_size = cfg.tp_size // attn_dp_size // view.attn_cp_size
if (
@@ -476,7 +487,9 @@ def handle_model_specific_adjustments(server_args: Any):
):
# Default attention backend selection moved to the override registry
# (arg_groups/overrides.py: _gemma4_overrides).
prefill_backend, decode_backend = server_args._resolved_attention_backends()
prefill_backend, decode_backend = attention_backends_of(
resolved_view(server_args)
)
accepted_backends = (
"trtllm_mha",
"triton",
@@ -585,9 +598,13 @@ def handle_model_specific_adjustments(server_args: Any):
def handle_model_capability_adjustments(server_args: Any):
from sglang.srt.arg_groups.cuda_graph_hook import (
generate_prefill_cuda_graph_batch_sizes,
)
from sglang.srt.arg_groups.kv_cache_hook import (
validate_prefill_only_disable_kv_cache_args,
)
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE:
@@ -597,7 +614,7 @@ def handle_model_capability_adjustments(server_args: Any):
run_post_process_pass,
)
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
hf_config = model_config.hf_config
# HRM-Text needs bidirectional prompt attention (prefill), which only
@@ -768,8 +785,8 @@ def handle_model_capability_adjustments(server_args: Any):
)
}
if (Phase.PREFILL, "bs") not in cuda_graph_config_locked:
sizing["bs"] = server_args._generate_prefill_cuda_graph_batch_sizes(
sizing["max_bs"]
sizing["bs"] = generate_prefill_cuda_graph_batch_sizes(
server_args, sizing["max_bs"]
)
declare_resolution(
server_args,
@@ -831,13 +848,15 @@ def handle_mamba_radix_cache(server_args: Any, model_arch: str):
validate_mamba_extra_buffer(
view,
model_arch,
mamba_cache_chunk_size_of=lambda: server_args.mamba_cache_chunk_size,
mamba_cache_chunk_size_of=lambda: mamba_cache_chunk_size(server_args),
)
else:
validate_mamba_no_buffer(view, model_arch)
def handle_language_model_only(server_args: Any):
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if not cfg.language_model_only:
return
@@ -858,7 +877,7 @@ def handle_language_model_only(server_args: Any):
"--language-model-only is incompatible with --disaggregation-mode "
"prefill/decode"
)
architectures = server_args.get_model_config().hf_config.architectures
architectures = model_config_of(server_args).hf_config.architectures
if not any(
a in server_args.LANGUAGE_MODEL_ONLY_ARCHITECTURES for a in architectures
):
@@ -3,6 +3,7 @@
from __future__ import annotations
import glob
import importlib
import logging
import os
@@ -177,7 +178,7 @@ def handle_load_format(server_args: Any):
load_format="gguf",
)
if cfg.load_format == "auto" and server_args._is_mistral_native_format():
if cfg.load_format == "auto" and is_mistral_native_format(server_args):
declare_resolution(
server_args,
"_handle_load_format",
@@ -306,3 +307,62 @@ def validate_transfer_engine(server_args: Any):
return False
else:
return True
def is_mistral_native_format(server_args: Any) -> bool:
"""True iff the checkpoint requires load_format=mistral.
Looks for consolidated*.safetensors with no competing
model*.safetensors; when both weight formats ship in the
same checkpoint (e.g. Mistral-7B-Instruct-v0.3) the HF path is
preferred to avoid loading Mistral-named weights into an
HF-named architecture.
Name override: mistral-large-3 / mistral-small-4 /
leanstral always treat as Mistral-native when params.json
is present -- those families need Mistral weight loading
regardless of which weight files happen to be present.
"""
cfg = resolving_view(server_args)
_MISTRAL_NATIVE_PATTERNS = (
"mistral-large-3",
"mistral-small-4",
"leanstral",
)
name_matches = any(
p in str(cfg.model_path).lower() for p in _MISTRAL_NATIVE_PATTERNS
)
def _check_format(has_params, has_consolidated, has_hf_weights) -> bool:
if has_params and name_matches:
return True
return has_consolidated and not has_hf_weights
if os.path.isdir(cfg.model_path):
return _check_format(
has_params=os.path.exists(os.path.join(cfg.model_path, "params.json")),
has_consolidated=bool(
glob.glob(os.path.join(cfg.model_path, "consolidated*.safetensors"))
),
has_hf_weights=bool(
glob.glob(os.path.join(cfg.model_path, "model*.safetensors"))
),
)
try:
from huggingface_hub import HfApi
files = {s.rfilename for s in HfApi().model_info(cfg.model_path).siblings}
return _check_format(
has_params="params.json" in files,
has_consolidated=any(
f.startswith("consolidated") and f.endswith(".safetensors")
for f in files
),
has_hf_weights=any(
f.startswith("model") and f.endswith(".safetensors") and "/" not in f
for f in files
),
)
except Exception:
return False
+32 -8
View File
@@ -8,7 +8,10 @@ import os
from typing import Any
from sglang.srt.arg_groups.overrides import (
cutedsl_moe_max_num_tokens,
declare_resolution,
max_prefill_buffer_tokens,
max_speculative_num_draft_tokens,
resolved_view,
resolving_view,
)
@@ -24,6 +27,8 @@ def handle_moe_kernel_config(server_args: Any):
# The quantization-driven runner resolutions moved to the pipeline
# (arg_groups/overrides.py: _moe_runner_backend_quant_constraints);
# the compatibility asserts and fusion writes stay below.
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_moe_runner_backend_quant_constraints,
@@ -50,7 +55,7 @@ def handle_moe_kernel_config(server_args: Any):
# modelopt_mixed with non-NVFP4 MoE layers is rejected at load time.
assert (
view.quantization in ["modelopt_fp4", "modelopt_mixed", "nvfp4_online"]
or server_args.get_model_config().nvfp4_moe_meta is not None
or model_config_of(server_args).nvfp4_moe_meta is not None
), f"Invalid quantization '{view.quantization}'. \nFlashInfer CuteDSL MOE currently supports only: 'modelopt_fp4', 'modelopt_mixed' (with NVFP4 MoE layers), 'nvfp4_online', or hybrid NVFP4 models."
assert view.ep_size in [
1,
@@ -116,6 +121,8 @@ def handle_a2a_moe(server_args: Any):
# the resolution pipeline (arg_groups/overrides.py:
# _a2a_backend_overrides / _a2a_ep_size); the per-backend logs,
# asserts, fusion/deepep_mode/env/cuda-graph writes stay below.
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_a2a_backend_overrides,
@@ -259,7 +266,7 @@ def handle_a2a_moe(server_args: Any):
logger.warning("--deepep-mode is ignored for Flashinfer MoE A2A")
if not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set() and (
resolved_view(server_args).quantization == "modelopt_fp4"
or server_args.get_model_config().nvfp4_moe_meta is not None
or model_config_of(server_args).nvfp4_moe_meta is not None
):
envs.SGLANG_MOE_NVFP4_DISPATCH.set(True)
logger.warning(
@@ -291,7 +298,7 @@ def handle_a2a_moe(server_args: Any):
# Skip validation if disaggregation mode is decode.
if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode":
assert (
server_args._required_mori_dispatch_tokens_per_rank()
required_mori_dispatch_tokens_per_rank(server_args)
) <= envs.SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(), (
"SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK (default 4096) "
"must be >= the per-rank MoRI dispatch tokens "
@@ -339,7 +346,7 @@ def handle_a2a_moe(server_args: Any):
# Skip validation if disaggregation mode is decode
if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode":
assert (
server_args._required_pplx_dispatch_tokens_per_rank()
required_pplx_dispatch_tokens_per_rank(server_args)
) <= envs.SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get(), (
"SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK (default 128) "
"must be >= the per-rank pplx dispatch tokens "
@@ -372,7 +379,7 @@ def validate_deepep_v2_dispatch_token_budget(server_args: Any) -> None:
capacity = envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
if view.disaggregation_mode != "decode":
prefill_tokens = server_args.max_prefill_buffer_tokens() or (
prefill_tokens = max_prefill_buffer_tokens(server_args) or (
view.max_prefill_tokens or 0
)
if prefill_tokens > capacity:
@@ -396,7 +403,7 @@ def validate_deepep_v2_dispatch_token_budget(server_args: Any) -> None:
per_rank_pool_bs = max(1, view.max_running_requests // attn_dp_size)
graph_bs = min(graph_bs, per_rank_pool_bs)
tokens_per_req = (
server_args.max_speculative_num_draft_tokens or 1
max_speculative_num_draft_tokens(server_args) or 1
if view.speculative_algorithm
else 1
)
@@ -413,6 +420,8 @@ def validate_deepep_v2_dispatch_token_budget(server_args: Any) -> None:
def validate_deepep_v2_model_architecture(server_args: Any) -> None:
"""Allow DeepEP v2 only where its model workflow is validated."""
from sglang.srt.arg_groups.overrides import model_config_of
if (
parse_connector_type(resolved_view(server_args).model_path)
== ConnectorType.INSTANCE
@@ -424,7 +433,7 @@ def validate_deepep_v2_model_architecture(server_args: Any) -> None:
)
architectures = (
getattr(server_args.get_model_config().hf_config, "architectures", None) or []
getattr(model_config_of(server_args).hf_config, "architectures", None) or []
)
architecture = architectures[0] if architectures else None
@@ -458,7 +467,7 @@ def validate_cutedsl_a2a_token_budget(server_args: Any):
and cfg.disaggregation_mode != "decode"
):
return
required_tokens = server_args.cutedsl_moe_max_num_tokens()
required_tokens = cutedsl_moe_max_num_tokens(server_args)
max_dispatch_tokens_per_rank = (
envs.SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() or 1024
)
@@ -481,3 +490,18 @@ def validate_cutedsl_a2a_token_budget(server_args: Any):
f"{required_per_rank}` or lower the relevant limit "
f"(e.g. --max-prefill-tokens) to <= {max_cutedsl_tokens}."
)
def required_mori_dispatch_tokens_per_rank(server_args: Any) -> int:
"""Max tokens a single rank dispatches through MoRI in one forward."""
cfg = resolving_view(server_args)
return cfg.chunked_prefill_size
def required_pplx_dispatch_tokens_per_rank(server_args: Any) -> int:
"""Max tokens a single rank dispatches through pplx in one forward."""
cfg = resolving_view(server_args)
required = cfg.chunked_prefill_size
if cfg.cuda_graph_max_bs_decode is not None:
required = max(required, cfg.cuda_graph_max_bs_decode)
return required
+346 -40
View File
@@ -38,12 +38,14 @@ import dataclasses
import inspect
import json
import logging
import math
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple
from sglang.srt.arg_groups.arg_utils import field_names, resolvable_fields
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.platforms import current_platform
from sglang.srt.utils.common import (
cpu_has_amx_support,
get_device_capability,
@@ -57,9 +59,11 @@ from sglang.srt.utils.common import (
is_flashinfer_available,
is_gfx95_supported,
is_hip,
is_hopper_with_cuda_12_3,
is_mnnvl_fabric_device,
is_mps,
is_musa,
is_no_spec_infer_or_topk_one,
is_npu,
is_sm90_supported,
is_sm100_supported,
@@ -94,7 +98,7 @@ def register_model_override(architecture: str):
The decorated callable receives ``(server_args, hf_config)``, must not
mutate either, and returns a ``{field: resolved_value}`` dict (possibly
empty when nothing applies). Providers needing derived model data beyond
the HF config go through ``server_args.get_model_config()`` (cached,
the HF config go through ``model_config_of(server_args)`` (cached,
read-only) — never anything mutating.
"""
@@ -293,7 +297,7 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
The stash *is* the resolution result: the bags are projected from it,
`resolution_result` answers from it, and no field is written. A resolver
reading a field another resolver may have decided must read `resolving_view`
(or `ServerArgs._resolved()`), which
(or `resolved_view(server_args)`), which
`test_resolution_reads_the_declarations` pins.
For resolvers inside ``__post_init__``; launcher-stage resolution goes
@@ -460,8 +464,7 @@ def resolved_view(server_args: Any) -> ResolvedView:
overlaid on the fields, snapshotted per call.
For mid-resolution code that is not a pass (``__post_init__`` handlers and
hooks), and for the record's own members that must answer with what
resolution decided -- a declaration-only resolver (a model-specific
hooks) that must answer with what resolution decided -- a declaration-only resolver (a model-specific
override, a registry entry) never writes the field, so a field read there
answers with the raw input."""
return ResolvedView(server_args, overlay=_declaration_overlay(server_args))
@@ -702,7 +705,7 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
if not (is_sm100_supported() and get_device_sm() in (100, 103)):
return {}
backends_unset = server_args.is_attention_backend_not_set()
backends_unset = is_attention_backend_not_set(cfg)
if cfg.speculative_algorithm != "DSPARK":
if not backends_unset:
return {}
@@ -809,7 +812,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
# Set attention backend for DeepSeek
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "dsa"
logger.info("Use dsa attention backend for DeepSeek with DSA.")
if not is_npu() and not is_xpu(): # CUDA or ROCm GPU
@@ -882,7 +885,7 @@ def _deepseek_family_overrides(server_args: Any, hf_config: Any) -> dict:
)
# MLA prefill CP auto-config. Mirrors the NSA CP block above
# (minus the in-seq/round-robin mode split, which MLA CP does not support)
if cfg.enable_prefill_cp and server_args.use_mla_backend():
if cfg.enable_prefill_cp and use_mla_backend(server_args):
logger.warning(
"MLA prefill context parallel is still experimental. "
"Verified on Hopper with the fa3 backend."
@@ -944,7 +947,7 @@ def _minimax_m2_overrides(server_args: Any, hf_config: Any) -> dict:
if (
is_sm100_supported()
and cfg.moe_runner_backend == "auto"
and server_args.get_model_config().quantization == "modelopt_fp4"
and model_config_of(server_args).quantization == "modelopt_fp4"
):
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
logger.info(
@@ -971,7 +974,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
quant_resolved = quant_method
if is_hip():
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "triton"
if cfg.moe_runner_backend == "auto" and quant_resolved == "mxfp8":
overrides["moe_runner_backend"] = "triton"
@@ -994,7 +997,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
if not aiter_fusion_resolved and not envs.SGLANG_M3_ALLOW_CUSTOM_AR.get():
overrides["disable_custom_all_reduce"] = True
elif is_sm100_supported():
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
if (
cfg.kv_cache_dtype == "fp8_e4m3"
and not envs.SGLANG_DISABLE_M3_FP8_ATTN_GEMM.get()
@@ -1025,7 +1028,7 @@ def _minimax_m3_overrides(server_args: Any, hf_config: Any) -> dict:
f"moe_runner_backend={overrides.get('moe_runner_backend', cfg.moe_runner_backend)}."
)
elif is_sm90_supported():
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = "fa3"
page_resolved = cfg.page_size
if (
@@ -1117,7 +1120,7 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
# Set attention backend for GPT-OSS
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
if is_sm100_supported():
overrides["attention_backend"] = "trtllm_mha"
elif is_sm90_supported():
@@ -1254,7 +1257,7 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
default_attention_backend = "trtllm_mha" if is_sm100_supported() else "triton"
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
logger.info(
f"Use {default_attention_backend} as default attention backend for Gemma4"
)
@@ -1265,7 +1268,7 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
elif cfg.attention_backend is None:
overrides["attention_backend"] = default_attention_backend
if is_sm100_supported() and cfg.moe_runner_backend == "auto":
if server_args.get_model_config().quantization == "modelopt_fp4":
if model_config_of(server_args).quantization == "modelopt_fp4":
overrides["quantization"] = "modelopt_fp4"
overrides["moe_runner_backend"] = "flashinfer_trtllm"
logger.info(
@@ -1278,12 +1281,12 @@ def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict:
@_register_for("MossVLForConditionalGeneration")
def _moss_vl_overrides(server_args: Any, hf_config: Any) -> dict:
overrides: Dict[str, Any] = {}
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(resolving_view(server_args)):
overrides["prefill_attention_backend"] = "flashinfer"
logger.info("Use flashinfer as default prefill attention backend for Moss-VL")
prefill_backend = (
overrides.get("prefill_attention_backend")
or server_args.get_attention_backends()[0]
or attention_backends_of(resolved_view(server_args))[0]
)
assert prefill_backend == "flashinfer", (
"MossVLForConditionalGeneration requires flashinfer prefill "
@@ -1323,7 +1326,7 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
if dense_decode is not None:
overrides["decode_attention_backend"] = dense_decode
elif has_sparse_attention:
uses_sparse_backend = cfg.is_attention_backend_not_set() or any(
uses_sparse_backend = is_attention_backend_not_set(cfg) or any(
backend in ("minicpm_flashattn", "minicpm_flashinfer")
for backend in (
cfg.attention_backend,
@@ -1335,7 +1338,7 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict:
raise ValueError(
"MiniCPM sparse attention does not support PD disaggregation"
)
if cfg.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
overrides["attention_backend"] = (
"minicpm_flashinfer"
if is_blackwell_supported()
@@ -1413,7 +1416,7 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict:
logger.info(f"Setting swa_full_tokens_ratio to 0.1 for {model_arch}.")
if cfg.moe_runner_backend == "auto":
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
# nvidia/DeepSeek-V4-Pro-NVFP4 uses the routed TRT-LLM runner.
if model_config.nvfp4_moe_meta is not None:
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
@@ -1477,7 +1480,7 @@ def _inkling_overrides(server_args: Any, hf_config: Any) -> dict:
# supported default when the user left every attention-backend flag unset
# (mirrors the MiniMax-M3 SM100 fa4-default above); an explicit
# --attention-backend / --prefill/decode-attention-backend still wins.
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
inkling_attn_backend = "fa4" if is_sm100_supported() else "triton"
overrides["attention_backend"] = inkling_attn_backend
logger.info(
@@ -1495,7 +1498,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
cache handling and the triton-backend assert stay in the arch branch)."""
cfg = resolving_view(server_args)
model_arch = hf_config.architectures[0]
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
overrides: Dict[str, Any] = {}
is_modelopt = model_config.quantization in [
@@ -1565,7 +1568,7 @@ def _nemotron_h_overrides(server_args: Any, hf_config: Any) -> dict:
else:
overrides["moe_runner_backend"] = "flashinfer_cutlass"
if is_blackwell_supported() and cfg.is_attention_backend_not_set():
if is_blackwell_supported() and is_attention_backend_not_set(cfg):
if cfg.speculative_algorithm is not None:
speculative_algorithm = cfg.speculative_algorithm.upper()
if is_sm100_supported() and cfg.speculative_eagle_topk in (
@@ -1607,13 +1610,14 @@ def _qwen3_5_hybrid_overrides(server_args: Any, hf_config: Any) -> dict:
return {}
sm100_default_attn_backend = "triton"
# trtllm_mha requires speculative_eagle_topk == 1 and page_size > 1.
# _get_default_attn_backend handles the eagle_topk check.
# get_default_attn_backend handles the eagle_topk check.
# There is only one case where page_size=1 is required,
# which is when radix cache is enabled and both extra_buffer
# and spec decoding are disabled.
default_attn_backend = server_args._get_default_attn_backend(
use_mla_backend=server_args.use_mla_backend(),
model_config=server_args.get_model_config(),
default_attn_backend = get_default_attn_backend(
server_args,
use_mla_backend=use_mla_backend(server_args),
model_config=model_config_of(server_args),
)
# The mamba radix-cache pass runs before this dispatch: read the
# declared strategy through the view (the legacy branch observed the
@@ -1746,7 +1750,7 @@ def _olmo2_overrides(server_args: Any, hf_config: Any) -> dict:
def _step3p_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
overrides: Dict[str, Any] = {}
if server_args.is_attention_backend_not_set():
if is_attention_backend_not_set(cfg):
if is_blackwell_supported():
logger.info("Auto-select fa4 attention backend for Step3p7 on Blackwell.")
overrides["attention_backend"] = "fa4"
@@ -1857,7 +1861,7 @@ def _mamba_radix_cache_resolution(view: Any) -> dict:
get_linear_attn_spec_by_arch,
)
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
model_arch = hf_config.architectures[0]
in_branch = model_arch in _MAMBA_RADIX_CACHE_ARCHS
@@ -1895,7 +1899,7 @@ def _dsa_kv_cache_dtype_default(view: Any) -> dict:
PRISTINE dsa split backends (their resolution runs after this pass)."""
from sglang.srt.configs.model_config import is_deepseek_dsa
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
if hf_config.architectures[0] not in _DEEPSEEK_FAMILY_ARCHS:
return {}
if not is_deepseek_dsa(hf_config):
@@ -1965,7 +1969,7 @@ def _dsa_split_backend_resolution(view: Any) -> dict:
capability. The hisparse arm takes precedence under --enable-hisparse."""
from sglang.srt.configs.model_config import is_deepseek_dsa
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
if hf_config.architectures[0] not in _DEEPSEEK_FAMILY_ARCHS:
return {}
if not is_deepseek_dsa(hf_config):
@@ -2065,7 +2069,7 @@ def _deepseek_moe_quant_resolution(view: Any) -> dict:
backend for DeepSeek"), NOT a dispatch-time declaration: the DSA
kv-cache-dtype default earlier in the branch must read the PRISTINE
quantization, so this resolution has to stay at its legacy slot."""
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
model_arch = hf_config.architectures[0]
if model_arch not in _DEEPSEEK_FAMILY_ARCHS:
return {}
@@ -2152,7 +2156,7 @@ def _deepseek_spec_moe_resolution(view: Any) -> dict:
quantization (after _deepseek_moe_quant_resolution) and the pre-a2a
ep_size, exactly like the legacy in-branch writes."""
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
model_arch = hf_config.architectures[0]
if model_arch not in _DEEPSEEK_FAMILY_ARCHS:
return {}
@@ -2198,7 +2202,7 @@ def _deepseek_v4_kv_cache_dtype(view: Any) -> dict:
"""Slot pass in the DeepSeek V4 hook: default the kv-cache dtype to FP8
(bfloat16 on NPU, where the pool geometry differs) and validate the
result. The NPU split-backend writes stay in the hook."""
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
model_arch = hf_config.architectures[0]
if model_arch != "DeepseekV4ForCausalLM":
return {}
@@ -2271,7 +2275,7 @@ def _flashinfer_allreduce_fusion_auto_enable(view: Any) -> dict:
single-node systems. Reads the mid-resolution enable_dp_attention /
moe_a2a_backend (after the DeepSeek CP and a2a declarations), exactly
like the legacy tail block."""
model_arch = view.get_model_config().hf_config.architectures[0]
model_arch = model_config_of(view).hf_config.architectures[0]
if envs.SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION.get() and model_arch in {
"Qwen3_5MoeForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
@@ -2345,7 +2349,7 @@ def _deterministic_is_deepseek_model(view: Any) -> bool:
if parse_connector_type(view.model_path) == ConnectorType.INSTANCE:
return False
try:
hf_config = view.get_model_config().hf_config
hf_config = model_config_of(view).hf_config
return hf_config.architectures[0] in [
"DeepseekV2ForCausalLM",
"DeepseekV3ForCausalLM",
@@ -2413,8 +2417,8 @@ def _attention_backend_default(view: Any) -> dict:
): # override the default attention backend
return {"attention_backend": view.prefill_attention_backend}
if view.attention_backend is None:
backend = view._get_default_attn_backend(
view.use_mla_backend(), view.get_model_config()
backend = get_default_attn_backend(
view, use_mla_backend(view), model_config_of(view)
)
logger.info(
f"Attention backend not specified. Use {backend} backend by default."
@@ -2600,7 +2604,7 @@ def _fa4_page_constraint(view: Any) -> dict:
or view.decode_attention_backend == "fa4"
or view.prefill_attention_backend == "fa4"
)
and not view.use_mla_backend()
and not use_mla_backend(view)
and is_sm100_supported()
# EAGLE topk>1 spec runs the two-pass page-tree cascade, which the FA4
# CUTLASS kernel aborts on at page_size>1. That path only works at
@@ -2641,7 +2645,7 @@ def _attention_backend_platform_fallbacks(view: Any) -> dict:
def _intel_xpu_page_constraint(view: Any) -> dict:
_, decode_backend = attention_backends_of(view)
if decode_backend == "intel_xpu":
if view.use_mla_backend():
if use_mla_backend(view):
supported_page_sizes = [16, 32, 64, 128]
msg = "Intel XPU attention backend for MLA Decode"
else:
@@ -2658,7 +2662,7 @@ def _intel_xpu_page_constraint(view: Any) -> dict:
@register_post_process
def _attention_backend_dual_chunk(view: Any) -> dict:
if (
getattr(view.get_model_config().hf_config, "dual_chunk_attention_config", None)
getattr(model_config_of(view).hf_config, "dual_chunk_attention_config", None)
is not None
):
if view.attention_backend is None:
@@ -3036,3 +3040,305 @@ def _hrm_text_attention_force(view: Any) -> dict:
"attention."
)
return {"attention_backend": "triton"}
def record_of(view: Any) -> Any:
"""The record a view reads through.
For the few helpers a view cannot serve: `get_default_attn_backend` reads
through *both* overlays, so it needs the record the two views are built
from rather than either one of them.
"""
return object.__getattribute__(view, "_server_args")
def is_attention_backend_not_set(cfg: Any):
"""None of the three attention backends has been decided yet.
Takes the view rather than the record: every read is a view read, and the
callers that hold a view (the override providers) would otherwise have to
reach back through it for a record.
"""
return (
cfg.attention_backend is None
and cfg.prefill_attention_backend is None
and cfg.decode_attention_backend is None
)
def get_default_attn_backend(server_args: Any, use_mla_backend: bool, model_config):
"""
Auto select the fastest attention backend.
1. Models with MHA Architecture (e.g: Llama, QWen)
1.1 We will turn on FA3 on hopper unless user use spec decode with topk > 1 or page_size > 1.
1.2 Use trtllm_mha for SM100/SM103 (Blackwell B200/GB200/B300) excluding spec with topk > 1.
Note: trtllm_mha does not support SM120, which will fall back to flashinfer.
1.3 In other cases, we will use flashinfer if available, otherwise use triton.
2. Models with MLA Architecture and using FA3
2.1 We will use FA3 backend on hopper.
2.2 We will use Flashinfer backend on blackwell.
2.3 Otherwise, we will use triton backend.
"""
cfg = resolving_view(server_args)
# OOT platforms provide their own default attention backend.
if current_platform.is_out_of_tree():
return current_platform.get_default_attention_backend()
# Whisper requires flashinfer for cross-attention CUDA graph support.
if "WhisperForConditionalGeneration" in (
model_config.hf_config.architectures or []
):
return "flashinfer"
if not use_mla_backend:
# MHA architecture
if is_hopper_with_cuda_12_3() and is_no_spec_infer_or_topk_one(
resolved_view(server_args)
):
# Note: flashinfer 0.6.1 caused performance regression on Hopper attention kernel
# Before the kernel is fixed, we choose fa3 as the default backend on Hopper MHA
# ref: https://github.com/sgl-project/sglang/issues/17411
return "fa3"
elif (
is_sm100_supported()
and is_no_spec_infer_or_topk_one(resolved_view(server_args))
and (
cfg.speculative_algorithm is None
or cfg.speculative_eagle_topk is not None
)
):
# trtllm_mha requires equal K/V row widths; fa4 carries
# v_head_dim through.
if model_config.has_asymmetric_kv:
return "fa4"
return "trtllm_mha"
elif is_hip():
return "aiter"
elif is_mps():
return "torch_native"
else:
# FlashInfer does not support attention sinks.
if is_flashinfer_available() and not model_config.has_attention_sinks:
return "flashinfer"
return "triton"
else:
# MLA architecture
if is_hopper_with_cuda_12_3():
return "fa3"
elif is_sm100_supported():
return "flashinfer"
elif is_hip():
head_num = model_config.get_num_kv_heads(cfg.tp_size)
# TODO current aiter only support head number 16 or 128 head number
if head_num == 128 or head_num == 16:
return "aiter"
else:
return "triton"
elif is_mps():
return "torch_native"
else:
return "triton"
def use_mla_backend(server_args: Any):
from sglang.srt.configs.model_config import AttentionArch
model_config = model_config_of(server_args)
return model_config.attention_arch == AttentionArch.MLA
def should_report_expert_balancedness(server_args: Any) -> bool:
cfg = resolving_view(server_args)
return cfg.expert_balancedness_report_mode != "off"
def model_config_of(server_args: Any):
"""The model configuration this record describes, built once and memoised.
Takes a view as readily as the record: a view is a read overlay of one
record, the memo has to live on that record either way, and the callers
that hold a view would otherwise all have to unwrap it themselves.
"""
if isinstance(server_args, (ResolvedView, ResolvingConfig)):
server_args = record_of(server_args)
# Lazy init to avoid circular import
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import ModelConfig
memo = getattr(server_args, "_model_config", None)
if memo is not None:
# The key is the path this record carried when the cache was
# filled. The GGUF and ModelScope handlers declare a different
# `model_path`, and a configuration built before them describes
# another checkpoint. `ModelConfig` re-points its own `model_path`
# at the local pull directory when the weights sit behind an
# object-store URI, so its field is not the key. A configuration a
# fixture supplied carries no key and is handed back as it is.
built_from = getattr(server_args, "_model_config_built_from", None)
if built_from is None or built_from == cfg.model_path:
return memo
model_config = ModelConfig.from_server_args(server_args)
server_args._model_config = model_config
server_args._model_config_built_from = cfg.model_path
if model_config.is_hybrid_swa:
logger.info(
"Hybrid SWA model detected. architectures=%s",
model_config.hf_config.architectures,
)
return model_config
def post_capture_kv_sizing_planned(server_args: Any) -> bool:
"""Whether the mem_fraction heuristic may skip the graph reserve; must be
False for any config the runtime won't post-capture-size, else it gets an
under-reserved fraction."""
cfg = resolving_view(server_args)
mla_enabled = use_mla_backend(server_args)
if not envs.SGLANG_ENABLE_POST_CAPTURE_KV_SIZING.get():
return False
if cfg.device != "cuda":
return False
if cfg.dcp_size != 1:
return False
if mla_enabled:
return False
if cfg.kv_cache_dtype == "fp4_e2m1":
return False
if cfg.prefill_only_disable_kv_cache:
return False
if cfg.enable_memory_saver:
return False
if envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get() is not None:
return False
if (
cfg.disaggregation_mode != "prefill"
and cfg.cuda_graph_config.decode.backend == Backend.DISABLED
):
return False
if cfg.disaggregation_mode != "decode":
prefill_cfg = cfg.cuda_graph_config.prefill
# We can only skip eager activation headroom when the largest
# prefill forward batch size is already graph-captured. Otherwise,
# an eager forward will need more memory and lead to OOM.
if (
prefill_cfg.backend == Backend.DISABLED
or cfg.chunked_prefill_size <= 0
or max_prefill_buffer_tokens(server_args) > max(prefill_cfg.bs or (0,))
):
return False
from sglang.srt.configs.model_config import is_deepseek_v4, is_minimax_sparse
hf_config = model_config_of(server_args).hf_config
if is_deepseek_v4(hf_config) or is_minimax_sparse(hf_config):
return False
return True
def cutedsl_moe_max_num_tokens(server_args: Any) -> int:
"""Largest number of tokens a single forward routes through a CuteDSL
MoE layer on one (DP) rank. Single source of truth for both the
standard-allgather wrapper buffers and the FlashInfer A2A dispatcher
budget. Max over the prefill (max_prefill_tokens), piecewise-prefill
capture, and decode/verify bounds; num_tokens_per_req is
speculative_num_draft_tokens under speculative decoding, else 1.
"""
cfg = resolving_view(server_args)
if cfg.speculative_algorithm:
num_tokens_per_req = cfg.speculative_num_draft_tokens or 1
else:
num_tokens_per_req = 1
prefill_tokens = cfg.max_prefill_tokens
cg_config = cfg.cuda_graph_config
if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE:
prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0)
decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0
decode_tokens = decode_max_bs * num_tokens_per_req
return max(prefill_tokens, decode_tokens)
def max_prefill_buffer_tokens(server_args: Any) -> int:
"""Prefill-buffer ceiling: chunked_prefill_size, except PP dynamic
chunking can grow chunks toward max_prefill_tokens and probe at 1.25x."""
cfg = resolving_view(server_args)
chunked = (
cfg.chunked_prefill_size
if cfg.chunked_prefill_size and cfg.chunked_prefill_size > 0
else 0
)
tokens = chunked
if cfg.enable_dynamic_chunking and cfg.pp_size > 1 and chunked:
tokens = max(tokens, cfg.max_prefill_tokens or 0, math.ceil(chunked * 1.25))
return tokens
def mamba_cache_chunk_size(server_args: Any) -> int:
# For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE
# (or mamba_chunk_size if it is defined in the model's config) and page_size.
# It is used to determine the caching point in a sequence during prefill.
# A pre-seeded `_mamba_cache_chunk_size` (fixtures supply one so a dummy
# model never loads an HF config) is honored as-is; otherwise the memo
# is only kept once the record is resolved, because `page_size` below
# is resolution-written.
from sglang.srt.arg_groups.overrides import model_config_of
if not hasattr(server_args, "_mamba_cache_chunk_size"):
try:
from sglang.kernels.ops.attention.fla.chunk_delta_h import (
CHUNK_SIZE as FLA_CHUNK_SIZE,
)
except ImportError:
# Must match sglang.kernels.ops.attention.fla.chunk_delta_h.CHUNK_SIZE
FLA_CHUNK_SIZE = 64
hf_config = model_config_of(server_args).hf_config
chunk_size = getattr(hf_config, "mamba_chunk_size", FLA_CHUNK_SIZE)
page_size = resolved_view(server_args).page_size
assert (
max(chunk_size, page_size) % min(chunk_size, page_size) == 0
), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}"
if not getattr(server_args, "_resolution_finished", False):
return max(chunk_size, page_size)
server_args._mamba_cache_chunk_size = max(chunk_size, page_size)
return server_args._mamba_cache_chunk_size
def max_speculative_num_draft_tokens(server_args: Any) -> Optional[int]:
"""Return the maximum draft-token count speculative decoding may use.
Memoized only once the record is resolved: an answer computed off a raw
record describes inputs resolution is about to rewrite (auto speculative
sizing fills `speculative_num_draft_tokens` in), and a cache filled that
early would keep answering with it.
"""
cfg = resolving_view(server_args)
memo = server_args.__dict__.get("_max_speculative_num_draft_tokens")
if memo is not None:
return memo
if cfg.speculative_num_draft_tokens is None:
result = None
elif not cfg.speculative_adaptive:
result = cfg.speculative_num_draft_tokens
else:
from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config,
)
candidate_steps = resolve_candidate_steps_from_config(
cfg_path=cfg.speculative_adaptive_config,
)
# TODO: adaptive spec currently requires topk=1, so each runtime
# state needs steps + 1 draft-token slots. Revisit this if topk>1
# is supported.
result = max(candidate_steps) + 1
if getattr(server_args, "_resolution_finished", False):
server_args._max_speculative_num_draft_tokens = result
return result
+11 -4
View File
@@ -11,6 +11,7 @@ from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
should_report_expert_balancedness,
)
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import envs
@@ -21,12 +22,14 @@ logger = logging.getLogger(__name__)
def handle_context_parallelism(server_args: Any):
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE:
from sglang.srt.configs.model_config import is_deepseek_dsa
from sglang.srt.layers.cp.utils import CP_V2_DEFAULT_MODEL_CLASSES
model_config = server_args.get_model_config()
model_config = model_config_of(server_args)
hf_config = model_config.hf_config
model_arch = hf_config.architectures[0]
if model_arch in CP_V2_DEFAULT_MODEL_CLASSES:
@@ -154,6 +157,10 @@ def handle_dcp_validation(server_args: Any):
def handle_data_parallelism(server_args: Any):
# The dp_size==1 resets moved to the resolution pipeline
# (arg_groups/overrides.py: _data_parallelism_defaults).
from sglang.srt.arg_groups.cuda_graph_hook import (
generate_prefill_cuda_graph_batch_sizes,
)
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_data_parallelism_defaults,
@@ -213,8 +220,8 @@ def handle_data_parallelism(server_args: Any):
):
clamped = {"max_bs": cfg.chunked_prefill_size}
if (Phase.PREFILL, "bs") not in server_args._cuda_graph_config_locked:
clamped["bs"] = server_args._generate_prefill_cuda_graph_batch_sizes(
clamped["max_bs"]
clamped["bs"] = generate_prefill_cuda_graph_batch_sizes(
server_args, clamped["max_bs"]
)
declare_resolution(
server_args,
@@ -636,7 +643,7 @@ def handle_expert_distribution_metrics(server_args: Any):
"prometheus, both."
)
if server_args.should_report_expert_balancedness() and (
if should_report_expert_balancedness(server_args) and (
cfg.expert_distribution_recorder_mode is None
):
declare_resolution(
@@ -186,6 +186,7 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None:
def handle_encoder_disaggregation(server_args: Any):
from sglang.srt.arg_groups.model_hook import handle_language_model_only
from sglang.srt.arg_groups.overrides import model_config_of
from sglang.srt.arg_groups.validation_hook import validate_ib_devices
from sglang.srt.server_args import resolve_encoder_transfer_backend
@@ -223,7 +224,7 @@ def handle_encoder_disaggregation(server_args: Any):
)
# Validate model type for encoder disaggregation
hf_config = server_args.get_model_config().hf_config
hf_config = model_config_of(server_args).hf_config
model_arch = hf_config.architectures[0]
if cfg.encoder_transfer_backend == "auto":
declare_resolution(
+4 -2
View File
@@ -751,6 +751,8 @@ def handle_multimodal_feature_transport(server_args: Any):
may still auto-select CUDA VMM. The legacy CUDA IPC flag and environment
variable remain supported so existing deployments map to this policy.
"""
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
requested_transport = cfg.mm_feature_transport
legacy_ipc_is_set = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.is_set()
@@ -785,7 +787,7 @@ def handle_multimodal_feature_transport(server_args: Any):
"--encoder-transfer-backend instead."
)
elif (
server_args.get_model_config().is_multimodal
model_config_of(server_args).is_multimodal
and is_cuda()
and cfg.disaggregation_mode == "null"
):
@@ -803,7 +805,7 @@ def handle_multimodal_feature_transport(server_args: Any):
supports_cuda_vmm_feature_transport,
)
if supports_cuda_vmm_feature_transport(server_args.get_model_config()):
if supports_cuda_vmm_feature_transport(model_config_of(server_args)):
requested_transport = "cuda_vmm"
logger.info(
"Multimodal feature transport auto-resolved to "
@@ -337,11 +337,12 @@ def _handle_dflash(server_args: ServerArgs) -> None:
def _target_checkpoint_bundles_dspark_draft(server_args: ServerArgs) -> bool:
from sglang.srt.arg_groups.overrides import model_config_of
from sglang.srt.speculative.dspark_components.dspark_config import (
checkpoint_bundles_dspark_draft,
)
return checkpoint_bundles_dspark_draft(server_args.get_model_config().hf_config)
return checkpoint_bundles_dspark_draft(model_config_of(server_args).hf_config)
def _handle_dspark(server_args: ServerArgs) -> None:
@@ -662,6 +663,8 @@ def _handle_frozen_kv_mtp(server_args: ServerArgs) -> None:
def _handle_eagle_family(server_args: ServerArgs) -> None:
from sglang.srt.arg_groups.overrides import model_config_of
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
@@ -706,7 +709,7 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
"eagle speculative decoding."
)
model_arch = server_args.get_model_config().hf_config.architectures[0]
model_arch = model_config_of(server_args).hf_config.architectures[0]
if model_arch in [
"DeepseekV32ForCausalLM",
"DeepseekV3ForCausalLM",
@@ -9,6 +9,7 @@ import os
from typing import Any, Dict, List, Optional
from sglang.srt.arg_groups.overrides import (
resolved_view,
resolving_view,
)
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
@@ -389,7 +390,7 @@ def validate_ib_devices(server_args: Any, device_str: Optional[str]) -> Optional
def validate_experimental_sgl_marlin(server_args: Any):
view = server_args._resolved()
view = resolved_view(server_args)
if view.moe_runner_backend != "experimental_sgl_marlin":
return
+7 -2
View File
@@ -48,7 +48,10 @@ import torch
import uvloop
import zmq
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.overrides import (
resolved_view,
resolving_view,
)
from sglang.srt.elastic_ep.expert_backup_manager import run_expert_backup_manager
from sglang.srt.entrypoints.engine_info_bootstrap_server import (
EngineInfoBootstrapServer,
@@ -1627,6 +1630,8 @@ class Engine(EngineScoreMixin, EngineBase):
def _set_envs_and_config(server_args: ServerArgs):
from sglang.srt.arg_groups.overrides import attention_backends_of
cfg = resolving_view(server_args)
# Set global environments
# MNNVL fabric (GB200/GB300) multi-node: cross-node NVLink needs NCCL's
@@ -1678,7 +1683,7 @@ def _set_envs_and_config(server_args: ServerArgs):
# Check flashinfer version
if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"):
if "flashinfer" in cfg.get_attention_backends():
if "flashinfer" in attention_backends_of(resolved_view(cfg)):
assert_pkg_version(
"flashinfer_python",
"0.6.17",
@@ -38,6 +38,7 @@ import einops
import torch
import torch.distributed
from sglang.srt.arg_groups.overrides import should_report_expert_balancedness
from sglang.srt.environ import envs
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.observability.metrics_collector import (
@@ -179,7 +180,7 @@ class _ExpertDistributionRecorderReal(ExpertDistributionRecorder):
for k in self._accumulator.get_single_pass_gatherer_keys()
}
if server_args.should_report_expert_balancedness():
if should_report_expert_balancedness(server_args):
logger.info(
"ExpertDistributionRecorder auto start record since "
f"expert_balancedness_report_mode={get_exec().moe.expert_balancedness_report_mode}"
@@ -718,7 +719,7 @@ class _UtilizationRateAccumulatorMixin(_Accumulator):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._enable = self._server_args.should_report_expert_balancedness()
self._enable = should_report_expert_balancedness(self._server_args)
if self._enable:
self.window_sizes = EPLB_BALANCEDNESS_WINDOW_SIZES
@@ -44,7 +44,7 @@ def set_default_server_args(args: "ServerArgs"):
"""
Set default server arguments for NPU backend.
"""
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.overrides import resolving_view, use_mla_backend
cfg = resolving_view(args)
@@ -148,7 +148,7 @@ def set_default_server_args(args: "ServerArgs"):
"set_default_server_args",
hicache_io_backend="kernel_ascend",
)
if args.use_mla_backend():
if use_mla_backend(args):
declare_resolution(
args,
"set_default_server_args",
@@ -74,7 +74,12 @@ def create_trtllm_mla_backend(runner):
if not runner.use_mla_backend:
raise ValueError("trtllm_mla backend can only be used with MLA models.")
if get_parallel().dcp_enabled and get_spec().speculative_algorithm is not None:
_, decode_backend = runner.server_args.get_attention_backends()
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
resolved_view,
)
_, decode_backend = attention_backends_of(resolved_view(runner.server_args))
if decode_backend == "trtllm_mla":
raise ValueError(
"trtllm_mla cannot serve decode context parallelism with speculative "
@@ -56,7 +56,7 @@ from sglang.srt.layers.attention.trtllm_mla_backend import (
)
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.layers.logits_processor import get_in_autotune_dummy_run
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import get_parallel, max_speculative_num_draft_tokens
from sglang.srt.utils import is_flashinfer_available, is_tokenspeed_mla_available
if is_flashinfer_available():
@@ -145,9 +145,7 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
self.device,
self.num_q_heads,
self.kv_lora_rank,
max_q_len=(
model_runner.server_args.max_speculative_num_draft_tokens or 1
),
max_q_len=(max_speculative_num_draft_tokens() or 1),
)
# Pre-JIT the prefill kernel variants. Each cute.compile takes 1-2
+5 -4
View File
@@ -21,6 +21,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional
import torch
from sglang.srt.arg_groups.overrides import resolved_view
from sglang.srt.layers.cp.base import get_cp_strategy
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
from sglang.srt.layers.cp.utils import (
@@ -42,11 +43,11 @@ if TYPE_CHECKING:
def supports_prefill_cp_bcg(server_args: ServerArgs) -> bool:
"""Return whether the selected prefill-CP configuration supports BCG."""
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.overrides import attention_backends_of, resolving_view
cfg = resolving_view(server_args)
resolved = server_args._resolved()
prefill_attention_backend, _ = server_args._resolved_attention_backends()
resolved = resolved_view(server_args)
prefill_attention_backend, _ = attention_backends_of(resolved_view(server_args))
return (
cfg.enable_prefill_cp
and resolved.attn_cp_size == cfg.tp_size
@@ -64,7 +65,7 @@ def filter_prefill_cp_bcg_capture_num_tokens(
capture_num_tokens: list[int], server_args: ServerArgs
) -> list[int]:
"""Keep only token buckets where the zigzag CP strategy can run."""
min_num_tokens = server_args._resolved().attn_cp_size * 2
min_num_tokens = resolved_view(server_args).attn_cp_size * 2
filtered = [size for size in capture_num_tokens if size >= min_num_tokens]
if not filtered:
raise ValueError(
@@ -76,8 +76,10 @@ def create_kt_config_from_server_args(
if server_args.kt_weight_path is None:
return None
from sglang.srt.arg_groups.overrides import model_config_of
num_layers = getattr(
server_args.get_model_config().hf_config, "num_hidden_layers", None
model_config_of(server_args).hf_config, "num_hidden_layers", None
)
return KTConfig(
@@ -8,7 +8,7 @@ from typing import Optional
import torch
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.arg_groups.overrides import cutedsl_moe_max_num_tokens, resolving_view
from sglang.srt.layers.communicator import (
CommunicateWithAllReduceAndLayerNormFn,
LayerCommunicator,
@@ -37,7 +37,7 @@ def resolve_max_m(model_runner) -> int:
decode_config = server_args.cuda_graph_config.decode
prefill_config = server_args.cuda_graph_config.prefill
candidates = [
server_args.cutedsl_moe_max_num_tokens(),
cutedsl_moe_max_num_tokens(model_runner.server_args),
model_runner.max_running_requests,
decode_config.max_bs,
prefill_config.max_bs,
@@ -29,8 +29,10 @@ logger = logging.getLogger(__name__)
def is_post_capture_kv_active(
*, server_args: ServerArgs, is_draft_worker: bool
) -> bool:
from sglang.srt.arg_groups.overrides import post_capture_kv_sizing_planned
return (
server_args.post_capture_kv_sizing_planned()
post_capture_kv_sizing_planned(server_args)
and current_platform.is_cuda()
and not is_draft_worker
)
@@ -44,6 +44,7 @@ from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
get_spec,
max_speculative_num_draft_tokens,
)
from sglang.srt.utils.common import (
ceil_align,
@@ -780,9 +781,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.swa_page_size = cfg.window_size
self.swa_ratio = get_schedule().swa_full_tokens_ratio
self.is_speculative = get_spec().speculative_algorithm is not None
self.online_c128_mtp_max_draft_tokens = (
kvc.server_args.max_speculative_num_draft_tokens or 0
)
self.online_c128_mtp_max_draft_tokens = max_speculative_num_draft_tokens() or 0
self.requested_max_running_requests_per_worker = (
get_schedule().max_running_requests // kvc.ps.attn_dp_size
if get_schedule().max_running_requests is not None
@@ -814,7 +813,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
if self.is_speculative:
# Ring is sized once here, so it must serve the largest adaptive tier.
self._assert_ring_serves_draft_tokens(
kvc.server_args.max_speculative_num_draft_tokens or 0
max_speculative_num_draft_tokens() or 0
)
self.bytes_per_full_token = self._get_bytes_per_full_token()
@@ -30,7 +30,9 @@ from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_model,
get_schedule,
get_spec,
max_prefill_buffer_tokens,
)
from sglang.srt.utils import empty_context, log_info_on_rank0
@@ -342,9 +344,7 @@ def maybe_flashinfer_autotune_extend(
mr = runner.model_runner
# Prefer the per-rank scheduler buffer while preserving the legacy ceiling
# when chunked prefill is disabled.
num_tokens = (
mr.server_args.max_prefill_buffer_tokens() or mr.server_args.max_prefill_tokens
)
num_tokens = max_prefill_buffer_tokens() or get_schedule().max_prefill_tokens
if num_tokens <= (decode_num_tokens or 0):
return # decode-shaped autotune already covered these buckets
is_pd_prefill_target = (
@@ -15,6 +15,8 @@ import sys
from pathlib import Path
from typing import Any
from sglang.srt.arg_groups.overrides import declare_resolution
METADATA_FORMAT_VERSION = 3
GGUF_SHARD_SUFFIX_RE = re.compile(r"-\d{5}-of-\d{5}\.gguf$")
DEEPSEEK_METADATA_FORMAT_VERSION = 4
@@ -215,7 +217,8 @@ def prepare_raw_kimi_server_args(
model_path,
tokenizer_dir=tokenizer_path,
)
server_args._declare(
declare_resolution(
server_args,
"prepare_raw_kimi_server_args",
model_path=str(assets["model_dir"]),
tokenizer_path=str(assets["model_dir"]),
@@ -538,7 +541,8 @@ def prepare_raw_deepseek_server_args(
config_sha256 = _deepseek_digest(
model_value.get("config_sha256"), "model.config_sha256"
)
server_args._declare(
declare_resolution(
server_args,
"prepare_raw_deepseek_server_args",
model_path=str(model_config.parent),
tokenizer_path=str(model_config.parent),
+17 -11
View File
@@ -1672,7 +1672,7 @@ def max_prefill_buffer_tokens() -> int:
Every input is a published leaf (``schedule`` plus the configured PP size),
so this derives from the bags and follows a post-publish override;
``ServerArgs.max_prefill_buffer_tokens`` is the pre-publish equivalent and
``overrides.max_prefill_buffer_tokens`` is the pre-publish equivalent and
``TestDerivedPredicatesAgreeAcrossTiers`` pins the two equal.
"""
import math
@@ -1726,17 +1726,19 @@ def pre_capture_activation_reserve_mb(gpu_mem: float | None) -> float:
# --- Derived config accessors ------------------------------------------------
#
# A few values are computed from several config fields plus the HF config, so
# they are ``ServerArgs`` members rather than namespace leaves. Business code
# must not reach for the startup record to get them: these accessors are the
# named home, and this module — which owns the slot — is the only place that
# reads it. Each one keeps the member's exact semantics, including which model
# they are derived accessors rather than namespace leaves. Business code must
# not reach for the startup record to get them: these accessors are the named
# home, and this module — which owns the slot — is the only place that reads
# it. Each one keeps the pre-publish function's exact semantics, including which model
# config it derives from (always the process's, i.e. the target's).
def mamba_cache_chunk_size() -> int:
"""The caching point granularity for mamba state: ``max(the model's mamba
chunk size, page_size)``. Cached on the config after the first call."""
return get_server_args().mamba_cache_chunk_size
from sglang.srt.arg_groups.overrides import mamba_cache_chunk_size as _of
return _of(get_server_args())
def mamba_checkpoint_grid(tree_page: int) -> int:
@@ -1759,7 +1761,7 @@ def max_speculative_num_draft_tokens() -> int | None:
"""The largest draft-token count speculative decoding may use.
All three inputs are ``spec`` leaves, so this derives from the bags and
follows a post-publish override; ``ServerArgs.max_speculative_num_draft_tokens``
follows a post-publish override; ``overrides.max_speculative_num_draft_tokens``
is the pre-publish equivalent. Adaptive spec resolves the count from its
candidate-step table instead of the flat field.
"""
@@ -1788,7 +1790,9 @@ def _adaptive_draft_token_bound(cfg_path: str | None) -> int:
def uses_mla_backend() -> bool:
"""Whether this process's model runs the MLA attention path."""
return get_server_args().use_mla_backend()
from sglang.srt.arg_groups.overrides import use_mla_backend
return use_mla_backend(get_server_args())
def attention_backends() -> tuple:
@@ -1796,7 +1800,7 @@ def attention_backends() -> tuple:
back to ``attention_backend``.
All three inputs are ``exec.kernel`` leaves, so this derives from the bags
and follows a post-publish override; ``ServerArgs.get_attention_backends``
and follows a post-publish override; ``overrides.attention_backends_of``
is the pre-publish equivalent the resolution pipeline uses. A built runner
stamps its own resolved pair (``ModelRunner.prefill_attention_backend_str``);
read that when there is a runner in hand.
@@ -1810,7 +1814,9 @@ def attention_backends() -> tuple:
def process_model_config():
"""The process's ``ModelConfig`` (built once from the published config)."""
return get_server_args().get_model_config()
from sglang.srt.arg_groups.overrides import model_config_of
return model_config_of(get_server_args())
def cutedsl_moe_max_num_tokens() -> int:
@@ -1818,7 +1824,7 @@ def cutedsl_moe_max_num_tokens() -> int:
Every input is a published leaf (``spec``, ``schedule``, ``exec.graph``), so
this derives from the bags and follows a post-publish override;
``ServerArgs.cutedsl_moe_max_num_tokens`` is the pre-publish equivalent the
``overrides.cutedsl_moe_max_num_tokens`` is the pre-publish equivalent the
resolution pipeline uses. Max over the prefill bound, the piecewise-prefill
capture, and the decode/verify bound.
"""
+1 -616
View File
@@ -37,11 +37,8 @@ import argparse
import copy
import dataclasses
import functools
import glob
import json
import logging
import math
import os
import tempfile
import uuid
from typing import Any, Callable, Dict, List, Literal, Optional, Union
@@ -56,11 +53,9 @@ from sglang.srt.arg_groups.argparse_actions import (
LoRAPathAction,
)
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
mamba_extra_buffer_lazy_of,
mamba_extra_buffer_of,
remote_instance_transfer_engine_of,
resolved_view,
resolving_view,
)
from sglang.srt.environ import envs
@@ -69,23 +64,14 @@ from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.model_executor.cuda_graph_config import (
Backend,
CudaGraphConfig,
Phase,
parse_cuda_graph_config_arg,
with_phase,
)
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.platforms import current_platform
from sglang.srt.speculative.decoupled_spec_io import DecoupledSpecIpcConfig
from sglang.srt.utils.common import (
LORA_TARGET_ALL_MODULES,
SUPPORTED_LORA_TARGET_MODULES,
human_readable_int,
is_flashinfer_available,
is_hip,
is_hopper_with_cuda_12_3,
is_mps,
is_no_spec_infer_or_topk_one,
is_sm100_supported,
json_list_type,
nullable_str,
)
@@ -3745,7 +3731,7 @@ class ServerArgs:
# Everything outside the fields, enumerated from the instance: the raw
# snapshot, the stash, and what resolution memoized -- including the
# `get_model_config()` memo, which the copy carries over rather than
# model-configuration memo, which the copy carries over rather than
# rebuild.
field_names = {field.name for field in dataclasses.fields(self)}
for name, value in vars(self).items():
@@ -3763,91 +3749,10 @@ class ServerArgs:
object.__setattr__(replacement, "_resolution_finished", True)
return replacement
def _declare(self, source: str, **fields: Any) -> None:
"""This record's handlers declaring their resolution writes.
See ``arg_groups.overrides.declare_resolution``, which the hooks these
handlers call reach directly.
"""
from sglang.srt.arg_groups.overrides import declare_resolution
declare_resolution(self, source, **fields)
# ------------------------------------------------------------------
# CUDA graph configuration resolution
# ------------------------------------------------------------------
def _apply_cuda_graph_disaggregation_roles(self):
cfg = resolving_view(self)
if cfg.disaggregation_mode == "prefill":
if (Phase.DECODE, "backend") not in self._cuda_graph_config_locked:
self._declare(
"_apply_cuda_graph_disaggregation_roles",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
elif cfg.disaggregation_mode == "decode":
if (Phase.PREFILL, "backend") not in self._cuda_graph_config_locked:
self._declare(
"_apply_cuda_graph_disaggregation_roles",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
def post_capture_kv_sizing_planned(self) -> bool:
"""Whether the mem_fraction heuristic may skip the graph reserve; must be
False for any config the runtime won't post-capture-size, else it gets an
under-reserved fraction."""
cfg = resolving_view(self)
# use_mla_backend is a method at args time but ModelRunner overwrites it
# with a bool on global_server_args (see the FIXME there) -- handle both.
use_mla = self.use_mla_backend
mla_enabled = use_mla() if callable(use_mla) else use_mla
if not envs.SGLANG_ENABLE_POST_CAPTURE_KV_SIZING.get():
return False
if cfg.device != "cuda":
return False
if cfg.dcp_size != 1:
return False
if mla_enabled:
return False
if cfg.kv_cache_dtype == "fp4_e2m1":
return False
if cfg.prefill_only_disable_kv_cache:
return False
if cfg.enable_memory_saver:
return False
if envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get() is not None:
return False
if (
cfg.disaggregation_mode != "prefill"
and cfg.cuda_graph_config.decode.backend == Backend.DISABLED
):
return False
if cfg.disaggregation_mode != "decode":
prefill_cfg = cfg.cuda_graph_config.prefill
# We can only skip eager activation headroom when the largest
# prefill forward batch size is already graph-captured. Otherwise,
# an eager forward will need more memory and lead to OOM.
if (
prefill_cfg.backend == Backend.DISABLED
or cfg.chunked_prefill_size <= 0
or self.max_prefill_buffer_tokens() > max(prefill_cfg.bs or (0,))
):
return False
from sglang.srt.configs.model_config import is_deepseek_v4, is_minimax_sparse
hf_config = self.get_model_config().hf_config
if is_deepseek_v4(hf_config) or is_minimax_sparse(hf_config):
return False
return True
def pre_capture_activation_reserve_mb(self, gpu_mem: Optional[float]) -> float:
# Runtime activation working-set reserve for eager decode above the captured
# max_bs and transient prefill/logits; also covers fixed state caches.
@@ -3870,347 +3775,13 @@ class ServerArgs:
reserved_mem = max(reserved_mem, 10 * 1024)
return reserved_mem
def reserve_for_graph_mb(self) -> float:
cfg = resolving_view(self)
decode_cuda_graph_config = cfg.cuda_graph_config.decode
prefill_cuda_graph_config = cfg.cuda_graph_config.prefill
reserved_mem = 0.0
if (
cfg.disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED
):
reserved_mem += decode_cuda_graph_config.max_bs * 2
if (
self._resolved().enable_dp_attention
and cfg.disaggregation_mode != "prefill"
):
# DP attention needs more padding for some operations, and much more for large
# cuda graph max bs (torch allocator / implementation inefficiencies).
reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 3
if decode_cuda_graph_config.max_bs > 300:
reserved_mem += decode_cuda_graph_config.max_bs * cfg.dp_size * 1.5
if (
cfg.disaggregation_mode != "decode"
and prefill_cuda_graph_config.backend != Backend.DISABLED
):
if not self.use_mla_backend():
# Only non-torch memory is counted; torch memory is reused by cuda graph capture.
reserved_mem += len(prefill_cuda_graph_config.bs) * 8
else:
# MLA backend overhead is much higher than expected with fa3.
reserved_mem += 1.5 * 1024
if (
prefill_cuda_graph_config.backend == Backend.BREAKABLE
and resolved_view(self).moe_a2a_backend == "deepep"
):
# Prefill-BCG DeepEP delta (bridge pool + NVL first-touch
# during capture); decode-side DeepEP is a baseline cost.
reserved_mem += 1 * 1024
return reserved_mem
def reserve_for_deepep_a2a_mb(self) -> float:
# DeepEP all-to-all buffers captured in the decode graph are real extra
# allocations, reserved on top of the floor.
cfg = resolving_view(self)
decode_cuda_graph_config = cfg.cuda_graph_config.decode
if (
cfg.disaggregation_mode != "prefill"
and decode_cuda_graph_config.backend != Backend.DISABLED
and resolved_view(self).moe_a2a_backend == "deepep"
):
return 2 * 1024
return 0.0
def _generate_decode_cuda_graph_batch_sizes(self, max_bs: int):
"""
Generate the list of batch sizes for CUDA graph capture based on max_bs.
This integrates the logic from cuda_graph_runner.py.
"""
cfg = resolving_view(self)
# Handle disable_cuda_graph_padding as the first condition for both spec and non-spec
if cfg.disable_cuda_graph_padding:
capture_bs = list(range(1, max_bs + 1))
elif cfg.speculative_algorithm is None:
# Normal case:
capture_bs = (
[1, 2, 4, 8, 12]
+ list(range(16, 257, 8))
+ list(range(272, 512, 16))
+ list(range(512, max_bs + 1, 32))
)
else:
# Spec decoding case: less padding for smaller batch sizes
capture_bs = (
list(range(1, 9, 1))
+ list(range(10, 33, 2))
+ list(range(40, 65, 4))
+ list(range(72, 257, 8))
+ list(range(272, max_bs + 1, 16))
)
capture_bs = [bs for bs in capture_bs if bs <= max_bs]
if max_bs not in capture_bs:
capture_bs.append(max_bs)
return capture_bs
def _generate_cpu_graph_batch_sizes(self):
"""
Generate the list of batch sizes for CPU graph capture based on torch_compile_max_bs.
"""
cfg = resolving_view(self)
if cfg.disable_cuda_graph_padding:
capture_bs = list(range(1, cfg.torch_compile_max_bs + 1))
else:
capture_bs = sorted(
set().union(
range(1, 17),
range(18, 31, 2),
range(32, 81, 4),
range(84, cfg.torch_compile_max_bs + 1, 8),
{cfg.torch_compile_max_bs},
)
)
capture_bs = [bs for bs in capture_bs if bs <= cfg.torch_compile_max_bs]
return capture_bs
def _generate_prefill_cuda_graph_batch_sizes(self, max_bs: int):
"""
Generate the list of batch sizes for prefill CUDA graph capture
based on max_bs. For tc_piecewise prefill, bs carries the
captured token count (one shape knob per phase).
"""
capture_sizes = (
list(range(4, 33, 4))
+ list(range(48, 257, 16))
+ list(range(288, 513, 32))
+ list(range(576, 1024 + 1, 64))
+ list(range(1280, 4096 + 1, 256))
+ list(range(4608, max_bs + 1, 512))
)
capture_sizes = [s for s in capture_sizes if s <= max_bs]
return capture_sizes
def _set_default_dsa_kv_cache_dtype(self, major: int, quantization: str) -> None:
# Moved to the resolution pipeline (arg_groups/overrides.py:
# _dsa_kv_cache_dtype_default), invoked here at its legacy slot.
from sglang.srt.arg_groups.overrides import (
_dsa_kv_cache_dtype_default,
run_post_process_pass,
)
run_post_process_pass(self, _dsa_kv_cache_dtype_default)
def _set_default_dsa_backends(self, major: int) -> None:
# Moved to the resolution pipeline (arg_groups/overrides.py:
# _dsa_split_backend_resolution), invoked here at its legacy slot.
from sglang.srt.arg_groups.overrides import (
_dsa_split_backend_resolution,
run_post_process_pass,
)
run_post_process_pass(self, _dsa_split_backend_resolution)
def _support_mamba_cache_extra_buffer(self, model_arch: str):
from sglang.srt.arg_groups.overrides import supports_mamba_cache_extra_buffer
return supports_mamba_cache_extra_buffer(self, model_arch)
def _get_default_attn_backend(self, use_mla_backend: bool, model_config):
"""
Auto select the fastest attention backend.
1. Models with MHA Architecture (e.g: Llama, QWen)
1.1 We will turn on FA3 on hopper unless user use spec decode with topk > 1 or page_size > 1.
1.2 Use trtllm_mha for SM100/SM103 (Blackwell B200/GB200/B300) excluding spec with topk > 1.
Note: trtllm_mha does not support SM120, which will fall back to flashinfer.
1.3 In other cases, we will use flashinfer if available, otherwise use triton.
2. Models with MLA Architecture and using FA3
2.1 We will use FA3 backend on hopper.
2.2 We will use Flashinfer backend on blackwell.
2.3 Otherwise, we will use triton backend.
"""
cfg = resolving_view(self)
# OOT platforms provide their own default attention backend.
if current_platform.is_out_of_tree():
return current_platform.get_default_attention_backend()
# Whisper requires flashinfer for cross-attention CUDA graph support.
if "WhisperForConditionalGeneration" in (
model_config.hf_config.architectures or []
):
return "flashinfer"
if not use_mla_backend:
# MHA architecture
if is_hopper_with_cuda_12_3() and is_no_spec_infer_or_topk_one(
resolved_view(self)
):
# Note: flashinfer 0.6.1 caused performance regression on Hopper attention kernel
# Before the kernel is fixed, we choose fa3 as the default backend on Hopper MHA
# ref: https://github.com/sgl-project/sglang/issues/17411
return "fa3"
elif (
is_sm100_supported()
and is_no_spec_infer_or_topk_one(resolved_view(self))
and (
cfg.speculative_algorithm is None
or cfg.speculative_eagle_topk is not None
)
):
# trtllm_mha requires equal K/V row widths; fa4 carries
# v_head_dim through.
if model_config.has_asymmetric_kv:
return "fa4"
return "trtllm_mha"
elif is_hip():
return "aiter"
elif is_mps():
return "torch_native"
else:
# FlashInfer does not support attention sinks.
if is_flashinfer_available() and not model_config.has_attention_sinks:
return "flashinfer"
return "triton"
else:
# MLA architecture
if is_hopper_with_cuda_12_3():
return "fa3"
elif is_sm100_supported():
return "flashinfer"
elif is_hip():
head_num = model_config.get_num_kv_heads(self.tp_size)
# TODO current aiter only support head number 16 or 128 head number
if head_num == 128 or head_num == 16:
return "aiter"
else:
return "triton"
elif is_mps():
return "torch_native"
else:
return "triton"
def cutedsl_moe_max_num_tokens(self) -> int:
"""Largest number of tokens a single forward routes through a CuteDSL
MoE layer on one (DP) rank. Single source of truth for both the
standard-allgather wrapper buffers and the FlashInfer A2A dispatcher
budget. Max over the prefill (max_prefill_tokens), piecewise-prefill
capture, and decode/verify bounds; num_tokens_per_req is
speculative_num_draft_tokens under speculative decoding, else 1.
"""
cfg = resolving_view(self)
if cfg.speculative_algorithm:
num_tokens_per_req = cfg.speculative_num_draft_tokens or 1
else:
num_tokens_per_req = 1
prefill_tokens = cfg.max_prefill_tokens
cg_config = cfg.cuda_graph_config
if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE:
prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0)
decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0
decode_tokens = decode_max_bs * num_tokens_per_req
return max(prefill_tokens, decode_tokens)
def max_prefill_buffer_tokens(self) -> int:
"""Prefill-buffer ceiling: chunked_prefill_size, except PP dynamic
chunking can grow chunks toward max_prefill_tokens and probe at 1.25x."""
cfg = resolving_view(self)
chunked = (
cfg.chunked_prefill_size
if cfg.chunked_prefill_size and cfg.chunked_prefill_size > 0
else 0
)
tokens = chunked
if cfg.enable_dynamic_chunking and cfg.pp_size > 1 and chunked:
tokens = max(tokens, cfg.max_prefill_tokens or 0, math.ceil(chunked * 1.25))
return tokens
def _required_mori_dispatch_tokens_per_rank(self) -> int:
"""Max tokens a single rank dispatches through MoRI in one forward."""
cfg = resolving_view(self)
return cfg.chunked_prefill_size
def _required_pplx_dispatch_tokens_per_rank(self) -> int:
"""Max tokens a single rank dispatches through pplx in one forward."""
cfg = resolving_view(self)
required = cfg.chunked_prefill_size
if cfg.cuda_graph_max_bs_decode is not None:
required = max(required, cfg.cuda_graph_max_bs_decode)
return required
# ===== END TO BE REFACTORED ====
def _is_mistral_native_format(self) -> bool:
"""True iff the checkpoint requires load_format=mistral.
Looks for consolidated*.safetensors with no competing
model*.safetensors; when both weight formats ship in the
same checkpoint (e.g. Mistral-7B-Instruct-v0.3) the HF path is
preferred to avoid loading Mistral-named weights into an
HF-named architecture.
Name override: mistral-large-3 / mistral-small-4 /
leanstral always treat as Mistral-native when params.json
is present -- those families need Mistral weight loading
regardless of which weight files happen to be present.
"""
cfg = resolving_view(self)
_MISTRAL_NATIVE_PATTERNS = (
"mistral-large-3",
"mistral-small-4",
"leanstral",
)
name_matches = any(
p in str(cfg.model_path).lower() for p in _MISTRAL_NATIVE_PATTERNS
)
def _check_format(has_params, has_consolidated, has_hf_weights) -> bool:
if has_params and name_matches:
return True
return has_consolidated and not has_hf_weights
if os.path.isdir(cfg.model_path):
return _check_format(
has_params=os.path.exists(os.path.join(cfg.model_path, "params.json")),
has_consolidated=bool(
glob.glob(os.path.join(cfg.model_path, "consolidated*.safetensors"))
),
has_hf_weights=bool(
glob.glob(os.path.join(cfg.model_path, "model*.safetensors"))
),
)
try:
from huggingface_hub import HfApi
files = {s.rfilename for s in HfApi().model_info(cfg.model_path).siblings}
return _check_format(
has_params="params.json" in files,
has_consolidated=any(
f.startswith("consolidated") and f.endswith(".safetensors")
for f in files
),
has_hf_weights=any(
f.startswith("model")
and f.endswith(".safetensors")
and "/" not in f
for f in files
),
)
except Exception:
return False
LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",)
# The strided-layout Triton requirement is enforced via
@@ -4587,52 +4158,6 @@ class ServerArgs:
return False
return True
def get_model_config(self):
# Lazy init to avoid circular import
cfg = resolving_view(self)
from sglang.srt.configs.model_config import ModelConfig
memo = getattr(self, "_model_config", None)
if memo is not None:
# The key is the path this record carried when the cache was
# filled. The GGUF and ModelScope handlers declare a different
# `model_path`, and a configuration built before them describes
# another checkpoint. `ModelConfig` re-points its own `model_path`
# at the local pull directory when the weights sit behind an
# object-store URI, so its field is not the key. A configuration a
# fixture supplied carries no key and is handed back as it is.
built_from = getattr(self, "_model_config_built_from", None)
if built_from is None or built_from == cfg.model_path:
return memo
model_config = ModelConfig.from_server_args(self)
self._model_config = model_config
self._model_config_built_from = cfg.model_path
if model_config.is_hybrid_swa:
logger.info(
"Hybrid SWA model detected. architectures=%s",
model_config.hf_config.architectures,
)
return model_config
def _resolved(self):
"""Read-only view of the resolving configuration: declared fields
resolve from the declaration stash."""
return resolved_view(self)
def _late_resolution(self, source: str, **fields) -> None:
"""Resolve fields at the launcher's validation stage (pre-publish).
See ``arg_groups.overrides.declare_late_resolution``: the decision goes
to this instance's declaration stash, so every holder of it carries the
decision and publishes bags that answer with it. Refused outright once
the config is published.
"""
from sglang.srt.arg_groups.overrides import declare_late_resolution
declare_late_resolution(self, source, **fields)
def __setattr__(self, name, value):
# Once resolution has finished the record is the READ-ONLY raw input
# the config bags were projected from. Resolved config changes go to the bags via
@@ -4651,153 +4176,17 @@ class ServerArgs:
)
object.__setattr__(self, name, value)
def _resolved_attention_backends(self):
"""Mid-resolution (prefill, decode) backends: reads through the pass
view so declared fields resolve from the declaration stash."""
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
)
return attention_backends_of(resolved_view(self))
def get_attention_backends(self):
"""The (prefill, decode) pair resolution decided.
Reads through the declaration stash, not the fields: the model-specific
overrides declare into the stash without writing the fields, so a field
read answers with what the operator typed.
"""
return attention_backends_of(resolved_view(self))
def use_mla_backend(self):
from sglang.srt.configs.model_config import AttentionArch
model_config = self.get_model_config()
return model_config.attention_arch == AttentionArch.MLA
def is_attention_backend_not_set(self):
cfg = resolving_view(self)
return (
cfg.attention_backend is None
and cfg.prefill_attention_backend is None
and cfg.decode_attention_backend is None
)
def enable_mamba_extra_buffer(self) -> bool:
return mamba_extra_buffer_of(resolving_view(self))
def enable_mamba_extra_buffer_lazy(self) -> bool:
return mamba_extra_buffer_lazy_of(resolving_view(self))
@property
def max_speculative_num_draft_tokens(self) -> Optional[int]:
"""Return the maximum draft-token count speculative decoding may use.
Memoized only once the record is resolved: an answer computed off a raw
record describes inputs resolution is about to rewrite (auto speculative
sizing fills `speculative_num_draft_tokens` in), and a cache filled that
early would keep answering with it.
"""
cfg = resolving_view(self)
memo = self.__dict__.get("_max_speculative_num_draft_tokens")
if memo is not None:
return memo
if cfg.speculative_num_draft_tokens is None:
result = None
elif not cfg.speculative_adaptive:
result = cfg.speculative_num_draft_tokens
else:
from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config,
)
candidate_steps = resolve_candidate_steps_from_config(
cfg_path=cfg.speculative_adaptive_config,
)
# TODO: adaptive spec currently requires topk=1, so each runtime
# state needs steps + 1 draft-token slots. Revisit this if topk>1
# is supported.
result = max(candidate_steps) + 1
if getattr(self, "_resolution_finished", False):
self._max_speculative_num_draft_tokens = result
return result
@property
def mamba_cache_chunk_size(self) -> int:
# For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE
# (or mamba_chunk_size if it is defined in the model's config) and page_size.
# It is used to determine the caching point in a sequence during prefill.
# A pre-seeded `_mamba_cache_chunk_size` (fixtures supply one so a dummy
# model never loads an HF config) is honored as-is; otherwise the memo
# is only kept once the record is resolved, because `page_size` below
# is resolution-written.
if not hasattr(self, "_mamba_cache_chunk_size"):
try:
from sglang.kernels.ops.attention.fla.chunk_delta_h import (
CHUNK_SIZE as FLA_CHUNK_SIZE,
)
except ImportError:
# Must match sglang.kernels.ops.attention.fla.chunk_delta_h.CHUNK_SIZE
FLA_CHUNK_SIZE = 64
hf_config = self.get_model_config().hf_config
chunk_size = getattr(hf_config, "mamba_chunk_size", FLA_CHUNK_SIZE)
page_size = resolved_view(self).page_size
assert (
max(chunk_size, page_size) % min(chunk_size, page_size) == 0
), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}"
if not getattr(self, "_resolution_finished", False):
return max(chunk_size, page_size)
self._mamba_cache_chunk_size = max(chunk_size, page_size)
return self._mamba_cache_chunk_size
def check_server_args(self):
from sglang.srt.arg_groups.validation_hook import check_server_args
check_server_args(self)
def adjust_mem_fraction_for_vlm(self, model_config):
cfg = resolving_view(self)
vision_config = getattr(model_config.hf_config, "vision_config", None)
if vision_config is None:
return
# roughly reduce the mem_fraction_static base on params of Vit
original_server_arg_mem_fraction = cfg.mem_fraction_static
# a base mem_fraction_static factor for regular Vit
base_mem_fraction_reduction_ratio = 0.95
vit_num_layers = getattr(vision_config, "num_hidden_layers", 24)
vit_hidden_size = getattr(vision_config, "hidden_size", 1024)
# baseline ViT params (ViT-L/14)
baseline_vit_layers = 24
baseline_vit_hidden_size = 1024
# weight params count
current_complexity_score = vit_num_layers * (vit_hidden_size**2)
baseline_complexity_score = baseline_vit_layers * (baseline_vit_hidden_size**2)
complexity_ratio = (
current_complexity_score / baseline_complexity_score
if baseline_complexity_score > 0
else 1.0
)
# every time the complexity grows 100%, adjust final factor for 10%
sensitivity_scale = 0.1
dynamic_adjustment_factor = 1.0 - sensitivity_scale * (complexity_ratio - 1.0)
dynamic_adjustment_factor = max(0.8, min(1.05, dynamic_adjustment_factor))
final_overall_factor = (
base_mem_fraction_reduction_ratio * dynamic_adjustment_factor
)
self._declare(
"adjust_mem_fraction_for_vlm",
mem_fraction_static=original_server_arg_mem_fraction * final_overall_factor,
)
@property
def _parsed_modelexpress_config(self) -> dict:
cache = getattr(self, "_mx_config_cache", None)
@@ -4952,10 +4341,6 @@ class ServerArgs:
descriptor["load_topic"] = LOAD_TOPIC
return descriptor
def should_report_expert_balancedness(self) -> bool:
cfg = resolving_view(self)
return cfg.expert_balancedness_report_mode != "off"
def should_log_expert_balancedness_to_server_log(self) -> bool:
cfg = resolving_view(self)
+11 -4
View File
@@ -2088,20 +2088,27 @@ def _wait_for_gpu_idle_in_ci(
pass
# Names the runner kits stamp onto a record that are not members of it.
# `ModelRunner` computes `use_mla_backend` on itself; the kits copy that bool
# onto the record they hand the runner, and `hasattr` cannot see it.
_RUNNER_WRITTEN_NAMES = frozenset({"use_mla_backend"})
def server_args_variant(server_args, **fields):
"""A modified deep copy of a config, for a test double whose fixture
differs from the (possibly published, read-only) config it starts from.
The receiver is untouched; the copy keeps its read-only guard.
A name may also shadow a method with a fixture value (the runner kits set
``use_mla_backend``, a method ModelRunner itself overwrites at init);
names that exist nowhere on the class fail loudly."""
A name may also be one the kits stamp on rather than a field (see
``_RUNNER_WRITTEN_NAMES``); names that exist nowhere fail loudly."""
variant = copy.deepcopy(server_args)
cls = type(variant)
unknown = {
name
for name in fields
if name not in cls.__dataclass_fields__ and not hasattr(cls, name)
if name not in cls.__dataclass_fields__
and not hasattr(cls, name)
and name not in _RUNNER_WRITTEN_NAMES
}
if unknown:
raise ValueError(f"unknown ServerArgs field(s): {sorted(unknown)}")
@@ -58,12 +58,12 @@ class TestServerArgsIBDeviceValidation(unittest.TestCase):
real_listdir = os.listdir
with patch(
"sglang.srt.server_args.os.path.isdir",
"sglang.srt.arg_groups.validation_hook.os.path.isdir",
side_effect=lambda path: (
True if path == "/sys/class/infiniband" else real_isdir(path)
),
), patch(
"sglang.srt.server_args.os.listdir",
"sglang.srt.arg_groups.validation_hook.os.listdir",
side_effect=lambda path: (
available_devices
if path == "/sys/class/infiniband"
@@ -88,8 +88,9 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
"sglang.srt.arg_groups.cuda_graph_hook"
".disable_tc_piecewise_cudagraph_if_incompatible"
) as disable_if_incompatible,
patch.object(
args, "_resolved_attention_backends", return_value=("fa3", "fa3")
patch(
"sglang.srt.arg_groups.overrides.attention_backends_of",
return_value=("fa3", "fa3"),
),
):
apply_cuda_graph_compatibility(args)
@@ -116,12 +117,11 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
args._cuda_graph_config_locked = set()
with (
patch.object(
args,
"_resolved_attention_backends",
patch(
"sglang.srt.arg_groups.overrides.attention_backends_of",
return_value=("trtllm_mla", "trtllm_mla"),
),
patch.object(args, "use_mla_backend", return_value=True),
patch("sglang.srt.arg_groups.overrides.use_mla_backend", return_value=True),
):
apply_cuda_graph_compatibility(args)
@@ -137,9 +137,8 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
)
args._cuda_graph_config_locked = {(Phase.PREFILL, "backend")}
with patch.object(
args,
"_resolved_attention_backends",
with patch(
"sglang.srt.arg_groups.overrides.attention_backends_of",
return_value=("trtllm_mla", "trtllm_mla"),
):
apply_cuda_graph_compatibility(args)
@@ -185,10 +184,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
args.disable_radix_cache = False
args.chunked_prefill_size = 2048
with (
patch.object(args, "get_model_config", return_value=args._model_config),
patch("sglang.srt.arg_groups.model_hook.is_cuda", return_value=True),
):
with (patch("sglang.srt.arg_groups.model_hook.is_cuda", return_value=True),):
handle_model_capability_adjustments(args)
self.assertTrue(resolution_result(args, "disable_radix_cache"))
@@ -215,7 +211,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
hf_config=SimpleNamespace(architectures=["BertModel"]),
)
with patch.object(args, "get_model_config", return_value=args._model_config):
if True: # the record already carries the seeded configuration
handle_model_capability_adjustments(args)
self.assertTrue(resolution_result(args, "is_embedding"))
@@ -87,7 +87,10 @@ from sglang.srt.mem_cache.unified_radix_cache import (
_OngoingPrefetch,
_OngoingWriteThrough,
)
from sglang.srt.runtime_context import get_server_args, get_serving
from sglang.srt.runtime_context import (
get_serving,
mamba_cache_chunk_size,
)
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import (
ServerArgs,
@@ -5225,7 +5228,7 @@ class UnifiedRadixCacheSuite:
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
cache, allocator, req_to_token_pool = build_fixture(self.cfg)
chunk_size = get_server_args().mamba_cache_chunk_size
chunk_size = mamba_cache_chunk_size()
tokens = self._make_seq(1, chunk_size + 1)
self._insert(cache, allocator, req_to_token_pool, tokens)
leaf = cache.match_prefix(
@@ -5258,7 +5261,7 @@ class UnifiedRadixCacheSuite:
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
cache, allocator, req_to_token_pool = self._build_hicache_fixture()
chunk_size = get_server_args().mamba_cache_chunk_size
chunk_size = mamba_cache_chunk_size()
prefix = self._make_seq(1, chunk_size)
tokens = prefix + self._make_seq(1000, chunk_size + 1)
self._insert(cache, allocator, req_to_token_pool, prefix)
@@ -5283,7 +5286,7 @@ class UnifiedRadixCacheSuite:
if not self.cfg.has_mamba or self.cfg.has_swa or self.cfg.page_size != 1:
self.skipTest("requires page_size=1 Full+Mamba")
cache, allocator, req_to_token_pool = self._build_hicache_fixture()
chunk_size = get_server_args().mamba_cache_chunk_size
chunk_size = mamba_cache_chunk_size()
prefix = self._make_seq(1, chunk_size)
tokens = prefix + self._make_seq(1000, chunk_size + 1)
self._insert(cache, allocator, req_to_token_pool, prefix)
@@ -50,10 +50,7 @@ def test_packed_speculative_extend_is_limited_to_pd_prefill_target(mode, error):
def test_chunked_prefill_disabled_uses_legacy_token_ceiling():
model_runner = SimpleNamespace(
server_args=SimpleNamespace(
max_prefill_buffer_tokens=Mock(return_value=0),
max_prefill_tokens=32768,
),
server_args=SimpleNamespace(),
is_generation=True,
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_speculative=lambda: False),
@@ -76,6 +73,12 @@ def test_chunked_prefill_disabled_uses_legacy_token_ceiling():
"get_disagg",
return_value=SimpleNamespace(disaggregation_mode="prefill"),
),
patch.object(flashinfer_autotune, "max_prefill_buffer_tokens", return_value=0),
patch.object(
flashinfer_autotune,
"get_schedule",
return_value=SimpleNamespace(max_prefill_tokens=32768),
),
patch.object(flashinfer_autotune, "run_flashinfer_autotune_forward"),
patch.object(flashinfer_autotune.torch.cuda, "empty_cache"),
):
@@ -1,4 +1,4 @@
"""`get_model_config()` caches, and the key is the path the record carried.
"""`model_config_of()` caches, and the key is the path the record carried.
Two movements of a `model_path` reach this cache, and only the first one means
the cached configuration describes the wrong checkpoint:
@@ -20,6 +20,7 @@ import tempfile
import unittest
from types import SimpleNamespace
from sglang.srt.arg_groups.overrides import declare_resolution, model_config_of
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.environ import EnvField, envs
from sglang.srt.server_args import ServerArgs
@@ -117,7 +118,7 @@ class TestTheModelConfigCache(CustomTestCase):
self.assertEqual(server_args.model_path, _OBJECT_STORE_URI)
self.assertEqual(cached.model_path, pulled)
self.assertIs(server_args.get_model_config(), cached)
self.assertIs(model_config_of(server_args), cached)
def test_a_declared_model_path_rebuilds_the_configuration(self):
"""The GGUF and ModelScope shape: the record's own path moved."""
@@ -125,14 +126,15 @@ class TestTheModelConfigCache(CustomTestCase):
second_checkpoint = self._checkpoint()
server_args = ServerArgs(model_path=first_checkpoint, device="cuda")
first = server_args.get_model_config()
first = model_config_of(server_args)
self.assertEqual(first.model_path, first_checkpoint)
server_args._declare(
declare_resolution(
server_args,
"test_a_declared_model_path_rebuilds_the_configuration",
model_path=second_checkpoint,
)
second = server_args.get_model_config()
second = model_config_of(server_args)
self.assertIsNot(second, first)
self.assertEqual(second.model_path, second_checkpoint)
@@ -151,11 +153,11 @@ class TestTheModelConfigCache(CustomTestCase):
model_path=second_checkpoint,
)
rebuilt = copy_.get_model_config()
rebuilt = model_config_of(copy_)
self.assertEqual(rebuilt.model_path, second_checkpoint)
self.assertIs(copy_.get_model_config(), rebuilt)
self.assertIs(model_config_of(copy_), rebuilt)
# The parent keeps the configuration it resolved with.
self.assertEqual(server_args.get_model_config().model_path, first_checkpoint)
self.assertEqual(model_config_of(server_args).model_path, first_checkpoint)
def test_a_supplied_configuration_is_handed_back(self):
"""A configuration nothing in here built carries no key, so nothing
@@ -164,7 +166,7 @@ class TestTheModelConfigCache(CustomTestCase):
stand_in = SimpleNamespace(model_path="somewhere/else")
server_args._model_config = stand_in
self.assertIs(server_args.get_model_config(), stand_in)
self.assertIs(model_config_of(server_args), stand_in)
if __name__ == "__main__":
@@ -30,7 +30,7 @@ _SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt"
# Two quantities sharing one name.
_READ_BEFORE_RESOLUTION = frozenset({"is_embedding"})
# Declared after the first `get_model_config()`, so the cached configuration
# Declared after the first `model_config_of()`, so the cached configuration
# holds the earlier value. Nothing reads the stale copy today (its one consumer
# is on the `is_draft_model` branch, built after resolution), and fixing it
# means moving the build or the hook. Pinned so a second field in this position
@@ -109,7 +109,7 @@ def _registry_collection_is_after_the_build():
Handler-local ordering only -- the caller still has to compare against the
pipeline-wide first build, which sits in an *earlier* step: hoisting the
collection above this handler's own `get_model_config()` call does not move
collection above this handler's own `model_config_of()` call does not move
it above the configuration another handler already cached.
"""
handler = None
@@ -145,7 +145,7 @@ def _registry_collection_is_after_the_build():
name = func.id
else:
continue
if name == "get_model_config" and build is None:
if name == "model_config_of" and build is None:
build = node.lineno
if name == "collect_model_override_declarations" and collect is None:
collect = node.lineno
@@ -192,11 +192,12 @@ def _server_args_names(tree, path):
and value.args[0].id in names
)
# `resolved = self._resolved()` is the same view, spelled as the
# record's own member.
# resolution vocabulary.
member = (
isinstance(value, ast.Call)
and isinstance(value.func, ast.Attribute)
and value.func.attr == "_resolved"
and isinstance(value.func, ast.Name)
and value.func.id == "resolved_view"
and isinstance(value.func.value, ast.Name)
and value.func.value.id in names
)
@@ -267,7 +268,7 @@ def _late_resolution_fields():
if isinstance(node.func, ast.Attribute)
else getattr(node.func, "id", "")
)
if called in ("_late_resolution", "declare_late_resolution"):
if called == "declare_late_resolution":
fields |= {kw.arg for kw in node.keywords if kw.arg}
return fields
@@ -488,14 +489,14 @@ def _declaration_positions():
wanted = _constructor_reads()
def build_site():
"""(step index, method name, line) of the first `get_model_config()`."""
"""(step index, method name, line) of the first `model_config_of()`."""
for index, step in enumerate(steps):
for method in reached[step]:
for node in ast.walk(methods[method]):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "get_model_config"
and isinstance(node.func, ast.Name)
and node.func.id == "model_config_of"
):
return index, step, method, node.lineno
return None
@@ -532,8 +533,8 @@ def _declaration_positions():
same_body = index == build_index and method == build_method
rank = 0 if same_body and node.lineno < build_line_in_body else 1
if (
isinstance(node.func, ast.Attribute)
and node.func.attr == "_declare"
isinstance(node.func, ast.Name)
and node.func.id == "declare_resolution"
):
fields = {kw.arg for kw in node.keywords if kw.arg}
# A handler that calls an imported hook (the Kimi and DeepSeek
@@ -669,8 +670,8 @@ class TestModelConfigReadsResolvedInput(CustomTestCase):
for method in reached[step]
if any(
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "get_model_config"
and isinstance(node.func, ast.Name)
and node.func.id == "model_config_of"
for node in ast.walk(methods[method])
)
)
@@ -61,7 +61,7 @@ class TestNoPublicNonFieldSlot(CustomTestCase):
written = _self_written_attributes()
self.assertGreater(
len(written),
5,
3,
f"only {len(written)} self-writes found; the scan is broken, not the "
"record",
)
@@ -30,6 +30,7 @@ under its own default configuration.
"""
import unittest
from unittest import mock
from sglang.srt.arg_groups.kv_cache_hook import handle_page_major_kv_layout
from sglang.srt.server_args import ServerArgs
@@ -66,13 +67,18 @@ def _accepts(
"mamba_backend": "triton",
}.items():
object.__setattr__(sa, name, value)
sa.use_mla_backend = lambda: use_mla
sa._resolved_attention_backends = lambda: [backend]
try:
handle_page_major_kv_layout(sa)
return True
except AssertionError:
return False
# `use_mla_backend` asks the model configuration, which this stand-in has
# no room for; the case under test is what the handler does with the answer.
# The handler imports it inside the function, so the source module is
# where the patch has to go.
with mock.patch(
"sglang.srt.arg_groups.overrides.use_mla_backend", return_value=use_mla
):
try:
handle_page_major_kv_layout(sa)
return True
except AssertionError:
return False
class TestPageMajorBackendAllowlist(unittest.TestCase):
@@ -24,7 +24,6 @@ import unittest
import unittest.mock
import sglang
from sglang.srt import server_args as server_args_module
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
@@ -147,7 +146,7 @@ def _late_resolvers():
if isinstance(node.func, ast.Attribute)
else getattr(node.func, "id", None)
)
if called in ("declare_late_resolution", "_late_resolution"):
if called == "declare_late_resolution":
return True
if called and reaches(called, seen):
return True
@@ -887,7 +886,9 @@ class TestResolutionDeclarations(CustomTestCase):
# The pipeline asks the platform other questions on the way through
# (whether it is out of tree, whether it supports piecewise capture),
# and which of those it reaches depends on the host.
class _Plugin(type(server_args_module.current_platform)):
from sglang.srt.platforms import current_platform
class _Plugin(type(current_platform)):
device_name = "oot"
def apply_server_args_defaults(self, server_args):
@@ -36,7 +36,7 @@ import unittest.mock
import torch
import sglang
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.overrides import model_config_of, resolution_result
from sglang.srt.environ import EnvField, envs
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda
@@ -519,7 +519,7 @@ class TestProgramsResolveBeforeReadingResolution(CustomTestCase):
from sglang.srt.server_args import ServerArgs as _ServerArgs
srt = pathlib.Path(next(iter(sglang.__path__))).resolve() / "srt"
declarers = {"_declare", "declare_resolution", "declare_late_resolution"}
declarers = {"declare_resolution", "declare_late_resolution"}
fields = set()
field_names = {field.name for field in _dataclasses.fields(_ServerArgs)}
# The record plus every module under `arg_groups/`: a handler declares
@@ -826,10 +826,10 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
def test_the_copy_carries_what_resolution_left_on_the_record(self):
"""Not just the stash and the flag.
`get_model_config()` memoizes on the record, and that cache is filled
`model_config_of()` memoizes on the record, and that cache is filled
during resolution. A copy that is marked resolved but arrives without it
cannot fill it -- the read-only guard refuses the cache write -- so the
first `get_model_config()` raises. That is what killed the Ray
first `model_config_of()` raises. That is what killed the Ray
schedulers, and it is why the carry is enumerated from the instance
rather than from a list of names.
"""
@@ -846,7 +846,7 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
[],
f"the copy did not carry what resolution left on the record: {missing}",
)
self.assertIsNotNone(copy_.get_model_config())
self.assertIsNotNone(model_config_of(copy_))
# Containers are copied, so the copy's declaration stays with it.
self.assertEqual(
len(parent._resolved_overrides) + 1, len(copy_._resolved_overrides)
@@ -4,8 +4,8 @@
nothing. The fields keep what the caller passed, so a resolver that reads a
field another resolver may have decided reads the raw input -- silently, and
only on the configurations where that other resolver fires. The whole pipeline
therefore reads through `resolving_view` (or `ServerArgs._resolved()`, which is
the same view spelled as the record's own member), and this pins that there is
therefore reads through `resolving_view` (or `resolved_view`, which is
the same view after resolution has finished), and this pins that there is
nothing left reading a field directly.
Subjects: every function in `arg_groups/` that takes a config, every
@@ -83,7 +83,6 @@ def _field_reads(fn, holders):
_DECLARERS = frozenset(
{
"_declare",
"declare_resolution",
"declare_late_resolution",
"declare_direct_writes",
@@ -472,7 +471,7 @@ class TestResolutionReadsTheDeclarations(CustomTestCase):
members = _record_members()
# The floor is here to catch the scan collapsing, not to pin the
# class's size.
self.assertGreater(len(members), 40, f"only {len(members)} members were found")
self.assertGreater(len(members), 25, f"only {len(members)} members were found")
offenders = []
for name, fn in sorted(members.items()):
holders = _holders(fn) | {"self"}
@@ -37,7 +37,10 @@ from sglang.srt.arg_groups.moe_hook import (
validate_deepep_v2_dispatch_token_budget,
validate_deepep_v2_speculative_draft,
)
from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.overrides import (
cutedsl_moe_max_num_tokens,
resolution_result,
)
from sglang.srt.arg_groups.parallel_hook import (
handle_context_parallelism,
handle_data_parallelism,
@@ -825,9 +828,7 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
)
defaults.update(kw)
view = ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
with (
patch("sglang.srt.configs.model_config.is_deepseek_dsa", return_value=True),
@@ -845,21 +846,21 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
),
}
@patch("sglang.srt.server_args.is_hip", return_value=False)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_sparse_on_cuda_bfloat16(self, _mock_is_hip):
resolved = self._resolve("bfloat16")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_sparse")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_sparse")
@patch("sglang.srt.server_args.is_hip", return_value=False)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=False)
def test_hisparse_defaults_to_flashmla_kv_on_cuda_fp8(self, _mock_is_hip):
resolved = self._resolve("fp8_e4m3")
self.assertEqual(resolved["dsa_prefill_backend"], "flashmla_kv")
self.assertEqual(resolved["dsa_decode_backend"], "flashmla_kv")
@patch("sglang.srt.server_args.is_hip", return_value=False)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=False)
def test_hisparse_accepts_flashinfer_sparse_mla_on_cuda_fp8(self, _mock_is_hip):
"""SM120 GLM DSA resolves both DSA backends to flashinfer_sparse_mla, so
the fp8 hisparse allow-set must admit it or --enable-hisparse cannot
@@ -876,14 +877,14 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
validate_hisparse_dsa_backend(server_args, "dsa_prefill_backend", "prefill")
validate_hisparse_dsa_backend(server_args, "dsa_decode_backend", "decode")
@patch("sglang.srt.server_args.is_hip", return_value=True)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=True)
def test_hisparse_defaults_to_tilelang_on_rocm(self, _mock_is_hip):
resolved = self._resolve("bfloat16")
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=True)
def test_hisparse_preserves_rocm_user_backend_and_defaults_missing_side(
self, _mock_is_hip
):
@@ -892,7 +893,7 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
self.assertEqual(resolved["dsa_prefill_backend"], "tilelang")
self.assertEqual(resolved["dsa_decode_backend"], "tilelang")
@patch("sglang.srt.server_args.is_hip", return_value=True)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=True)
def test_hisparse_accepts_aiter_backend_on_rocm(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
@@ -905,7 +906,7 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
validate_hisparse_dsa_backend(server_args, "dsa_prefill_backend", "prefill")
validate_hisparse_dsa_backend(server_args, "dsa_decode_backend", "decode")
@patch("sglang.srt.server_args.is_hip", return_value=True)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=True)
def test_hisparse_rejects_cuda_backend_on_rocm(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
@@ -917,7 +918,7 @@ class TestHiSparseDsaBackendPolicy(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "tilelang"):
validate_hisparse_dsa_backend(server_args, "dsa_prefill_backend", "prefill")
@patch("sglang.srt.server_args.is_hip", return_value=False)
@patch("sglang.srt.arg_groups.hisparse_hook._is_hip", return_value=False)
def test_hisparse_rejects_rocm_backend_on_cuda(self, _mock_is_hip):
server_args = ServerArgs(
model_path="dummy",
@@ -969,7 +970,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
args.prefill_attention_backend = prefill
args.decode_attention_backend = decode
args.page_size = page_size
# Short-circuit get_model_config(): the fa4 page_size branch only needs
# Short-circuit model_config_of(): the fa4 page_size branch only needs
# use_mla_backend() (mocked) and is_sm100_supported() (mocked), not a
# real model_config. Pre-set the attribute so get_model_config returns
# early without touching ModelConfig.from_server_args.
@@ -978,7 +979,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
return args
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
@patch("sglang.srt.arg_groups.overrides.use_mla_backend", return_value=False)
def test_combined_attention_backend_fa4_forces_page_size_128(
self, _mock_mla, _mock_sm100
):
@@ -993,7 +994,7 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
self.assertEqual(resolved_view(args).page_size, 128)
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@patch("sglang.srt.server_args.ServerArgs.use_mla_backend", return_value=False)
@patch("sglang.srt.arg_groups.overrides.use_mla_backend", return_value=False)
def test_explicit_prefill_fa4_forces_page_size_128(self, _mock_mla, _mock_sm100):
# `--prefill-attention-backend fa4`: the previously-covered path.
args = self._make_args(attention_backend=None, prefill="fa4", page_size=1)
@@ -1678,7 +1679,7 @@ class TestAdaptiveSpecArgs(CustomTestCase):
args.speculative_adaptive = True
args.speculative_adaptive_config = f.name
args.device = "cuda"
args.get_model_config = lambda: SimpleNamespace(
args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(
architectures=["LlamaForCausalLM"],
get_text_config=lambda: SimpleNamespace(),
@@ -1870,7 +1871,9 @@ class TestCudaGraphDisaggregationRoles(CustomTestCase):
)
with (
patch("sglang.srt.utils.is_cuda", return_value=True),
patch.object(ServerArgs, "use_mla_backend", return_value=False),
patch(
"sglang.srt.arg_groups.overrides.use_mla_backend", return_value=False
),
):
handle_cuda_graph_config(args)
return args
@@ -1943,7 +1946,9 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
)
with (
patch("sglang.srt.utils.is_cuda", return_value=True),
patch.object(ServerArgs, "use_mla_backend", return_value=False),
patch(
"sglang.srt.arg_groups.overrides.use_mla_backend", return_value=False
),
):
handle_cuda_graph_config(args)
return args
@@ -2007,7 +2012,9 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
)
with (
patch("sglang.srt.utils.is_cuda", return_value=True),
patch.object(ServerArgs, "use_mla_backend", return_value=False),
patch(
"sglang.srt.arg_groups.overrides.use_mla_backend", return_value=False
),
):
handle_cuda_graph_config(args)
return args
@@ -2096,7 +2103,7 @@ class TestCutedslMoeMaxNumTokens(CustomTestCase):
return server_args
def test_prefill_dominates_in_default_config(self):
self.assertEqual(self._args().cutedsl_moe_max_num_tokens(), 16384)
self.assertEqual(cutedsl_moe_max_num_tokens(self._args()), 16384)
def test_speculative_decoding_scales_decode_bound(self):
# decode bound 512 * 8 dominates the small prefill/piecewise bounds
@@ -2106,7 +2113,7 @@ class TestCutedslMoeMaxNumTokens(CustomTestCase):
speculative_algorithm="EAGLE",
speculative_num_draft_tokens=8,
)
self.assertEqual(args.cutedsl_moe_max_num_tokens(), 4096)
self.assertEqual(cutedsl_moe_max_num_tokens(args), 4096)
def test_piecewise_bound_excluded_when_disabled(self):
args = self._args(
@@ -2114,7 +2121,7 @@ class TestCutedslMoeMaxNumTokens(CustomTestCase):
disable_piecewise_cuda_graph=True,
cuda_graph_max_bs=64,
)
self.assertEqual(args.cutedsl_moe_max_num_tokens(), 512)
self.assertEqual(cutedsl_moe_max_num_tokens(args), 512)
class TestSamplingBackendTokenOracleEnvGate(CustomTestCase):
@@ -2466,10 +2473,9 @@ class TestDeepEPv2Args(CustomTestCase):
dp_size=8,
enable_dp_attention=True,
)
with patch.object(
ServerArgs,
"max_speculative_num_draft_tokens",
new=property(lambda _self: 16),
with patch(
"sglang.srt.arg_groups.moe_hook.max_speculative_num_draft_tokens",
return_value=16,
):
with envs.SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK.override(128):
with self.assertRaisesRegex(ValueError, "tokens/request=16"):
@@ -20,7 +20,7 @@ def _make_spec_args(device: str, algorithm: str = "EAGLE", **overrides) -> Serve
args.speculative_num_steps = 3
args.speculative_eagle_topk = 1
args.speculative_num_draft_tokens = 4
args.get_model_config = lambda: SimpleNamespace(
args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(
architectures=["LlamaForCausalLM"],
get_text_config=lambda: SimpleNamespace(),
+18 -11
View File
@@ -40,7 +40,7 @@ _OWNERS = ("server_args.py", "runtime_context.py", "arg_groups/")
# startup default wherever it is written, and `benchmark/` ships too.
_READS_SCANNED = _PACKAGE
_DECLARERS = ("_declare", "declare_resolution", "declare_late_resolution")
_DECLARERS = ("declare_resolution", "declare_late_resolution")
def _declared_by_keyword():
@@ -190,10 +190,10 @@ def _declared_by_registry_and_passes():
def _declared_by_late_resolution():
"""Keywords of `self._late_resolution(...)`, the fourth declarer spelling.
"""Keywords of `declare_late_resolution(record, ...)`, the late spelling.
It forwards `**fields` to `declare_late_resolution`, so the keywords sit at
its call sites and a scan for the declarer's own name finds none of them.
The fields sit at the call sites rather than in the declarer, so a scan
that only knew the declarer's own definition would find none of them.
"""
# The record plus `arg_groups/`: a hook calls it on the record it was
# handed, so scanning the record's file alone finds nothing.
@@ -203,8 +203,8 @@ def _declared_by_late_resolution():
for node in ast.walk(ast.parse(source.read_text(encoding="utf-8-sig"))):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "_late_resolution"
and isinstance(node.func, ast.Name)
and node.func.id == "declare_late_resolution"
):
fields |= {keyword.arg for keyword in node.keywords if keyword.arg}
return fields
@@ -540,13 +540,20 @@ class TestNoChainReadsOfResolvedConfig(CustomTestCase):
len(by_late),
3,
f"only {len(by_late)} fields are declared late; the "
"`_late_resolution` keyword scan broke",
"`declare_late_resolution` keyword scan broke",
)
# The three mechanisms are not the same set: if any became a subset of
# the keyword scan, that scan would be doing all the work and a
# regression in the others would be invisible.
# The data channel is not the keyword scan's subset: if it became one,
# that scan would be doing all the work and a regression here would be
# invisible. The late channel *is* a subset, and deliberately so --
# `declare_late_resolution` is a keyword declarer like the others now
# that the record hosts no forwarding member, so its own floor above is
# what pins it.
self.assertTrue(by_data - by_keyword, "the data channel adds nothing")
self.assertTrue(by_late - by_keyword, "late resolution adds nothing")
self.assertTrue(
by_late <= by_keyword,
"late resolution declares outside the keyword channel; it is the "
"same spelling, so the two cannot disagree",
)
def test_nothing_reads_a_resolved_field_off_a_borrowed_record(self):
found = _chain_reads(_resolution_written())
+49 -74
View File
@@ -15,6 +15,7 @@ from types import SimpleNamespace
from typing import Optional
from unittest.mock import patch
from sglang.srt.arg_groups import attention_hook
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.overrides import (
@@ -24,6 +25,7 @@ from sglang.srt.arg_groups.overrides import (
validate_declarations,
)
from sglang.srt.configs.minicpm import MiniCPMHybridConfig
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.environ import envs
from sglang.srt.runtime_context import (
get_context,
@@ -357,14 +359,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
enable_dp_attention=enable_dp_attention,
enable_hierarchical_cache=enable_hierarchical_cache,
)
args.is_attention_backend_not_set = lambda: all(
backend is None
for backend in (
args.attention_backend,
args.prefill_attention_backend,
args.decode_attention_backend,
)
)
mixer_types = []
if sparse_attention:
mixer_types.append("minicpm4")
@@ -456,7 +450,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
disaggregation_mode="null",
enable_dp_attention=False,
enable_hierarchical_cache=False,
is_attention_backend_not_set=lambda: True,
)
config = SimpleNamespace(
has_minicpm_sparse_attention=True,
@@ -630,7 +623,11 @@ class TestGoldenModelOverrides(_IsolatedPublish):
def test_minimax_m2_sm10x_nvfp4_uses_routed_trtllm(self):
"""MiniMax-M2 NVFP4 auto must avoid the unsupported plain TRT-LLM path."""
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
# Every module that asks: the attention handler validates what the
# override family picks, and each holds its own import.
with patch.object(
overrides_module, "is_sm100_supported", return_value=True
), patch.object(attention_hook, "is_sm100_supported", return_value=True):
explicit = self._construct(
"MiniMaxM2ForCausalLM",
"llama",
@@ -765,8 +762,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
speculative_draft_attention_backend=None,
page_size=None,
mamba_radix_cache_strategy="auto",
is_attention_backend_not_set=lambda: True,
get_model_config=lambda: model_config,
_model_config=model_config,
),
hf_config,
)
@@ -971,7 +967,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
server_args.speculative_algorithm = "DFLASH"
server_args.prefill_attention_backend = "triton"
server_args.speculative_draft_attention_backend = "fa3"
server_args.is_attention_backend_not_set = lambda: False
with (
patch.object(overrides_module, "is_blackwell_supported", return_value=True),
@@ -1093,7 +1088,9 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_gpt_oss_overrides(
SimpleNamespace(
dtype="float16",
is_attention_backend_not_set=lambda: False,
attention_backend="triton",
prefill_attention_backend=None,
decode_attention_backend=None,
),
SimpleNamespace(architectures=["GptOssForCausalLM"]),
)
@@ -1313,7 +1310,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
# dual-chunk config: mismatched explicit backend raises verbatim
def _mc(dual):
return SimpleNamespace(
get_model_config=lambda: SimpleNamespace(
_model_config=SimpleNamespace(
hf_config=SimpleNamespace(dual_chunk_attention_config=dual)
),
attention_backend="fa3",
@@ -1473,9 +1470,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
swa_full_tokens_ratio=ServerArgs.swa_full_tokens_ratio,
moe_a2a_backend="none",
moe_runner_backend="auto",
get_model_config=lambda: SimpleNamespace(
is_fp4_experts=True, nvfp4_moe_meta=None
),
_model_config=SimpleNamespace(is_fp4_experts=True, nvfp4_moe_meta=None),
)
defaults.update(kw)
return SimpleNamespace(**defaults)
@@ -1527,12 +1522,10 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
# FP8 checkpoints and non-CUDA platforms keep their platform-specific
# auto-resolution paths.
fp8_model_config = lambda: SimpleNamespace(
is_fp4_experts=False, nvfp4_moe_meta=None
)
fp8_model_config = SimpleNamespace(is_fp4_experts=False, nvfp4_moe_meta=None)
self.assertNotIn(
"moe_runner_backend",
_deepseek_v4_overrides(_args(get_model_config=fp8_model_config), hf),
_deepseek_v4_overrides(_args(_model_config=fp8_model_config), hf),
)
self.assertNotIn(
"moe_runner_backend",
@@ -1569,7 +1562,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
self.assertEqual(
_deepseek_v4_overrides(
_args(
get_model_config=lambda: SimpleNamespace(
_model_config=SimpleNamespace(
is_fp4_experts=False, nvfp4_moe_meta=object()
)
),
@@ -1604,15 +1597,10 @@ class TestGoldenModelOverrides(_IsolatedPublish):
speculative_draft_attention_backend=None,
page_size=None,
mamba_radix_cache_strategy="auto",
get_model_config=lambda: mc,
_model_config=mc,
)
defaults.update(kw)
args = SimpleNamespace(**defaults)
args.is_attention_backend_not_set = lambda: (
args.attention_backend is None
and args.prefill_attention_backend is None
and args.decode_attention_backend is None
)
return args
hf = _hf()
@@ -1725,9 +1713,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
with (
@@ -1821,9 +1807,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
with (
@@ -1952,15 +1936,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
ns = SimpleNamespace(**defaults)
ns.is_attention_backend_not_set = lambda: (
ns.attention_backend is None
and ns.prefill_attention_backend is None
and ns.decode_attention_backend is None
)
ns.get_attention_backends = lambda: (
ns.prefill_attention_backend or ns.attention_backend,
ns.decode_attention_backend or ns.attention_backend,
)
return ns
# nothing set: prefill defaults to flashinfer
@@ -1991,9 +1966,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
with (
@@ -2037,9 +2010,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
defaults = dict(kv_cache_dtype="auto", device="cuda")
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
self.assertEqual(
@@ -2078,9 +2049,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
with patch.object(overrides_module, "is_hip", return_value=True):
@@ -2148,9 +2117,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
return ResolvedView(
SimpleNamespace(
get_model_config=lambda: SimpleNamespace(hf_config=hf), **defaults
)
SimpleNamespace(_model_config=SimpleNamespace(hf_config=hf), **defaults)
)
# arch guard: non-mamba arch declares nothing
@@ -2255,17 +2222,28 @@ class TestGoldenModelOverrides(_IsolatedPublish):
def _args(default_backend, **kw):
defaults = dict(
attention_backend=None,
_get_default_attn_backend=lambda **_: default_backend,
use_mla_backend=lambda: False,
get_model_config=lambda: None,
prefill_attention_backend=None,
decode_attention_backend=None,
mamba_radix_cache_strategy="auto",
disable_radix_cache=False,
speculative_algorithm=None,
)
defaults.update(kw)
return SimpleNamespace(**defaults)
args = SimpleNamespace(**defaults)
args.default_backend_for_test = default_backend
return args
with patch.object(overrides_module, "is_sm100_supported", return_value=True):
with patch.object(
overrides_module, "is_sm100_supported", return_value=True
), patch.object(
overrides_module,
"get_default_attn_backend",
lambda server_args, **_: server_args.default_backend_for_test,
), patch.object(
overrides_module, "use_mla_backend", return_value=False
), patch.object(
overrides_module, "model_config_of", return_value=None
):
# radix on + no extra buffer + no spec -> page_size=1 path
self.assertEqual(
_qwen3_5_hybrid_overrides(_args("trtllm_mha"), None),
@@ -2422,11 +2400,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
)
defaults.update(kw)
ns = SimpleNamespace(**defaults)
ns.is_attention_backend_not_set = lambda: (
ns.attention_backend is None
and ns.prefill_attention_backend is None
and ns.decode_attention_backend is None
)
return ns
hf = SimpleNamespace()
@@ -2475,6 +2448,9 @@ class TestGoldenModelOverrides(_IsolatedPublish):
prefill_attention_backend=None,
speculative_draft_attention_backend=None,
page_size=1,
# `use_mla_backend` reads the model configuration; a non-MLA
# one keeps these assertions about the page constraints.
_model_config=SimpleNamespace(attention_arch=None),
)
defaults.update(kw)
return ResolvedView(SimpleNamespace(**defaults))
@@ -2533,7 +2509,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_fa4_page_constraint(
_view(
attention_backend="fa4",
use_mla_backend=lambda: False,
speculative_eagle_topk=None,
)
),
@@ -2543,7 +2518,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_fa4_page_constraint(
_view(
attention_backend="fa4",
use_mla_backend=lambda: False,
speculative_eagle_topk=2, # EAGLE topk>1 keeps default
)
),
@@ -2554,7 +2528,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_intel_xpu_page_constraint(
_view(
decode_attention_backend="intel_xpu",
use_mla_backend=lambda: False,
)
),
{"page_size": 128},
@@ -2563,7 +2536,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_intel_xpu_page_constraint(
_view(
decode_attention_backend="intel_xpu",
use_mla_backend=lambda: True,
_model_config=SimpleNamespace(attention_arch=AttentionArch.MLA),
page_size=16, # MLA decode accepts 16
)
),
@@ -2585,7 +2558,8 @@ class TestGoldenModelOverrides(_IsolatedPublish):
defaults = dict(
device="cuda",
attention_backend=None,
is_attention_backend_not_set=lambda: True,
prefill_attention_backend=None,
decode_attention_backend=None,
# keep the (now-absorbed) quant/moe blocks inert so these
# assertions stay attention-only
moe_runner_backend="triton",
@@ -2674,7 +2648,7 @@ class TestGoldenModelOverrides(_IsolatedPublish):
_quantization_explicitly_unset=False,
moe_a2a_backend="none",
moe_runner_backend="auto",
get_model_config=lambda: SimpleNamespace(
_model_config=SimpleNamespace(
hf_config=SimpleNamespace(
architectures=[arch], quantization_config=quant_cfg
)
@@ -2742,7 +2716,6 @@ class TestGoldenModelOverrides(_IsolatedPublish):
def _args(**kw):
defaults = dict(
is_attention_backend_not_set=lambda: True,
attention_backend=None,
prefill_attention_backend=None,
decode_attention_backend=None,
@@ -2869,7 +2842,9 @@ class TestGoldenModelOverrides(_IsolatedPublish):
defaults = dict(
speculative_algorithm=None,
enable_hierarchical_cache=False,
is_attention_backend_not_set=lambda: False,
attention_backend="triton",
prefill_attention_backend=None,
decode_attention_backend=None,
)
defaults.update(kw)
return SimpleNamespace(**defaults)
+15 -3
View File
@@ -16,6 +16,18 @@ from unittest.mock import patch
import sglang as _sglang
import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups.arg_utils import NS, A, Arg
from sglang.srt.arg_groups.overrides import (
attention_backends_of,
)
from sglang.srt.arg_groups.overrides import (
mamba_cache_chunk_size as mamba_cache_chunk_size_of,
)
from sglang.srt.arg_groups.overrides import (
max_prefill_buffer_tokens as max_prefill_buffer_tokens_of,
)
from sglang.srt.arg_groups.overrides import (
resolved_view,
)
from sglang.srt.runtime_context import (
Flags,
ParallelContext,
@@ -404,7 +416,7 @@ class TestServerArgsScopedOverride(_IsolatedServerArgs):
published = (
get_context().override_server_args(_mamba_cache_chunk_size=64).install()
)
self.assertEqual(published.mamba_cache_chunk_size, 64)
self.assertEqual(mamba_cache_chunk_size_of(published), 64)
def test_installed_config_arms_the_strict_guard(self):
# The published dummy must behave like a resolved config: bare writes
@@ -1156,7 +1168,7 @@ class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs):
)
get_context().set_server_args(args)
self.assertEqual(
ServerArgs.max_prefill_buffer_tokens(args),
max_prefill_buffer_tokens_of(args),
max_prefill_buffer_tokens(),
)
@@ -1243,7 +1255,7 @@ class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs):
)
get_context().set_server_args(args)
self.assertEqual(
ServerArgs.get_attention_backends(args),
attention_backends_of(resolved_view(args)),
attention_backends(),
)
@@ -34,6 +34,7 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
import sglang
from sglang.srt.arg_groups.overrides import attention_backends_of, resolved_view
_PACKAGE_ROOT = Path(next(iter(sglang.__path__))) / "srt"
@@ -181,7 +182,7 @@ class TestSplitBackendsReachTheDecisions(CustomTestCase):
("decode_attention_backend", "flashinfer"),
):
object.__setattr__(args, name, value)
self.assertIn("flashinfer", args.get_attention_backends())
self.assertIn("flashinfer", attention_backends_of(resolved_view(args)))
def test_support_triton_is_the_regression_being_guarded(self):
from sglang.srt.utils.common import support_triton
@@ -483,8 +483,8 @@ class TestSuppliedInstanceExposure(CustomTestCase):
tgts = [node.target]
elif (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "_declare"
and isinstance(node.func, ast.Name)
and node.func.id == "declare_resolution"
):
targets |= {
kw.arg
@@ -622,7 +622,7 @@ class TestSuppliedInstanceExposure(CustomTestCase):
for path in sorted(root.rglob("*.py")):
rel = path.relative_to(root).as_posix()
source = path.read_text(encoding="utf-8-sig")
if "_late_resolution" not in source:
if "declare_late_resolution" not in source:
continue
try:
tree = ast.parse(source)