config: the resolution pipeline moves out of the record (#36789)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-28 10:17:24 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 726665e08e
commit c2928e86d7
30 changed files with 6863 additions and 5639 deletions
@@ -0,0 +1,627 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the attention backends."""
from __future__ import annotations
import logging
import os
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
from sglang.srt.utils.common import (
is_cuda,
is_hip,
is_sm90_supported,
is_sm100_or_sm110_supported,
is_sm100_supported,
is_sm120_supported,
parse_connector_type,
)
logger = logging.getLogger(__name__)
def handle_attention_backend_compatibility(server_args: Any):
cfg = resolving_view(server_args)
model_config = server_args.get_model_config()
# The attention_backend write clusters of this handler moved to the
# resolution pipeline (arg_groups/overrides.py), each invoked below at
# its legacy slot; the interleaved non-attention adjustments stay.
from sglang.srt.arg_groups.overrides import (
_attention_backend_default,
_attention_backend_dual_chunk,
_attention_backend_fa3_fp8_fallback,
_attention_backend_platform_fallbacks,
_fa4_page_constraint,
_intel_xpu_page_constraint,
_mla_backend_page_constraints,
run_post_process_pass,
)
# Split-backend override + default fill.
run_post_process_pass(server_args, _attention_backend_default)
# Torch native and flex attention backends
attention_backend = resolved_view(server_args).attention_backend
if attention_backend == "torch_native":
logger.warning(
"Cuda graph is disabled because of using torch native attention backend"
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
if attention_backend == "flex_attention":
logger.warning(
"Cuda graph is disabled because of using torch Flex Attention backend"
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
assert (
cfg.speculative_algorithm is None
), "Speculative decoding is currently not supported with Flex Attention backend"
# Whisper's encoder token padding conflicts with prefix caching.
# Only disable for Whisper; other encoder-decoder models (e.g., mllama) use radix cache.
if (
model_config.is_encoder_decoder
and not cfg.disable_radix_cache
and "WhisperForConditionalGeneration"
in (model_config.hf_config.architectures or [])
):
logger.info("Radix cache is disabled for Whisper")
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
disable_radix_cache=True,
)
# Major NVIDIA platforms backends: the page-size snaps of this family
# moved to the resolution pipeline (arg_groups/overrides.py:
# _mla_backend_page_constraints); the raises and the cutedsl prefill
# fallback stay below.
run_post_process_pass(server_args, _mla_backend_page_constraints)
# The TRT-LLM / tokenspeed MLA kv-dtype validations moved to the
# resolution pipeline (arg_groups/overrides.py:
# _mla_kv_cache_dtype_checks), invoked here at their legacy slot.
from sglang.srt.arg_groups.overrides import _mla_kv_cache_dtype_checks
run_post_process_pass(server_args, _mla_kv_cache_dtype_checks)
# The CuteDSL MLA validation + prefill fill moved to the resolution
# pipeline (arg_groups/overrides.py: _cutedsl_prefill_backend_fill),
# invoked here at its legacy slot.
from sglang.srt.arg_groups.overrides import _cutedsl_prefill_backend_fill
run_post_process_pass(server_args, _cutedsl_prefill_backend_fill)
prefill_backend, decode_backend = server_args._resolved_attention_backends()
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()
):
raise ValueError(
"TRTLLM MHA backend for prefill requires Hopper (SM90), Blackwell (SM100), or SM120 GPUs. "
"Please use a different prefill backend."
)
if (
prefill_backend == "trtllm_mha"
and is_sm120_supported()
and (
cfg.kv_cache_dtype == "fp8_e4m3"
or (
envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get() or 0.0
)
> 0
)
):
raise ValueError(
"TRTLLM FMHAv2 prefill on SM120 does not support "
"fp8_e4m3 KV cache or skip-softmax."
)
if decode_backend == "trtllm_mha" and not (
is_sm90_supported() or is_sm100_supported() or is_sm120_supported()
):
raise ValueError(
"TRTLLM MHA backend for decode is only supported on Hopper (SM90), Blackwell (SM100) and (SM120) GPUs. Please use a different decode backend."
)
if (
prefill_backend == "trtllm_mha"
and not is_sm100_supported()
and (cfg.enable_prefill_context_parallel or cfg.attn_cp_size > 1)
):
raise ValueError(
"Prefill context parallelism with the TRTLLM MHA prefill backend "
"requires SM100 (trtllm-gen context kernel): the SM90/SM120 "
"fmha_v2 prefill path does not implement CP shard masking."
)
run_post_process_pass(server_args, _attention_backend_fa3_fp8_fallback)
run_post_process_pass(server_args, _fa4_page_constraint)
# AMD platforms backends
if resolved_view(server_args).attention_backend == "aiter":
if model_config.context_len > 8192:
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
mem_fraction_static=cfg.mem_fraction_static * 0.85,
)
# 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":
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."
)
run_post_process_pass(server_args, _intel_xpu_page_constraint)
# Dual chunk flash attention backend
run_post_process_pass(server_args, _attention_backend_dual_chunk)
if resolved_view(server_args).attention_backend == "dual_chunk_flash_attn":
logger.warning(
"Mixed chunk and radix cache are disabled when using dual-chunk flash attention backend"
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
enable_mixed_chunk=False,
)
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
disable_radix_cache=True,
)
def handle_linear_attn_backend(server_args: Any):
cfg = resolving_view(server_args)
import torch
# SM100+: default to FlashInfer GDN decode (and MTP verify, via pool API)
# when the user hasn't explicitly chosen a decode backend and
# mamba-ssm-dtype is bf16 (required by FlashInfer GDN on SM100+).
# Fixed in FlashInfer v0.6.7: flashinfer-ai/flashinfer#2810
if (
cfg.linear_attn_decode_backend is None
and cfg.linear_attn_backend != "helion"
and is_sm100_supported()
and cfg.mamba_ssm_dtype == "bfloat16"
# Stage 4: flashinfer's recurrent_kda compiles the state slot stride
# as a free int64, so it reads the page-major/unified envelope-strided
# state correctly — the unified-memory skip is no longer needed (the
# page-major gate now allows flashinfer for linear-attn decode).
):
declare_resolution(
server_args,
"_handle_linear_attn_backend",
linear_attn_decode_backend="flashinfer",
)
logger.info(
"SM100+ detected with mamba-ssm-dtype=bfloat16, "
"defaulting --linear-attn-decode-backend to flashinfer."
)
# SM100+ FlashInfer GDN decode requires bf16 state; SM90 uses float32.
decode = cfg.linear_attn_decode_backend or cfg.linear_attn_backend
# FlashKDA is a prefill-only KDA kernel (no decode kernel) but shares the
# backend choice list, so guard it from being selected for decode: error
# on an explicit --linear-attn-decode-backend flashkda, and fall back to
# triton decode when it was only inherited from base=flashkda (prefill
# keeps FlashKDA).
if decode == "flashkda":
if cfg.linear_attn_decode_backend == "flashkda":
raise ValueError(
"--linear-attn-decode-backend flashkda is not supported: "
"FlashKDA is prefill-only. Use "
"--linear-attn-prefill-backend flashkda (decode stays on triton)."
)
declare_resolution(
server_args,
"_handle_linear_attn_backend",
linear_attn_decode_backend="triton",
)
decode = "triton"
logger.info(
"FlashKDA is prefill-only; using triton for KDA decode "
"(FlashKDA stays on prefill)."
)
if (
decode == "flashinfer"
and cfg.mamba_ssm_dtype != "bfloat16"
and is_cuda()
and torch.cuda.get_device_capability()[0] >= 10
):
raise ValueError(
"--linear-attn-decode-backend flashinfer on SM100+ requires "
"--mamba-ssm-dtype bfloat16, "
f"got {cfg.mamba_ssm_dtype!r}"
)
verify = cfg.linear_attn_verify_backend
if verify is None and decode == "flashinfer":
verify = "flashinfer"
if (
verify == "flashinfer"
and cfg.mamba_ssm_dtype != "bfloat16"
and is_cuda()
and torch.cuda.get_device_capability()[0] >= 10
):
raise ValueError(
"--linear-attn-verify-backend flashinfer on SM100+ requires "
"--mamba-ssm-dtype bfloat16, "
f"got {cfg.mamba_ssm_dtype!r}"
)
# SM100+ FlashInfer GDN prefill requires CUDA 13+ (CuTe DSL kernel)
# for correctness and best performance.
prefill = cfg.linear_attn_prefill_backend or cfg.linear_attn_backend
cuda_version = torch.version.cuda
cuda_major = int(cuda_version.split(".")[0]) if cuda_version is not None else 0
if (
prefill == "flashinfer"
and is_cuda()
and torch.cuda.get_device_capability()[0] >= 10
and cuda_major < 13
):
raise ValueError(
"--linear-attn-prefill-backend flashinfer on SM100+ requires CUDA 13+, "
f"got CUDA {cuda_version or 'unknown'}"
)
# ReplaySSM buffered decode guards. Runs on Triton, or Helion for KDA.
# cuda-graph is supported (slice 1b: CUDA-graph-safe static
# write-cursor buffers). The RADIX prefix cache is now supported (slice
# 2b: the decode kernel force-flushes the ring into temporal[slot] on
# the radix track boundary `seq_lens % mamba_track_interval == 0`, and
# the COW copy-into-slot path resets the ring cursor) -- so the
# --disable-radix-cache requirement is dropped.
#
# Slice 2b only wires the no_buffer mamba scheduler strategy (the
# default). The extra_buffer strategy donates the track snapshot via
# `donate_mamba_ping_pong_slot` with a separate ping-pong slot swap that
# does NOT route through MambaPool.copy_from, so the ReplaySSM ring
# cursor of the donated/kept slot would not be reset there. Handling
# that donation path is a follow-up; for now require no_buffer.
if cfg.enable_linear_replayssm:
if decode not in {"triton", "helion"}:
raise ValueError(
"--enable-linear-replayssm requires Triton, or Helion for "
"KDA, as the linear-attn decode backend; got "
f"--linear-attn-decode-backend={decode!r}."
)
from sglang.srt.arg_groups.overrides import (
mamba_extra_buffer_of,
)
if mamba_extra_buffer_of(resolved_view(server_args)):
raise ValueError(
"--enable-linear-replayssm requires --mamba-radix-cache-strategy "
"no_buffer (the default); the extra_buffer ping-pong "
"donation path is not yet supported (follow-up). Got "
f"--mamba-radix-cache-strategy={cfg.mamba_radix_cache_strategy!r}."
)
if cfg.disaggregation_mode != "null":
# The disaggregated decode pool (HybridMambaDecodeReqToTokenPool)
# is not wired for the ReplaySSM ring, so the flag would silently
# no-op there; disagg also runs a different cache/coordination
# flow that is not yet validated for ReplaySSM (follow-up).
raise ValueError(
"--enable-linear-replayssm is not supported under PD "
"disaggregation yet (follow-up). Got "
f"--disaggregation-mode={cfg.disaggregation_mode!r}."
)
if cfg.linear_replayssm_cache_len < 1:
raise ValueError(
"--linear-replayssm-cache-len must be >= 1, got "
f"{cfg.linear_replayssm_cache_len}."
)
# ReplaySSM spec-verify (Part B of #28511): linear-chain target verify via
# fold-every-commit -- the verify stores each draft step's raw inputs into
# the per-slot (rawv, rawk, g, beta) window and the commit replays the
# accepted prefix into the fp32 checkpoint. The intra-window interaction
# uses a strictly-lower causal mask, so it is valid ONLY for a linear
# draft chain (speculative_eagle_topk in {None, 1}, i.e. NEXTN / MTP);
# EAGLE tree verify (topk > 1) must fall back to the recurrent verify.
# GDN sizes the window to the draft maximum; KDA (kda_backend) keeps a
# --linear-replayssm-cache-len window and folds via its own fused
# verify ring-write + commit_kda_replayssm_after_verify.
if cfg.enable_linear_replayssm_spec:
if cfg.speculative_eagle_topk not in (None, 1):
raise ValueError(
"--enable-linear-replayssm-spec requires a linear draft chain "
"(--speculative-eagle-topk in {None, 1}); the chunked verify "
"kernel uses a strictly-lower causal mask and is invalid for "
"EAGLE tree verify. Got "
f"--speculative-eagle-topk={cfg.speculative_eagle_topk!r}."
)
if decode not in ("triton", "flashinfer"):
raise ValueError(
"--enable-linear-replayssm-spec requires the triton or "
"flashinfer linear-attn decode backend, got "
f"--linear-attn-decode-backend={decode!r}."
)
from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode,
read_ragged_verify_mode,
)
ragged_mode = read_ragged_verify_mode()
if ragged_mode is not RaggedVerifyMode.STATIC:
# Ragged ring-writes need the KDA fold-every-commit family
# (DSPARK/DFLASH) + the triton verify kernel (nv_cutedsl falls
# back to it for ragged layouts). The GDN ring-write kernels do
# not take the ragged layout and the flashinfer verify kernel
# never writes the ring -> a stale ring would be folded; keep
# refusing those combinations.
_algo = (cfg.speculative_algorithm or "").upper()
verify = cfg.linear_attn_verify_backend
if _algo not in ("DSPARK", "DFLASH") or verify not in (
"triton",
"nv_cutedsl",
):
raise ValueError(
"--enable-linear-replayssm-spec with "
f"SGLANG_RAGGED_VERIFY_MODE={ragged_mode.value} requires the "
"KDA fold-every-commit family (DSPARK/DFLASH) and a "
"ring-writing verify kernel (--linear-attn-verify-backend "
"triton or nv_cutedsl); got "
f"algorithm={cfg.speculative_algorithm!r}, "
f"verify={verify!r}. Use SGLANG_RAGGED_VERIFY_MODE=static."
)
if cfg.disaggregation_mode == "prefill":
raise ValueError(
"--enable-linear-replayssm-spec is not supported on a PD "
"prefill server: the ring is spec-verify-only scratch and "
"the prefill server never runs spec verify."
)
if cfg.enable_linear_replayssm:
raise ValueError(
"--enable-linear-replayssm-spec and --enable-linear-replayssm are "
"mutually exclusive: they share the ring storage but drive it "
"with incompatible cursor protocols (per-decode-forward vs "
"per-verify-commit advance)."
)
if cfg.mamba_ssm_dtype is None:
logger.info(
"--enable-linear-replayssm-spec: setting --mamba-ssm-dtype "
"float32 (the closed-loop exact fold keeps the SSM checkpoint "
"bit-identical to the recurrent baseline)."
)
declare_resolution(
server_args,
"_handle_linear_attn_backend",
mamba_ssm_dtype="float32",
)
elif cfg.mamba_ssm_dtype != "float32":
logger.warning(
"--enable-linear-replayssm-spec with --mamba-ssm-dtype=%s: the "
"closed-loop fold re-quantizes the committed state each "
"commit/flush (fp32 keeps it bit-exact to the fp32 recurrent "
"baseline), so it may drift over long sequences. Validate "
"accuracy for your model.",
cfg.mamba_ssm_dtype,
)
def handle_multi_item_scoring(server_args: Any):
"""Setup and validate multi-item scoring constraints.
Auto-disables settings incompatible with MIS mechanics (CUDA graph,
radix cache, chunked prefill). Asserts on attention backend since
changing it silently could surprise users who intentionally picked
a non-flashinfer backend.
"""
cfg = resolving_view(server_args)
if not cfg.enable_mis:
return
if cfg.cuda_graph_config.decode.backend != Backend.DISABLED:
logger.warning("CUDA graph is disabled because --enable-mis is set.")
declare_resolution(
server_args,
"_handle_multi_item_scoring",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_multi_item_scoring",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
if not cfg.disable_radix_cache:
logger.warning("Radix cache is disabled because --enable-mis is set.")
declare_resolution(
server_args,
"_handle_multi_item_scoring",
disable_radix_cache=True,
)
if cfg.chunked_prefill_size != -1:
logger.warning("Chunked prefill is disabled because --enable-mis is set.")
declare_resolution(
server_args,
"_handle_multi_item_scoring",
chunked_prefill_size=-1,
)
prefill_backend, decode_backend = server_args._resolved_attention_backends()
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. "
f"Current backends: prefill={prefill_backend}, decode={decode_backend}"
)
def handle_deterministic_inference(server_args: Any):
from sglang.srt.server_args import (
RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND,
)
cfg = resolving_view(server_args)
if cfg.rl_on_policy_target is not None:
logger.warning("Enable deterministic inference because of rl_on_policy_target.")
declare_resolution(
server_args,
"_handle_deterministic_inference",
enable_deterministic_inference=True,
)
# For VLM
envs.SGLANG_VLM_CACHE_SIZE_MB.set(0)
# TODO remove this environment variable as a whole
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set(True)
if cfg.enable_deterministic_inference:
if cfg.enable_aiter_allreduce_fusion:
logger.warning(
"Disable --enable-aiter-allreduce-fusion because deterministic inference is enabled."
)
declare_resolution(
server_args,
"_handle_deterministic_inference",
enable_aiter_allreduce_fusion=False,
)
# Moved to the resolution pipeline (arg_groups/overrides.py:
# _deterministic_allreduce_fusion_disable), invoked here at its
# legacy slot.
from sglang.srt.arg_groups.overrides import (
_deterministic_allreduce_fusion_disable,
run_post_process_pass,
)
run_post_process_pass(server_args, _deterministic_allreduce_fusion_disable)
# The forced-pytorch sampling write and the attention backend
# fill/validation moved to the resolution pipeline
# (arg_groups/overrides.py), invoked at their legacy slots.
from sglang.srt.arg_groups.overrides import (
_deterministic_attention_backend,
_deterministic_sampling_backend,
run_post_process_pass,
)
run_post_process_pass(server_args, _deterministic_sampling_backend)
is_deepseek_model = False
if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE:
try:
hf_config = server_args.get_model_config().hf_config
model_arch = hf_config.architectures[0]
is_deepseek_model = model_arch in [
"DeepseekV2ForCausalLM",
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"MistralLarge3ForCausalLM",
"PixtralForConditionalGeneration",
"GlmMoeDsaForCausalLM",
"Glm4MoeLiteForCausalLM",
]
except Exception:
pass
# Check attention backend
run_post_process_pass(server_args, _deterministic_attention_backend)
attention_backend = resolved_view(server_args).attention_backend
if is_deepseek_model:
if attention_backend not in RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND:
raise ValueError(
f"Currently only {RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND} attention backends are supported for deterministic inference with absorbed-MLA models. But you're using {attention_backend}."
)
if attention_backend == "fa4" and not is_sm100_or_sm110_supported():
raise ValueError(
"Deterministic inference with absorbed-MLA models on the fa4 "
"attention backend requires SM100/SM110: it runs "
"absorbed MLA, whose qv argument flash_attn.cute only "
"implements on those archs."
)
if attention_backend not in RADIX_SUPPORTED_DETERMINISTIC_ATTENTION_BACKEND:
# Currently, only certain backends support radix cache. Support for other backends is in progress
declare_resolution(
server_args,
"_handle_deterministic_inference",
disable_radix_cache=True,
)
logger.warning(
f"Currently radix cache is not compatible with {attention_backend} attention backend for deterministic inference. It will be supported in the future."
)
# Check TP size
if cfg.tp_size > 1:
if is_hip():
# AMD: use 1-stage all-reduce kernel which is inherently deterministic
# (each GPU reads all data from all GPUs, reduces locally in fixed order)
logger.info("AMD/ROCm: Using 1-stage all-reduce kernel (deterministic)")
else:
# CUDA: use NCCL tree algorithm
os.environ["NCCL_ALGO"] = "allreduce:tree"
# Not declared: set_default_server_args() writes this field
# too, through its `args` parameter, so a declaration here
# would be a second source for one field.
declare_resolution(
server_args,
"_handle_deterministic_inference",
disable_custom_all_reduce=True,
)
# should_torch_symm_mem_allreduce() takes the
# symmetric-memory path only below a byte threshold, so
# which reduce runs would follow the token count.
declare_resolution(
server_args,
"_handle_deterministic_inference",
enable_torch_symm_mem=False,
)
# Each channel carries a differently shaped tree and the
# channel count is picked from the message size, so a
# token's reduction order would follow the token count.
nchannels = str(envs.SGLANG_DETERMINISTIC_NCCL_NCHANNELS.get())
os.environ["NCCL_MIN_NCHANNELS"] = nchannels
os.environ["NCCL_MAX_NCHANNELS"] = nchannels
logger.warning(
"NCCL_ALGO is set to 'allreduce:tree', the NCCL channel count is pinned, and custom and symmetric-memory all reduce are disabled for deterministic inference when TP size > 1."
)
@@ -0,0 +1,455 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the CUDA-graph capture configuration."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.connector import ConnectorType
from sglang.srt.model_executor.cuda_graph_config import (
ALLOWED_BACKENDS_PER_PHASE,
Backend,
CudaGraphConfig,
Phase,
default_cuda_graph_config,
with_phase,
)
from sglang.srt.platforms import current_platform
from sglang.srt.utils.common import (
is_cpu,
is_hip,
is_mps,
is_npu,
is_xpu,
parse_connector_type,
)
from sglang.srt.utils.hf_transformers_utils import check_gguf_file
logger = logging.getLogger(__name__)
def parse_cuda_graph_config(server_args: Any):
"""Resolve cuda_graph_config from explicit JSON, per-phase
convenience flags, legacy global flags, and defaults.
Precedence (highest first): explicit JSON > convenience > legacy > defaults.
Also populates server_args._cuda_graph_config_locked — the set of
(phase, key) tuples that came from non-default sources; the
auto-disable cascade respects this lock (the old
--enforce-piecewise-cuda-graph semantics generalized).
"""
cfg = resolving_view(server_args)
raw_input = cfg.cuda_graph_config
if isinstance(raw_input, CudaGraphConfig):
explicit_input = raw_input.to_dict()
else:
explicit_input = raw_input or {}
config = default_cuda_graph_config()
locked: set = set()
def _set(phase: str, key: str, value: Any) -> None:
setattr(getattr(config, phase), key, value)
locked.add((phase, key))
# ---- Legacy global flags (lowest precedence above defaults) ----
if cfg.disable_cuda_graph:
_set(Phase.DECODE, "backend", Backend.DISABLED)
_set(Phase.PREFILL, "backend", Backend.DISABLED)
# ---- Boolean per-phase off-switches ----
# Below the explicit backend selectors so --cuda-graph-backend-*
# wins if both are given.
if cfg.disable_prefill_cuda_graph:
_set(Phase.PREFILL, "backend", Backend.DISABLED)
if cfg.disable_decode_cuda_graph:
_set(Phase.DECODE, "backend", Backend.DISABLED)
# ---- Per-phase convenience flags ----
if cfg.cuda_graph_backend_decode is not None:
_set(Phase.DECODE, "backend", cfg.cuda_graph_backend_decode)
if cfg.cuda_graph_backend_prefill is not None:
_set(Phase.PREFILL, "backend", cfg.cuda_graph_backend_prefill)
if cfg.cuda_graph_max_bs_decode is not None:
_set(Phase.DECODE, "max_bs", cfg.cuda_graph_max_bs_decode)
if cfg.cuda_graph_max_bs_prefill is not None:
_set(Phase.PREFILL, "max_bs", cfg.cuda_graph_max_bs_prefill)
if cfg.cuda_graph_bs_decode is not None:
_set(Phase.DECODE, "bs", cfg.cuda_graph_bs_decode)
if cfg.cuda_graph_bs_prefill is not None:
_set(Phase.PREFILL, "bs", cfg.cuda_graph_bs_prefill)
if cfg.cuda_graph_tc_compiler is not None:
# Written to both phases so the value is in place when TC_PIECEWISE
# decode is implemented; today decode ignores it.
_set(Phase.DECODE, "tc_compiler", cfg.cuda_graph_tc_compiler)
_set(Phase.PREFILL, "tc_compiler", cfg.cuda_graph_tc_compiler)
# ---- Explicit JSON config (highest precedence) ----
for phase, phase_config in explicit_input.items():
if not isinstance(phase_config, dict):
continue
for key, value in phase_config.items():
_set(phase, key, value)
declare_resolution(
server_args,
"_parse_cuda_graph_config",
cuda_graph_config=config,
)
server_args._cuda_graph_config_locked = locked
def apply_cuda_graph_compatibility(server_args: Any):
"""Auto-disable prefill cuda graph for incompatible configs.
Rules are split per backend — TcPiecewise and Breakable have
different constraints. Skipped when the user explicitly set the
prefill backend (this folds in the old
--enforce-piecewise-cuda-graph contract).
"""
cfg = resolving_view(server_args)
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
return
# Breakable is the CUDA default but not multimodal-compatible;
# piecewise-allowlisted archs run their validated decoder prefill
# there instead. Archs also on the breakable allowlist keep it --
# 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
# 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"
):
logger.info(
"Using tc_piecewise CUDA graph for validated multimodal " "decoder prefill."
)
declare_resolution(
server_args,
"_apply_cuda_graph_compatibility",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.TC_PIECEWISE
),
)
if cfg.cuda_graph_config.prefill.backend == Backend.TC_PIECEWISE:
server_args._disable_tc_piecewise_cudagraph_if_incompatible()
elif cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE:
server_args._disable_breakable_cudagraph_if_incompatible()
elif cfg.cuda_graph_config.prefill.backend == Backend.FULL:
server_args._disable_full_prefill_cudagraph_if_incompatible()
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.
"""
cfg = resolving_view(server_args)
rules = [
(
"model-arch blacklist",
lambda: server_args.get_model_config().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),
("pipeline parallelism (pp_size > 1)", lambda: cfg.pp_size > 1),
(
"non-CUDA hardware (HIP/NPU/CPU/MPS/XPU)",
lambda: is_hip() or is_npu() or is_cpu() or is_mps() or is_xpu(),
),
(
"OOT platform without piecewise support",
lambda: current_platform.is_out_of_tree()
and not current_platform.support_piecewise_cuda_graph(),
),
(
"MoE A2A backend",
lambda: resolved_view(server_args).moe_a2a_backend != "none",
),
# Dynamo blocks LoRA under tc_piecewise (per-batch LoRABatchInfo
# rebinds break guards); breakable/full support LoRA.
("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,
),
(
"GGUF quantization",
lambda: cfg.load_format == "gguf"
or resolved_view(server_args).quantization == "gguf"
or check_gguf_file(cfg.model_path),
),
("DLLM (diffusion LLM)", lambda: cfg.dllm_algorithm is not None),
(
"CPU offload / hierarchical cache",
lambda: cfg.cpu_offload_gb > 0 or cfg.enable_hierarchical_cache,
),
(
"deterministic inference",
lambda: cfg.enable_deterministic_inference,
),
("PD disaggregation", lambda: cfg.disaggregation_mode != "null"),
("symmetric memory", lambda: cfg.enable_symm_mem),
(
"expert distribution recorder",
lambda: cfg.enable_eplb
or cfg.expert_distribution_recorder_mode is not None,
),
(
"context parallel (attn_cp_size > 1)",
lambda: resolved_view(server_args).attn_cp_size > 1,
),
("CUDA graph debug mode", lambda: cfg.debug_cuda_graph),
(
"DSA prefill context parallelism",
lambda: cfg.enable_dsa_prefill_context_parallel,
),
# Capture builds a dummy extend forward with attn_dcp_metadata=None.
(
"decode context parallel (dcp_size > 1)",
lambda: cfg.dcp_size > 1,
),
]
for _name, predicate in rules:
if predicate():
declare_resolution(
server_args,
"_disable_tc_piecewise_cudagraph_if_incompatible",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
# One decision, one declaration: every rule declares the same
# value, so a later match would only append a duplicate entry.
break
def disable_breakable_cudagraph_if_incompatible(server_args: Any):
"""Breakable (segmented capture, no torch.compile). Breakable enforces
memory-saver rejection in its own __init__; config-time rules can be
added here as they're discovered.
"""
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
rules = [
# DSV4 is BCG-compatible but introduces heavy memory pressure: the
# 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),
),
# CP all_gather replay size mismatch under BCG.
(
"context parallel (attn_cp_size > 1)",
lambda: resolved_view(server_args).attn_cp_size > 1
and not supports_prefill_cp_bcg(server_args),
),
# Capture builds a dummy extend forward with attn_dcp_metadata=None.
(
"decode context parallel (dcp_size > 1)",
lambda: cfg.dcp_size > 1,
),
# TBO capture is unsupported.
(
"two-batch overlap",
lambda: cfg.enable_two_batch_overlap,
),
(
"unvalidated a2a backend",
lambda: resolved_view(server_args).moe_a2a_backend
not in ("none", "deepep", "megamoe", "flashinfer"),
),
# 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,
),
]
for name, predicate in rules:
if predicate():
logger.warning(
"Breakable CUDA graph is incompatible with %s; "
"disabling prefill CUDA graph.",
name,
)
declare_resolution(
server_args,
"_disable_breakable_cudagraph_if_incompatible",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
return
def disable_full_prefill_cudagraph_if_incompatible(server_args: Any):
"""Full prefill CG: empty rule list today; see the experimental warning."""
cfg = resolving_view(server_args)
rules = []
for name, predicate in rules:
if predicate():
logger.warning(
"Full prefill CUDA graph is incompatible with %s; "
"disabling prefill CUDA graph.",
name,
)
declare_resolution(
server_args,
"_disable_full_prefill_cudagraph_if_incompatible",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
return
def disable_prefill_cuda_graph_for_deepseek_trtllm_mla(server_args: Any):
"""Disable prefill CUDA graph for dsr1 by default when using the trtllm_mla
attention backend. Under any captured prefill CUDA graph (tc_piecewise or
breakable) trtllm_mla falls back to FlashAttention for prefill and regresses
performance, so disable whichever prefill graph backend is in effect.
"""
cfg = resolving_view(server_args)
if (Phase.PREFILL, "backend") in server_args._cuda_graph_config_locked:
return
if cfg.cuda_graph_config.prefill.backend == Backend.DISABLED:
return
if (
"DeepseekV3ForCausalLM"
not in server_args.get_model_config().hf_config.architectures
):
return
prefill_attention_backend, _ = server_args._resolved_attention_backends()
if prefill_attention_backend != "trtllm_mla":
return
logger.warning(
"Disabling prefill CUDA graph (%s) by default for the DeepSeek-V3 arch on "
"the trtllm_mla attention backend (a captured prefill graph forces a "
"FlashAttention fallback that regresses prefill). Set the prefill cuda graph "
"backend explicitly (e.g. --cuda-graph-backend-prefill tc_piecewise) to override.",
cfg.cuda_graph_config.prefill.backend,
)
declare_resolution(
server_args,
"_disable_prefill_cuda_graph_for_deepseek_trtllm_mla",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
def apply_deepep_adjustments(server_args: Any):
"""Config adjustments required by the DeepEP a2a backend."""
cfg = resolving_view(server_args)
if resolved_view(server_args).moe_a2a_backend != "deepep":
return
# Non-multiple-of-8 prefill buckets can hang DeepEP a2a capture under
# breakable CUDA graph
if cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE:
bs = cfg.cuda_graph_config.prefill.bs
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)
aligned = sorted({((b + 7) // 8) * 8 for b in bs})
if aligned != sorted(bs):
logger.info(
"Breakable prefill CUDA graph with DeepEP requires bucket "
"sizes divisible by 8; aligning %s -> %s.",
sorted(bs),
aligned,
)
declare_resolution(
server_args,
"_apply_deepep_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config,
Phase.PREFILL,
bs=aligned,
max_bs=aligned[-1],
),
)
def apply_inkling_prefill_cuda_graph_default(server_args: Any):
"""Inkling opts into full-graph prefill CUDA-graph capture. Must run
before _handle_cuda_graph_config: the generic breakable default is
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."""
cfg = resolving_view(server_args)
if (
cfg.cuda_graph_backend_prefill is not None
or cfg.disable_prefill_cuda_graph
or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE
):
return
arch = server_args.get_model_config().hf_config.architectures[0]
if arch in (
"InklingForConditionalGeneration",
"InklingForConditionalGenerationMTP",
):
declare_resolution(
server_args,
"_apply_inkling_prefill_cuda_graph_default",
cuda_graph_backend_prefill=Backend.FULL,
)
def apply_muse_glimmer_prefill_cuda_graph_max_bs_default(server_args: Any):
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]
if arch in ("MuseGlimmerForCausalLM", "MuseGlimmerForConditionalGeneration"):
declare_resolution(
server_args,
"_apply_muse_glimmer_prefill_cuda_graph_max_bs_default",
cuda_graph_max_bs_prefill=512,
)
def handle_cuda_graph_config(server_args: Any):
cfg = resolving_view(server_args)
server_args._parse_cuda_graph_config()
server_args._apply_cuda_graph_compatibility()
server_args._apply_deepep_adjustments()
server_args._apply_cuda_graph_disaggregation_roles()
server_args._validate_cuda_graph_config()
# Warn on the final resolved config (not inside the compat cascade —
# that path is skipped when the user explicitly sets the backend,
# which is the only way to get 'full' for prefill today).
if cfg.cuda_graph_config.prefill.backend == Backend.FULL:
logger.warning(
"cuda_graph_config[prefill].backend='full' is experimental. "
"Use breakable or tc_piecewise for production workloads."
)
def validate_cuda_graph_config(server_args: Any):
cfg = resolving_view(server_args)
if cfg.cuda_graph_config is None:
return
for phase in Phase.ALL:
backend = getattr(cfg.cuda_graph_config, phase).backend
if backend not in ALLOWED_BACKENDS_PER_PHASE[phase]:
raise ValueError(
f"--cuda-graph-config[{phase}].backend={backend!r} not allowed; "
f"allowed: {ALLOWED_BACKENDS_PER_PHASE[phase]}"
)
+124
View File
@@ -0,0 +1,124 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for diffusion-LM inference."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolving_view,
)
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
from sglang.srt.utils.common import is_hip
logger = logging.getLogger(__name__)
def handle_dllm_inference(server_args: Any):
cfg = resolving_view(server_args)
if cfg.dllm_algorithm is None:
return
# On AMD/HIP, disable cuda graph for DLLM (the attention_backend
# resolution moved to the pipeline: arg_groups/overrides.py
# _dllm_attention_backend, invoked below at its legacy slot).
if is_hip():
if (
cfg.cuda_graph_config.decode.backend != Backend.DISABLED
or cfg.cuda_graph_config.prefill.backend != Backend.DISABLED
):
logger.warning(
"Cuda graph is disabled for diffusion LLM inference on AMD GPUs"
)
declare_resolution(
server_args,
"_handle_dllm_inference",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_dllm_inference",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
from sglang.srt.arg_groups.overrides import (
_dllm_attention_backend,
_dllm_overlap_disable,
run_post_process_pass,
)
run_post_process_pass(server_args, _dllm_attention_backend)
run_post_process_pass(server_args, _dllm_overlap_disable)
# The page-size alignment + block-size cap for dllm moved to the
# resolution pipeline (arg_groups/overrides.py: _dllm_page_size).
# Invoked outside the radix gate: the alignment fill keeps its radix
# gate inside the pass, the block-size cap applies regardless (it
# replaces the unconditional scheduler-init fallback).
from sglang.srt.arg_groups.overrides import _dllm_page_size
run_post_process_pass(server_args, _dllm_page_size)
if not cfg.disable_radix_cache:
if cfg.enable_hierarchical_cache:
logger.warning(
"Hierarchical cache is disabled because of using diffusion LLM inference"
)
declare_resolution(
server_args,
"_handle_dllm_inference",
enable_hierarchical_cache=False,
)
if cfg.enable_lmcache:
logger.warning(
"LMCache is disabled because of using diffusion LLM inference"
)
declare_resolution(
server_args, "_handle_dllm_inference", enable_lmcache=False
)
if cfg.enable_flexkv:
logger.warning(
"FlexKV is disabled because of using diffusion LLM inference"
)
declare_resolution(
server_args, "_handle_dllm_inference", enable_flexkv=False
)
if cfg.pp_size > 1:
logger.warning(
"Pipeline parallelism is disabled because of using diffusion LLM inference"
)
declare_resolution(
server_args,
"_handle_dllm_inference",
pp_size=1,
)
if cfg.enable_lora:
logger.warning("Currently LoRA is not supported by diffusion LLM inference.")
declare_resolution(server_args, "_handle_dllm_inference", enable_lora=False)
if cfg.disaggregation_mode != "null":
logger.warning(
"Currently disaggregation is not supported by diffusion LLM inference."
)
declare_resolution(
server_args,
"_handle_dllm_inference",
disaggregation_mode="null",
)
if cfg.enable_mixed_chunk:
logger.warning(
"Mixed chunked prefill is disabled because of using diffusion LLM inference."
)
declare_resolution(
server_args,
"_handle_dllm_inference",
enable_mixed_chunk=False,
)
@@ -0,0 +1,209 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the hierarchical KV cache."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolving_view,
)
logger = logging.getLogger(__name__)
def handle_hicache(server_args: Any):
"""Normalize hicache-related knobs into a valid runtime configuration.
Resolution order:
1) Layout <-> I/O compatibility for direct conflicts.
2) Storage <-> layout compatibility (may rewrite layout).
"""
cfg = resolving_view(server_args)
# Skip all normalization when neither hicache nor decode-offload path is active.
if not (
cfg.enable_hierarchical_cache
or cfg.disaggregation_decode_enable_offload_kvcache
or (
cfg.disaggregation_mode == "decode"
and cfg.disaggregation_decode_retraction_backup in (None, "host_pool")
)
):
return
server_args._validate_hicache_host_memory_mode()
# Step 1: Initial layout-io compatibility normalization.
server_args._resolve_layout_io_compatibility()
# Step 2: Storage-layout normalization without changing io backend.
server_args._resolve_storage_layout_compatibility()
# Step 3: DCP compatibility for the L2 (device<->host) path.
server_args._resolve_hicache_dcp_compatibility()
def handle_hicache_ratio_default(server_args: Any):
"""Default the host/device ratio per host memory mode.
Runs before the dummy-model boundary: direct HostKVCache consumers
(unit fixtures, dummy-model launches) must never see a None ratio.
buffer_only stages in flight rather than retaining, so it needs only
enough to cover the write backlog plus parked prefetches.
A decode server keeps the ratio unset here: kv_cache_builder resolves
it against the retraction-backup backend (1.0 for host_pool, else 2.0).
"""
cfg = resolving_view(server_args)
if cfg.hicache_ratio is None and cfg.disaggregation_mode != "decode":
declare_resolution(
server_args,
"_handle_hicache_ratio_default",
hicache_ratio=(
1.2 if cfg.hicache_host_memory_mode == "buffer_only" else 2.0
),
)
def resolve_hicache_dcp_compatibility(server_args: Any):
cfg = resolving_view(server_args)
if cfg.dcp_size <= 1 or not cfg.enable_hierarchical_cache:
return
if cfg.hicache_storage_backend is not None:
raise NotImplementedError(
"--hicache-storage-backend (L3) with --dcp-size > 1 is not "
"supported yet: under DCP each rank holds a distinct "
"interleaved MLA KV shard, so the rank-0-only replicated-MLA "
"backup and the storage keys must become dcp_rank-aware "
"first. Run HiCache+DCP with L1/L2 only."
)
if cfg.speculative_algorithm not in (None, "DSPARK"):
raise NotImplementedError(
"HiCache with --dcp-size > 1 only supports DSPARK speculative "
"decoding; other draft-model host pools have no DCP index "
"translation."
)
if cfg.enable_lmcache:
raise NotImplementedError(
"--enable-lmcache with --dcp-size > 1 is not supported: "
"LMCache has no DCP-aware index translation."
)
if cfg.enable_hisparse:
raise NotImplementedError(
"--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():
raise NotImplementedError(
"HiCache with --dcp-size > 1 is only supported for MLA models: "
"the index translation lives in MLATokenToKVPoolHost, and the "
"MHA host pool has none."
)
logger.info(
"HiCache + DCP enabled (L1/L2 only): host pool uses widened "
"logical slot accounting with per-rank physical translation at "
"the transfer boundary (dcp_size=%d).",
cfg.dcp_size,
)
def resolve_layout_io_compatibility(server_args: Any):
cfg = resolving_view(server_args)
if (
cfg.hicache_mem_layout == "page_first_direct"
and cfg.hicache_io_backend == "kernel"
):
declare_resolution(
server_args,
"_resolve_layout_io_compatibility",
hicache_io_backend="direct",
)
logger.warning(
"Kernel io backend does not support page first direct layout, switching to direct io backend"
)
if cfg.hicache_mem_layout == "page_first" and cfg.hicache_io_backend == "direct":
declare_resolution(
server_args,
"_resolve_layout_io_compatibility",
hicache_mem_layout="page_first_direct",
)
logger.warning(
"Page first layout is not supported with direct IO backend, switching to page first direct layout"
)
def resolve_storage_layout_compatibility(server_args: Any):
cfg = resolving_view(server_args)
if (
cfg.hicache_storage_backend != "mooncake"
or cfg.hicache_mem_layout != "layer_first"
):
return
if cfg.hicache_io_backend == "direct":
new_layout = "page_first_direct"
elif cfg.hicache_io_backend == "kernel":
new_layout = "page_first"
else:
# Keep current behavior for unknown backends (e.g., kernel_ascend).
new_layout = cfg.hicache_mem_layout
declare_resolution(
server_args,
"_resolve_storage_layout_compatibility",
hicache_mem_layout=new_layout,
)
logger.warning(
f"Mooncake storage backend does not support layer_first layout, "
f"switching to {new_layout} layout for {cfg.hicache_io_backend} io backend"
)
def validate_hicache_host_memory_mode(server_args: Any):
cfg = resolving_view(server_args)
if cfg.hicache_host_memory_mode not in ("cache", "buffer_only"):
raise ValueError(
"hicache_host_memory_mode must be 'cache' or 'buffer_only', "
f"got {cfg.hicache_host_memory_mode!r}"
)
# Both modes are defaulted upstream (a decode server resolves the
# ratio later, in kv_cache_builder), so this fires only if that
# defaulting regresses -- never build an unsized host pool.
if (
cfg.hicache_size <= 0
and cfg.hicache_ratio is None
and cfg.disaggregation_mode != "decode"
):
raise ValueError(
f"--hicache-host-memory-mode {cfg.hicache_host_memory_mode} "
"requires a host pool size: pass --hicache-size or "
"--hicache-ratio."
)
if cfg.hicache_host_memory_mode == "cache":
return
if cfg.hicache_storage_backend is None:
raise ValueError(
"--hicache-host-memory-mode buffer_only requires a storage backend "
"(--hicache-storage-backend): host memory is only a staging buffer "
"and all cached data lives in storage."
)
if cfg.hicache_write_policy == "write_back":
raise ValueError(
"--hicache-host-memory-mode buffer_only does not support "
"--hicache-write-policy write_back; use write_through or "
"write_through_selective."
)
if cfg.disaggregation_mode == "decode":
raise ValueError(
"--hicache-host-memory-mode buffer_only is not supported on "
"decode instances: the decode-side prefetch and offload paths "
"bypass the buffer-mode pipeline, fetching without its prefix "
"context and never consuming its staged holds. Prefill "
"instances share the standard scheduler path and are supported."
)
@@ -0,0 +1,425 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for KV-cache dtype and pool compatibility."""
from __future__ import annotations
import logging
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
from sglang.srt.utils.common import (
is_blackwell_supported,
is_cuda,
is_sm100_supported,
is_sm120_supported,
)
logger = logging.getLogger(__name__)
def handle_mxfp8_kv_cache_compatibility(server_args: Any) -> None:
"""MXFP8 KV cache uses operands available only on SM100+ (Blackwell)."""
cfg = resolving_view(server_args)
if cfg.kv_cache_dtype != "mxfp8":
return
if not is_blackwell_supported():
raise ValueError(
"--kv-cache-dtype mxfp8 requires an SM100+ (Blackwell) GPU for the "
"block-scaled operands used by the FA4 MXFP8 attention path."
)
def handle_kv4_compatibility(server_args: Any) -> None:
"""Check FP4 KV cache compatibility with the attention 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()
attention_backend = resolved_view(server_args).attention_backend
if is_cuda():
if cfg.kv_cache_dtype == "nvfp4" and not (
is_sm100_supported() or is_sm120_supported()
):
raise RuntimeError(
"--kv-cache-dtype=nvfp4 requires Blackwell SM100 or SM120. "
"Use --kv-cache-dtype=fp4_mx_block16 for the block-size-16 FP4 recipe."
)
if (
prefill_backend != decode_backend and prefill_backend != "fa4"
): # Take care of prefill=fa4 later
logger.warning(
f"Attention: Using KV4 with PREFILL = {prefill_backend} "
f"and DECODE = {decode_backend}. "
f"Compatibility issues are unlikely, but may occur in rare edge cases."
)
else:
if prefill_backend == "fa4":
if use_mla_backend: # FA4 + MLA
KV4_FA4_MLA_BACKEND_CHOICES = [
"cutlass_mla",
"flashinfer",
"trtllm_mla",
]
assert decode_backend in KV4_FA4_MLA_BACKEND_CHOICES, (
f"KV4 FA4 MLA expects decode_attention_backend to be one of "
f"{KV4_FA4_MLA_BACKEND_CHOICES}, but got {decode_backend}"
)
else: # FA4 + MHA
KV4_FA4_MHA_BACKEND_CHOICES = [
"triton",
"torch_native",
"flex_attention",
]
assert decode_backend in KV4_FA4_MHA_BACKEND_CHOICES, (
f"KV4 FA4 MHA expects decode_attention_backend to be one of "
f"{KV4_FA4_MHA_BACKEND_CHOICES}, but got {decode_backend}"
)
else:
if use_mla_backend: # !FA4 + MLA
KV4_ATTENTION_MLA_BACKEND_CHOICES = [
"cutlass_mla",
"flashinfer",
"trtllm_mla",
]
assert attention_backend in KV4_ATTENTION_MLA_BACKEND_CHOICES, (
f"KV4 MLA expects attention_backend to be one of "
f"{KV4_ATTENTION_MLA_BACKEND_CHOICES}, but got {attention_backend}"
)
else: # !FA4 + MHA
KV4_ATTENTION_MHA_BACKEND_CHOICES = [
"triton",
"torch_native",
"flex_attention",
"trtllm_mha",
]
assert attention_backend in KV4_ATTENTION_MHA_BACKEND_CHOICES, (
f"KV4 MHA expects attention_backend to be one of "
f"{KV4_ATTENTION_MHA_BACKEND_CHOICES}, but got {attention_backend}"
)
else:
raise RuntimeError("KV4 is not tested on non-CUDA platforms.")
def handle_prefill_only_disable_kv_cache(server_args: Any) -> None:
"""Validate --prefill-only-disable-kv-cache backend constraint.
Must run after _handle_attention_backend_compatibility() (which fills
the default attention_backend if unset) and _handle_multi_item_scoring()
(which may further mutate it). The assertion below guards against
accidental call-site reordering: if the resolved attention_backend is
still None, backends haven't settled yet and the resolved (prefill,
decode) pair would be a stale (None, None).
"""
cfg = resolving_view(server_args)
if not cfg.prefill_only_disable_kv_cache:
return
assert resolved_view(server_args).attention_backend is not None, (
"_handle_prefill_only_disable_kv_cache must run after "
"_handle_attention_backend_compatibility() so the prefill backend is resolved."
)
prefill_backend, _ = server_args._resolved_attention_backends()
if prefill_backend not in ("fa3", "fa4"):
raise ValueError(
"--prefill-only-disable-kv-cache currently requires the FA prefill backend "
f"(fa3/fa4), but got prefill backend {prefill_backend!r}. Other prefill-only "
"workloads and backends may be supported in a future change."
)
def handle_cache_compatibility(server_args: Any) -> None:
cfg = resolving_view(server_args)
if (
cfg.disaggregation_decode_retraction_backup == "host_pool"
and cfg.disaggregation_mode != "decode"
):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool is only "
"supported on a PD decode server."
)
if cfg.disaggregation_decode_retraction_backup == "host_pool" and cfg.dcp_size > 1:
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool does not "
"support --dcp-size > 1."
)
if (
cfg.disaggregation_decode_retraction_backup == "host_pool"
and cfg.enable_priority_scheduling
and not cfg.disable_priority_preemption
):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool requires "
"--disable-priority-preemption when priority scheduling is enabled."
)
if cfg.enable_hierarchical_cache and cfg.disable_radix_cache:
raise ValueError(
"The arguments enable-hierarchical-cache and disable-radix-cache are mutually exclusive "
"and cannot be used at the same time. Please use only one of them."
)
if cfg.disaggregation_decode_enable_offload_kvcache:
if cfg.disaggregation_mode != "decode":
raise ValueError(
"The argument disaggregation-decode-enable-offload-kvcache is only supported for decode side."
)
if cfg.hicache_storage_backend is None:
raise ValueError(
"The argument disaggregation-decode-enable-offload-kvcache is only supported when hicache-storage-backend is provided."
)
if cfg.disaggregation_decode_retraction_backup == "host_pool":
raise ValueError(
"The arguments disaggregation-decode-enable-offload-kvcache and "
"disaggregation-decode-retraction-backup=host_pool are mutually exclusive: "
"both build a decode host pool."
)
# Validate the effective ratio: model branches may declare a reset
# (e.g. Step3p forces 1.0 under hierarchical cache) that supersedes
# the user input before it ever takes effect.
if not (0 < resolved_view(server_args).swa_full_tokens_ratio <= 1.0):
raise ValueError("--swa-full-tokens-ratio should be in range (0, 1.0].")
def handle_unified_memory_pool(server_args: Any) -> None:
cfg = resolving_view(server_args)
if not cfg.enable_unified_memory:
return
if cfg.disaggregation_mode != "null":
# Constraints of the whole-envelope transfer; see
# UnifiedMLATokenToKVPool.get_contiguous_buf_infos.
assert cfg.disaggregation_transfer_backend == "mooncake", (
"--enable-unified-memory with PD disaggregation supports only "
"the mooncake transfer backend; got "
f"{cfg.disaggregation_transfer_backend!r}."
)
assert cfg.pp_size == 1, (
"--enable-unified-memory with PD disaggregation does not support "
"pipeline parallelism (whole-envelope transfer has no per-layer "
"entries to subset)."
)
assert not envs.SGLANG_DISABLE_LAZY_COMPACTION.get(), (
"--enable-unified-memory with PD disaggregation requires lazy "
"compaction; unset SGLANG_DISABLE_LAZY_COMPACTION."
)
assert not cfg.enable_hisparse, (
"--enable-unified-memory with PD disaggregation is not compatible "
"with --enable-hisparse: the decode-side HiSparse prealloc path "
"ships host/C4 rows straight from the allocator, bypassing the "
"virtual->physical translation the unified pool needs."
)
assert cfg.speculative_algorithm in (None, "DSPARK"), (
"--enable-unified-memory only supports --speculative-algorithm "
"DSPARK (chain draft); other speculative algorithms are not yet "
"audited for the unified pool's virtual/dense loc translation. Got "
f"--speculative-algorithm={cfg.speculative_algorithm!r}."
)
if cfg.speculative_algorithm == "DSPARK":
assert cfg.speculative_eagle_topk in (None, 1), (
"--enable-unified-memory + DSPARK supports a linear draft "
"chain only (--speculative-eagle-topk in {None, 1}); tree "
"verify is not audited for the unified pool. Got "
f"--speculative-eagle-topk={cfg.speculative_eagle_topk!r}."
)
# 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.discard(None)
assert spec_backends <= spec_allowed, (
"--enable-unified-memory + DSPARK requires spec-verify-audited "
f"attention backends {sorted(spec_allowed)} for both prefill "
f"and decode; got {sorted(spec_backends)}. flashinfer / fa3 do "
"not translate speculative verify indices to the unified "
"pool's dense space yet."
)
assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), (
"--enable-unified-memory is not yet compatible with hierarchical / "
"host-tiered KV cache (--enable-hierarchical-cache / --enable-lmcache): "
"the unified-memory-pool init wires up no host pools, and its device mamba / "
"full-attention slots are VIRTUAL — the host-offload path does not "
"translate them to physical."
)
assert cfg.dcp_size == 1, (
"--enable-unified-memory is not yet compatible with decode context "
"parallelism (--dcp-size > 1): the pool has no DCP-aware masked write "
"path (UnifiedMHATokenToKVPool.set_kv_buffer asserts dcp_kv_mask is None), "
"so a DCP run would boot and then fail on the first KV write."
)
# Only monolithic decode cuda-graph capture is wired; piecewise prefill
# capture is not. Guard when the user opts into it.
_cg_cfg = cfg.cuda_graph_config
if _cg_cfg is not None and _cg_cfg.prefill.backend == Backend.TC_PIECEWISE:
raise ValueError(
"--enable-unified-memory supports monolithic (decode) "
"cuda-graph capture only; disable piecewise prefill capture "
"(e.g. --cuda-graph-backend-prefill=disabled)."
)
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).
cfg = resolving_view(server_args)
if cfg.enable_unified_memory:
declare_resolution(
server_args,
"_handle_page_major_kv_layout",
enable_page_major_kv_layout=True,
)
if not cfg.enable_page_major_kv_layout:
return
# Only the Triton attention kernels read the strided 4-D envelope K/V
# views; FA3 / FlashInfer do not. EXCEPTION: the unified-memory MLA pool
# 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
# 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():
allowed_full = {
"triton",
"fa3",
"trtllm_mla",
"flashinfer",
"cutedsl_mla",
"tokenspeed_mla",
}
else:
allowed_full = {"triton"}
backends = set(server_args._resolved_attention_backends())
backends.discard(None)
assert backends <= allowed_full, (
"--enable-page-major-kv-layout requires the Triton attention backend "
"for the full-attention layers (unified-memory MLA also allows the "
f"paged MLA backends); got {sorted(backends)}, allowed "
f"{sorted(allowed_full)}. Pass a compatible --attention-backend."
)
# The Mamba/KDA state is stored in envelope-strided views; only
# stride-audited kernels may read it (Stage 4 audit, per slot):
# - decode: triton; flashinfer (recurrent_kda compiles the state slot
# stride as a free int64); helion (specializes KDA state strides 0-3
# and rejects a non-unit innermost stride); cutedsl (KDA fused sigmoid-
# gating update is stride-safe) on KDA-hybrid models only.
# - prefill: triton; flashkda (the wrapper gathers/scatters a contiguous
# per-slot copy); helion; cutedsl (kernel_h compiles h0/ht with dynamic
# int64 strides), with the same KDA-only caveat.
# - mamba (mamba2/short-conv state): triton only.
# use_mla_backend() distinguishes the KDA-hybrid family (K3/KimiLinear
# 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():
decode_allowed.update({"cutedsl", "helion"})
prefill_allowed.update({"cutedsl", "helion"})
resolved_linear_decode = cfg.linear_attn_decode_backend or cfg.linear_attn_backend
resolved_linear_prefill = cfg.linear_attn_prefill_backend or cfg.linear_attn_backend
assert resolved_linear_decode in decode_allowed | {None}, (
"--enable-page-major-kv-layout: linear-attention DECODE backend must "
f"be one of {sorted(decode_allowed)} for the strided conv/SSM state; "
f"got {resolved_linear_decode!r}."
)
assert resolved_linear_prefill in prefill_allowed | {None}, (
"--enable-page-major-kv-layout: linear-attention PREFILL backend must "
f"be one of {sorted(prefill_allowed)} for the strided conv/SSM state; "
f"got {resolved_linear_prefill!r}."
)
assert cfg.mamba_backend in (None, "triton"), (
"--enable-page-major-kv-layout requires the Triton Mamba kernels for "
f"the strided conv/SSM state; got {cfg.mamba_backend!r}. Pass "
"--mamba-backend triton."
)
def validate_prefill_only_disable_kv_cache_args(server_args: Any):
"""Validate --prefill-only-disable-kv-cache flag/precondition constraints.
Backend resolution is checked separately by
_handle_prefill_only_disable_kv_cache after backends settle.
"""
cfg = resolving_view(server_args)
if not cfg.prefill_only_disable_kv_cache:
return
# This flag is intentionally scoped to embedding mode for now. Other
# prefill-only paths (for example scoring and MIS) can benefit from
# the same idea later, but some of them still stage K/V through the
# paged cache today.
if not cfg.is_embedding:
raise ValueError(
"--prefill-only-disable-kv-cache currently requires --is-embedding. "
"Other prefill-only workloads may be supported in a future change once "
"their attention paths stop reading or writing the paged KV cache."
)
if cfg.kv_cache_dtype in ("nvfp4", "fp4_mx_block16"):
raise ValueError(
"--prefill-only-disable-kv-cache does not currently support "
"--kv-cache-dtype=nvfp4 or --kv-cache-dtype=fp4_mx_block16 because "
"the FP4 pool uses a separate allocation path."
)
if cfg.kv_cache_dtype == "mxfp8":
raise ValueError(
"--prefill-only-disable-kv-cache does not currently support "
"--kv-cache-dtype=mxfp8 because the MXFP8 pool stores separate "
"scale-factor buffers."
)
# Structural preconditions for the FA backend's fa_skip_kv_cache path,
# which is the only embedding path that doesn't read or write the pool:
# - chunked_prefill_size == -1 keeps a request in a single forward,
# so K/V never has to be reused across prefill chunks.
# - disable_radix_cache stops the prefix cache from indexing pool
# slots that no longer hold real data.
if cfg.chunked_prefill_size != -1:
raise ValueError(
"--prefill-only-disable-kv-cache requires --chunked-prefill-size=-1 so the FA "
"backend takes the fa_skip_kv_cache path; otherwise the pool would be touched "
"between prefill chunks."
)
if not cfg.disable_radix_cache:
raise ValueError(
"--prefill-only-disable-kv-cache requires --disable-radix-cache because the "
"radix cache indexes KV pool slots that no longer hold real data."
)
# 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:
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, "
"which the no-op pool intentionally rejects."
)
if cfg.enable_prefill_cp:
raise ValueError(
"--prefill-only-disable-kv-cache is incompatible with "
"--enable-prefill-cp: the prefill-CP path stages K/V through "
"the paged cache, which the no-op pool does not support."
)
# HiSparse selects a different pool class (HiSparseDSATokenToKVPool /
# HiSparseTokenToKVPoolAllocator) that is not the no-op pool.
if cfg.enable_hisparse:
raise ValueError(
"--prefill-only-disable-kv-cache is incompatible with --enable-hisparse: "
"HiSparse uses a dedicated pool family that is not the no-op MHA pool."
)
+220
View File
@@ -0,0 +1,220 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the LoRA adapters."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
resolving_view,
)
from sglang.srt.environ import envs
from sglang.srt.lora.lora_registry import LoRARef
logger = logging.getLogger(__name__)
def check_lora_server_args(server_args: Any):
cfg = resolving_view(server_args)
assert cfg.max_loras_per_batch > 0, "max_loras_per_batch must be positive"
# 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)
logger.warning(
"--enable-lora is set to True because --lora-paths is provided."
)
elif cfg.enable_lora is False:
logger.warning(
"--enable-lora is set to False, any provided lora_paths will be ignored."
)
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
)
if cfg.enable_lora_overlap_loading:
# TODO (glenliu21): use some sort of buffer with eviction instead of enforcing a limit
max_loaded_loras_limit = cfg.max_loras_per_batch * 2
assert (
cfg.max_loaded_loras is not None
and cfg.max_loaded_loras <= max_loaded_loras_limit
), (
"Enabling LoRA overlap loading requires pinning LoRA adapter weights in CPU memory, "
f"so --max-loaded-loras must be less than or equal to double --max-loras-per-batch: {max_loaded_loras_limit}"
)
# Validate compatibility with speculative decoding
server_args._check_lora_speculative_compatibility()
# Parse lora_paths
if isinstance(cfg.lora_paths, list):
parsed_lora_paths = []
for lora_path in cfg.lora_paths:
if isinstance(lora_path, str):
if "=" in lora_path:
name, path = lora_path.split("=", 1)
lora_ref = LoRARef(
lora_id=LoRARef.deterministic_id(name, path),
lora_name=name,
lora_path=path,
pinned=False,
)
else:
lora_ref = LoRARef(
lora_id=LoRARef.deterministic_id(lora_path, lora_path),
lora_name=lora_path,
lora_path=lora_path,
pinned=False,
)
elif isinstance(lora_path, dict):
assert (
"lora_name" in lora_path and "lora_path" in lora_path
), f"When providing LoRA paths as a list of dict, each dict should contain 'lora_name' and 'lora_path' keys. Got: {lora_path}"
lora_ref = LoRARef(
lora_id=LoRARef.deterministic_id(
lora_path["lora_name"], lora_path["lora_path"]
),
lora_name=lora_path["lora_name"],
lora_path=lora_path["lora_path"],
pinned=lora_path.get("pinned", False),
)
else:
raise ValueError(
f"Invalid type for item in --lora-paths list: {type(lora_path)}. "
"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
)
elif isinstance(cfg.lora_paths, dict):
server_args._late_resolution(
"check_lora_server_args",
lora_paths=[
LoRARef(
lora_id=LoRARef.deterministic_id(k, v),
lora_name=k,
lora_path=v,
pinned=False,
)
for k, v in cfg.lora_paths.items()
],
)
elif cfg.lora_paths is None:
server_args._late_resolution("check_lora_server_args", lora_paths=[])
else:
raise ValueError(
f"Invalid type for --lora-paths: {type(cfg.lora_paths)}. "
"Expected a list or a dictionary."
)
# 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(
"check_lora_server_args",
lora_target_modules=set(cfg.lora_target_modules),
)
if "all" in cfg.lora_target_modules:
assert (
len(cfg.lora_target_modules) == 1
), "If 'all' is specified in --lora-target-modules, it should be the only module specified."
# Ensure sufficient information is provided for LoRA initialization.
assert cfg.lora_paths or (
cfg.max_lora_rank and cfg.lora_target_modules
), "When no initial --lora-paths is provided, you need to specify both --max-lora-rank and --lora-target-modules for LoRA initialization."
# Validate max_loaded_loras
if cfg.max_loaded_loras is not None:
assert cfg.max_loaded_loras >= cfg.max_loras_per_batch, (
"max_loaded_loras should be greater than or equal to max_loras_per_batch. "
f"max_loaded_loras={cfg.max_loaded_loras}, max_loras_per_batch={cfg.max_loras_per_batch}"
)
assert len(cfg.lora_paths) <= cfg.max_loaded_loras, (
"The number of LoRA paths should not exceed max_loaded_loras. "
f"max_loaded_loras={cfg.max_loaded_loras}, lora_paths={len(cfg.lora_paths)}"
)
if cfg.max_lora_chunk_size is not None:
assert (
16 <= cfg.max_lora_chunk_size <= 128
and (cfg.max_lora_chunk_size & (cfg.max_lora_chunk_size - 1)) == 0
), "--max-lora-chunk-size must be a power of 2 between 16 and 128."
if cfg.lora_use_virtual_experts:
logger.info("Virtual expert computation enabled.")
assert (
cfg.lora_drain_wait_threshold >= 0.0
), "--lora-drain-wait-threshold must be non-negative."
def check_lora_speculative_compatibility(server_args: Any):
"""Validate LoRA + speculative decoding combinations.
Adapters apply to the target only; a shared draft runs unadapted.
Matches resolved algorithm names (NEXTN has collapsed to EAGLE).
"""
cfg = resolving_view(server_args)
if cfg.speculative_algorithm in ["NGRAM", None]:
return
# These algorithms present a uniform per-request token width during
# verify, which is what the LoRA segment layout assumes.
lora_spec_algorithms = ("EAGLE", "EAGLE3", "DFLASH", "DSPARK")
if cfg.speculative_algorithm not in lora_spec_algorithms:
promoted = (
" (NEXTN/EAGLE with a Gemma4 assistant draft is automatically "
"promoted to FROZEN_KV_MTP, which does not support LoRA)"
if cfg.speculative_algorithm == "FROZEN_KV_MTP"
else ""
)
raise ValueError(
"LoRA is only compatible with NGRAM, EAGLE, NEXTN, EAGLE3, "
"DFLASH, or DSPARK speculative decoding, not "
f"{cfg.speculative_algorithm}{promoted}."
)
ragged_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get()
# Each entry: (is unsupported, why). Reasons are appended to a shared
# prefix so the message names the combination, not just the flag.
unsupported = [
(
cfg.speculative_algorithm == "DSPARK" and ragged_mode != "static",
f"does not support SGLANG_RAGGED_VERIFY_MODE={ragged_mode!r}: "
"the per-request verify lengths it schedules break the "
"uniform-width LoRA segment layout",
),
(
cfg.speculative_adaptive,
"does not support --speculative-adaptive: the draft is built "
"from a static ServerArgs snapshot, and the runtime-state "
"swap does not rebuild LoRA cuda-graph metadata",
),
(
"experimental_sgl_trtllm"
in (cfg.moe_runner_backend, cfg.speculative_moe_runner_backend),
"does not support the experimental_sgl_trtllm MoE runner: its "
"TopK reads the LoRA config per forward, which the draft "
"resolves against the target's after its own publish ended",
),
(
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(),
"does not support SGLANG_ENABLE_OVERLAP_PLAN_STREAM=1: LoRA "
"batch preparation would run on the plan stream, unordered "
"against in-flight forwards",
),
]
for is_unsupported, reason in unsupported:
if is_unsupported:
raise ValueError(
f"LoRA with EAGLE/NEXTN/EAGLE3 speculative decoding {reason}."
)
+154
View File
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the Mamba / linear-attention backends."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
resolving_view,
)
from sglang.srt.utils.common import (
is_cuda,
is_flashinfer_available,
is_hip,
is_musa,
is_npu,
is_sm100_supported,
is_xpu,
)
logger = logging.getLogger(__name__)
def handle_mamba_backend(server_args: Any):
cfg = resolving_view(server_args)
if cfg.mamba_cache_philox_rounds < 0:
raise ValueError("--mamba-cache-philox-rounds must be non-negative.")
if cfg.mamba_max_states_per_path == 0 or cfg.mamba_max_states_per_path < -1:
raise ValueError(
"--mamba-max-states-per-path must be -1 (unlimited) or a positive "
f"integer, got {cfg.mamba_max_states_per_path}."
)
if cfg.enable_mamba_cache_stochastic_rounding:
if cfg.mamba_ssm_dtype != "float16":
raise ValueError(
"Stochastic rounding for the Mamba SSM cache requires "
f"--mamba-ssm-dtype float16, got {cfg.mamba_ssm_dtype!r}. "
"Run with --mamba-ssm-dtype float16 or disable "
"--enable-mamba-cache-stochastic-rounding."
)
if not is_cuda():
raise ValueError(
"Stochastic rounding for the Mamba SSM cache is only "
"supported on NVIDIA CUDA platforms. Disable "
"--enable-mamba-cache-stochastic-rounding on this platform."
)
if cfg.mamba_backend == "triton" and not is_sm100_supported():
raise ValueError(
"Stochastic rounding for the Mamba SSM cache with "
"--mamba-backend triton requires SM100 with CUDA >= 12.8 "
"because it uses the cvt.rs.f16x2.f32 PTX instruction. On "
"H100/SM90, run with --mamba-backend flashinfer "
"--mamba-ssm-dtype float16, or disable "
"--enable-mamba-cache-stochastic-rounding."
)
if cfg.mamba_backend == "flashinfer":
flashinfer_error = (
"FlashInfer mamba module not available, please check the "
"FlashInfer installation."
)
if cfg.enable_mamba_cache_stochastic_rounding:
flashinfer_error += (
" Stochastic rounding with --mamba-backend flashinfer "
"requires FlashInfer Mamba and --mamba-ssm-dtype float16."
)
if is_flashinfer_available():
try:
import flashinfer.mamba # noqa: F401
logger.info("Successfully imported FlashInfer mamba module")
except (ImportError, AttributeError):
raise ValueError(flashinfer_error)
else:
raise ValueError(flashinfer_error)
def handle_int8_mamba_checkpoint(server_args: Any):
# The int8 mamba checkpoint pool is only wired into the built-in
# MambaRadixCache. The host-offload path (enabled by
# --enable-hierarchical-cache) and custom radix-cache backends are NOT
# int8-aware: they would read int8 checkpoint slots as bf16 active slots
# (wrong pool / out-of-range). Reject the combination up front rather than
# silently corrupting state.
cfg = resolving_view(server_args)
if not cfg.enable_int8_mamba_checkpoint:
return
if cfg.enable_hierarchical_cache:
raise ValueError(
"--enable-int8-mamba-checkpoint is not supported together with "
"--enable-hierarchical-cache: the host-offload path "
"is not int8-aware. Disable one of them."
)
if cfg.radix_cache_backend is not None:
raise ValueError(
"--enable-int8-mamba-checkpoint only supports the built-in mamba "
f"radix cache; --radix-cache-backend={cfg.radix_cache_backend!r} "
"is not int8-aware. Omit --radix-cache-backend."
)
def validate_mamba_extra_buffer(view, model_arch: str, *, mamba_cache_chunk_size_of):
from sglang.srt.arg_groups.overrides import supports_mamba_cache_extra_buffer
assert supports_mamba_cache_extra_buffer(
view, model_arch
), f"extra_buffer is not supported for {model_arch}; use no_buffer."
assert (
is_cuda() or is_musa() or is_npu() or is_hip() or is_xpu()
), "extra_buffer needs CUDA/MUSA/NPU/ROCm/XPU (FLA)."
if view.mamba_radix_cache_strategy == "extra_buffer_lazy":
# The PD-disagg decode pool is not wired for lazy slots.
assert view.disaggregation_mode == "null", (
"extra_buffer_lazy unsupported under PD disaggregation; use "
"--mamba-radix-cache-strategy extra_buffer."
)
# eagle/ngram/dspark/dflash all verify through
# prepare_mamba_track_for_verify (lazy plan wired); dflash gained
# the hook in DFlashVerifyInput.prepare_for_verify.
if view.speculative_num_draft_tokens is not None:
assert view.mamba_track_interval >= view.speculative_num_draft_tokens
if view.page_size is not None:
assert view.mamba_track_interval % view.page_size == 0
# Called here and not passed in: `mamba_cache_chunk_size` derives from
# `page_size`, which resolution writes after this validator runs, so
# evaluating it at the call site raises on the unresolved `None`.
mamba_cache_chunk_size = mamba_cache_chunk_size_of()
assert mamba_cache_chunk_size is not None
if (
view.chunked_prefill_size is not None
and 0 < view.chunked_prefill_size < mamba_cache_chunk_size
):
logger.warning(
"Mamba radix extra-buffer is enabled with chunked_prefill_size=%s "
"smaller than mamba_cache_chunk_size=%s. This can make "
"mamba_track_mask false for unfinished chunked-prefill handoff "
"and skip Mamba state checkpoints.",
view.chunked_prefill_size,
mamba_cache_chunk_size,
)
def validate_mamba_no_buffer(view, model_arch: str):
assert view.page_size in (1, None), "no_buffer only supports page_size=1."
assert (
view.disable_overlap_schedule
), "no_buffer do not support overlap schedule. Try to set disable_overlap_schedule=True."
assert (
view.attention_backend != "trtllm_mha"
), "no_buffer do not support trtllm_mha attention backend."
+268
View File
@@ -0,0 +1,268 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the GPU memory budget."""
from __future__ import annotations
import copy
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolving_view,
)
from sglang.srt.environ import envs
logger = logging.getLogger(__name__)
def handle_gpu_memory_settings(server_args: Any, gpu_mem):
"""
Configure GPU memory-dependent settings including
chunked_prefill_size, cuda_graph_config[decode].max_bs, and mem_fraction_static.
Here are our heuristics:
- Set chunked_prefill_size and cuda_graph_config[decode].max_bs based on the GPU memory capacity.
This is because GPUs with more memory are generally more powerful, we need to use a larger
chunked_prefill_size and a larger decode max_bs to fully utilize the GPU.
- Then set mem_fraction_static based on chunked_prefill_size and decode max_bs.
GPU memory capacity = model weights + KV cache pool + activations + cuda graph buffers
The argument mem_fraction_static is defined as (model weights + KV cache pool) / GPU memory capacity,
or equivalently, mem_fraction_static = (GPU memory capacity - activations - cuda graph buffers) / GPU memory capacity.
In order to compute mem_fraction_static, we need to estimate the size of activations and cuda graph buffers.
The activation memory is proportional to the chunked_prefill_size.
The cuda graph memory is proportional to the decode max_bs.
We use reserved_mem = chunked_prefill_size * 1.5 + max_bs * 2 to estimate the size of activations and cuda graph buffers in GB,
and set mem_fraction_static = (GPU memory capacity - reserved_mem) / GPU memory capacity.
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.
"""
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)
decode_cuda_graph_config = cuda_graph_config.decode
prefill_cuda_graph_config = cuda_graph_config.prefill
if gpu_mem is not None:
if gpu_mem < 20 * 1024:
# T4, 4080
# (chunked_prefill_size 2k, max_bs 8)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=2048,
)
if decode_cuda_graph_config.max_bs is None:
decode_cuda_graph_config.max_bs = 8
elif gpu_mem < 35 * 1024:
# A10, 4090, 5090
# (chunked_prefill_size 2k, max_bs 24 if tp < 4 else 80)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=2048,
)
if decode_cuda_graph_config.max_bs is None:
if cfg.tp_size < 4:
decode_cuda_graph_config.max_bs = 24
else:
decode_cuda_graph_config.max_bs = 80
elif gpu_mem < 60 * 1024:
# A100 (40GB), L40,
# (chunked_prefill_size 4k, max_bs 32 if tp < 4 else 160)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=4096,
)
if decode_cuda_graph_config.max_bs is None:
if cfg.tp_size < 4:
decode_cuda_graph_config.max_bs = 32
else:
decode_cuda_graph_config.max_bs = 160
elif gpu_mem < 90 * 1024:
# H100, A100
# (chunked_prefill_size 8k, max_bs 256 if tp < 4 else 512)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=8192,
)
if decode_cuda_graph_config.max_bs is None:
if cfg.tp_size < 4:
decode_cuda_graph_config.max_bs = 256
else:
decode_cuda_graph_config.max_bs = 512
elif gpu_mem < 160 * 1024:
# H20, H200
# (chunked_prefill_size 8k, max_bs 256 if tp < 4 else 512)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=8192,
)
if decode_cuda_graph_config.max_bs is None:
if cfg.tp_size < 4:
decode_cuda_graph_config.max_bs = 256
else:
decode_cuda_graph_config.max_bs = 512
else:
# B200, MI300
# (chunked_prefill_size 16k, max_bs 512)
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=16384,
)
if decode_cuda_graph_config.max_bs is None:
decode_cuda_graph_config.max_bs = 512
else:
# Fallback defaults when gpu_mem is None
if cfg.chunked_prefill_size is None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
chunked_prefill_size=4096,
)
if decode_cuda_graph_config.max_bs is None:
decode_cuda_graph_config.max_bs = 160
# 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
)
)
else:
decode_cuda_graph_config.max_bs = max(decode_cuda_graph_config.bs)
else:
# Reuse decode_cuda_graph_config.bs for cpu graph and use torch_compile_max_bs for cpu graph batch size limit,
# as cpu graph is based on torch.compile
if decode_cuda_graph_config.bs is not None:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
torch_compile_max_bs=max(decode_cuda_graph_config.bs),
)
else:
# If decode_cuda_graph_config.bs is not set, we will preferentially use torch_compile_max_bs
# to generate decode_cuda_graph_config.bs
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
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()
assert (
cfg.torch_compile_max_bs > 0
), "cuda_graph_config[decode].bs should contain positive batch sizes"
decode_cuda_graph_config.max_bs = cfg.torch_compile_max_bs
if prefill_cuda_graph_config.max_bs is None:
# 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():
prefill_cuda_graph_config.max_bs = cfg.chunked_prefill_size
else:
prefill_cuda_graph_config.max_bs = 2048
# If max_total_tokens is set, cap prefill max_bs to not exceed max_total_tokens.
if cfg.max_total_tokens is not None:
prefill_cuda_graph_config.max_bs = min(
prefill_cuda_graph_config.max_bs, cfg.max_total_tokens
)
# For Llama2 series models, max_bs is limited to 4096.
# TODO(yuwei): remove this after the issue is fixed
if "llama-2" in cfg.model_path.lower():
prefill_cuda_graph_config.max_bs = min(
prefill_cuda_graph_config.max_bs, 4096
)
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
)
)
if cuda_graph_config != cfg.cuda_graph_config:
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
cuda_graph_config=cuda_graph_config,
)
if cfg.mem_fraction_static is None:
if server_args.post_capture_kv_sizing_planned():
# Post-capture sizing measures free memory after graph capture, so
# skip the graph/activation reserve; keep only the floor + parallel slack.
reserved_mem = 1536
reserved_mem += cfg.tp_size * cfg.pp_size / 8 * 1024
else:
# Tokens the activation working set scales with (per serving mode).
if cfg.disaggregation_mode == "decode":
running_requests = (
cfg.max_running_requests or decode_cuda_graph_config.max_bs or 1
)
draft_tokens = cfg.speculative_num_draft_tokens or 1
activation_tokens = max(running_requests * draft_tokens, 2048)
elif cfg.chunked_prefill_size > 0:
activation_tokens = max(cfg.chunked_prefill_size, 2048)
else:
activation_tokens = max(cfg.max_prefill_tokens, 2048)
# Constant meta data (e.g., from attention backend) + activation slack.
reserved_mem = 512
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()
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()
declare_resolution(
server_args,
"_handle_gpu_memory_settings",
mem_fraction_static=(
round((gpu_mem - reserved_mem) / gpu_mem, 3)
if gpu_mem is not None
else 0.88
),
)
# Multimodal models need more memory for the image processing,
# 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()
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)
# 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():
envs.SGLANG_SYMM_MEM_PREALLOC_GB_SIZE.set(4)
logger.warning(
"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."
)
+856
View File
@@ -0,0 +1,856 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for per-model and per-capability adjustments."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch
from sglang.srt.connector import ConnectorType
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, Phase, with_phase
from sglang.srt.utils.common import (
get_quantization_config,
is_cuda,
is_hip,
is_mps,
is_npu,
is_sm90_supported,
is_sm100_supported,
is_sm120_supported,
is_xpu,
parse_connector_type,
)
logger = logging.getLogger(__name__)
def handle_model_specific_adjustments(server_args: Any):
cfg = resolving_view(server_args)
from sglang.srt.configs.model_config import (
get_mimo_v2_fused_qkv_expected_tp_size,
is_deepseek_dsa,
)
if cfg.enable_deterministic_inference:
declare_resolution(
server_args,
"_handle_model_specific_adjustments",
enforce_disable_flashinfer_allreduce_fusion=True,
)
declare_resolution(
server_args,
"_handle_model_specific_adjustments",
uses_mamba_radix_cache=False,
)
if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE:
# No model overrides for an instance connector: no hf_config to
# key them on.
return
model_config = server_args.get_model_config()
hf_config = model_config.hf_config
model_arch = hf_config.architectures[0]
if model_arch == "InternS2MobiusForConditionalGeneration":
unsupported = []
if cfg.pp_size != 1:
unsupported.append("pipeline parallelism (--pp-size must be 1)")
if cfg.ep_size != 1:
unsupported.append("expert parallelism (--ep-size must be 1)")
if unsupported:
raise ValueError(
"Intern-S2-Mobius does not support: " + "; ".join(unsupported) + "."
)
if cfg.enable_dsa_cache_layer_split and not is_deepseek_dsa(hf_config):
raise ValueError(
"--enable-dsa-cache-layer-split is only supported for DSA "
"(DeepSeek Sparse Attention) models."
)
if cfg.enable_cp_decode_attn_tp:
from sglang.srt.layers.cp.cp_decode_attn_tp import (
CP_DECODE_ATTN_TP_SUPPORTED_ARCHS,
)
if model_arch not in CP_DECODE_ATTN_TP_SUPPORTED_ARCHS:
raise ValueError(
"--enable-cp-decode-attn-tp is only supported for models "
"whose attention linears are replicated across CP ranks "
f"(attn_tp_size=1). Got {model_arch}; supported: "
f"{sorted(CP_DECODE_ATTN_TP_SUPPORTED_ARCHS)}."
)
_hybrid_spec = get_linear_attn_spec_by_arch(model_arch)
if _hybrid_spec is not None and _hybrid_spec.uses_mamba_radix_cache:
server_args._handle_mamba_radix_cache(model_arch=model_arch)
# Collect the declarative model overrides (registry) on the
# pristine config and stash them for publish-time flags resolution;
# server_args is never mutated — mid-resolution readers see the
# declared values through resolved_view, runtime readers through the
# flags tier.
from sglang.srt.arg_groups.overrides import (
collect_model_override_declarations,
validate_declarations,
)
model_overrides = collect_model_override_declarations(
model_arch, server_args, hf_config
)
validate_declarations(server_args, model_overrides)
server_args._resolved_overrides.extend(model_overrides)
if model_arch in (
"KimiLinearForCausalLM",
"KimiK3ForConditionalGeneration",
):
from sglang.srt.arg_groups.kimi_k3_hook import (
apply_kimi_k3_linear_attn_defaults,
apply_kimi_k3_spec_backend_defaults,
)
apply_kimi_k3_linear_attn_defaults(server_args)
apply_kimi_k3_spec_backend_defaults(server_args)
if model_arch in [
"DeepseekV4ForCausalLM",
]:
from sglang.srt.arg_groups.deepseek_v4_hook import (
apply_deepseek_v4_defaults,
)
apply_deepseek_v4_defaults(server_args, model_arch)
if model_arch in [
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"KimiK25ForConditionalGeneration",
"MistralLarge3ForCausalLM",
"PixtralForConditionalGeneration",
"GlmMoeDsaForCausalLM",
"LongcatFlashForCausalLM",
"Dots3NoteForCausalLM",
]:
# Set attention backend for DeepSeek
if is_deepseek_dsa(hf_config): # DeepSeek 3.2/GLM 5
if envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.is_set():
logger.warning(
f"Dense attention kv len threshold is manually set to {envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get()} for DSA. Caution: This may cause performance regression if the threshold is larger than the index topk of model."
)
else:
# When threshold is not manually set, set it to the index topk of model
from sglang.srt.configs.model_config import get_dsa_index_topk
envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.set(
get_dsa_index_topk(hf_config)
)
logger.warning(
f"Set dense attention kv len threshold to model index_topk={envs.SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD.get()} for DeepSeek with DSA."
)
# The "dsa" attention fill moved to the override registry
# (arg_groups/overrides.py: _deepseek_family_overrides).
index_topk_freq = getattr(hf_config, "index_topk_freq", 1) or 1
index_topk_pattern = getattr(hf_config, "index_topk_pattern", None)
if cfg.enable_two_batch_overlap and (
index_topk_freq > 1
or (index_topk_pattern is not None and "S" in index_topk_pattern)
):
raise ValueError(
"--enable-two-batch-overlap is not supported with DSA "
"index-topk sharing (index_topk_freq > 1 or an "
"index_topk_pattern containing shared layers): the TBO op "
"path does not propagate topk indices across layers, so "
"shared layers would run sparse attention without indices."
)
if not is_npu() and not is_xpu(): # CUDA or ROCm GPU
if cfg.enable_prefill_cp:
# The DSA CP field declarations moved to the override
# registry (arg_groups/overrides.py:
# _deepseek_family_overrides).
declare_resolution(
server_args,
"_handle_model_specific_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config,
Phase.PREFILL,
backend=Backend.DISABLED,
),
)
else:
# Pure TP and partial DP Attention mode is active for DSA, logging a warning
if cfg.dp_size < cfg.tp_size:
logger.warning(
f"DSA with TP mode is active, dp_size={cfg.dp_size}, tp_size={cfg.tp_size}, "
f"attn_tp_size={cfg.tp_size}, attention weights will be sharded across {cfg.tp_size} ranks."
)
# The DSA page-size selection moved to the override registry
# (arg_groups/overrides.py: _deepseek_family_overrides).
import torch
major, _ = torch.cuda.get_device_capability()
server_args._set_default_dsa_kv_cache_dtype(
major, resolved_view(server_args).quantization
)
server_args._set_default_dsa_backends(major)
if cfg.enable_prefill_cp:
assert (
cfg.disaggregation_mode != "decode"
), "CP is only supported for prefill when PD disaggregation, please remove --enable-prefill-cp."
if (
cfg.enable_dsa_cache_layer_split
and cfg.disaggregation_mode != "prefill"
):
if cfg.disaggregation_mode == "decode":
raise ValueError(
"--enable-dsa-cache-layer-split is not supported on "
"decode workers. This flag is a prefill-CP "
"optimization; decode receives full cache shards "
"through PD transfer."
)
raise ValueError(
"--enable-dsa-cache-layer-split is only supported on PD "
"prefill workers. Non-PD workers also run decode and "
"require ordinary local decode cache semantics."
)
if cfg.enable_dsa_cache_layer_split and (
not cfg.enable_prefill_cp or cfg.cp_strategy != "interleave"
):
raise ValueError(
"--enable-dsa-cache-layer-split requires "
"--enable-prefill-cp and --cp-strategy interleave "
"(or legacy --enable-nsa-prefill-context-parallel with "
"--nsa-prefill-cp-mode round-robin-split)."
)
# Layer split relies on the mooncake all-CP-rank KV/indexer
# transfer path. mori/nixl support is a temporary limitation
# and will be added later by the community.
if (
cfg.enable_dsa_cache_layer_split
and cfg.disaggregation_transfer_backend != "mooncake"
):
raise ValueError(
"--enable-dsa-cache-layer-split currently only supports "
"the mooncake transfer backend (mooncake / mooncake_tcp). "
f"Got --disaggregation-transfer-backend "
f"{cfg.disaggregation_transfer_backend!r}. mori/nixl "
"support will be added later by the community."
)
if cfg.enable_dsa_cache_layer_split and cfg.pp_size > 1:
raise ValueError(
"--enable-dsa-cache-layer-split is not supported with "
"pipeline parallelism (pp_size > 1) yet. It requires "
"prefill context parallelism, and CP + PP has not been "
"validated for this feature."
)
else:
# DeepSeek V3/R1/V3.1
if cfg.cuda_graph_config.prefill.backend != Backend.DISABLED:
logger.info("Piecewise CUDA graph is enabled, use MLA for prefill.")
# The sm100 trtllm_mla fill moved to the override registry
# (arg_groups/overrides.py: _deepseek_family_overrides).
# 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():
declare_resolution(
server_args,
"_handle_model_specific_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config,
Phase.PREFILL,
backend=Backend.DISABLED,
),
)
# Set moe backend for DeepSeek: the sm100 quant/moe resolution
# moved to the resolution pipeline (arg_groups/overrides.py:
# _deepseek_moe_quant_resolution -- a slot pass, because the DSA
# kv-cache-dtype default above must read the pristine
# quantization). The HIP arm (fusion log + spec_moe writes, the
# latter awaiting the speculative-hook migration) stays below.
from sglang.srt.arg_groups.overrides import (
_deepseek_moe_quant_resolution,
run_post_process_pass,
)
run_post_process_pass(server_args, _deepseek_moe_quant_resolution)
if is_hip():
if is_deepseek_dsa(hf_config):
# The fused top-k v2 kernel (topk_transform_512_v2) is a
# CUDA/Hopper-only path: its JIT source includes
# <cooperative_groups.h> and uses cg::this_cluster()
# (thread-block clusters), neither of which exists on ROCm,
# so it fails to JIT-compile on gfx9xx during CUDA-graph
# capture. DeepSeek-V4 already disables it on HIP; mirror that
# 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:
# TODO (Hubert): Put this back later
# server_args.enable_aiter_allreduce_fusion = True
logger.info("Enable Aiter AllReduce Fusion for DeepseekV3ForCausalLM")
# The fp4-checkpoint draft spec-MoE resolution moved to the
# resolution pipeline (arg_groups/overrides.py:
# _deepseek_spec_moe_resolution), invoked here at its legacy
# slot.
from sglang.srt.arg_groups.overrides import (
_deepseek_spec_moe_resolution,
)
run_post_process_pass(server_args, _deepseek_spec_moe_resolution)
elif model_arch in [
"DeepseekV4ForCausalLM",
]:
from sglang.srt.arg_groups.deepseek_v4_hook import (
validate_deepseek_v4_cp,
validate_deepseek_v4_mega_moe_token_budget,
)
validate_deepseek_v4_cp(server_args)
validate_deepseek_v4_mega_moe_token_budget(server_args)
if is_sm120_supported():
# SM120 lacks tcgen05/TMEM: disable features that depend on
# DeepGEMM or require >99KB SMEM (topk_v2).
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
envs.SGLANG_OPT_USE_TOPK_V2.set(False)
envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.set(False)
if not envs.SGLANG_OPT_FUSE_MHC_POST_PRE.is_set():
envs.SGLANG_OPT_FUSE_MHC_POST_PRE.set(True)
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.set(True)
# Prefer TileLang over the Torch fallback.
envs.SGLANG_OPT_USE_TILELANG_INDEXER.set(True)
elif is_hip():
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.set(False)
envs.SGLANG_OPT_USE_TOPK_V2.set(True)
envs.SGLANG_OPT_USE_AITER_INDEXER.set(True)
envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.set(False)
envs.SGLANG_OPT_USE_TILELANG_MHC_POST.set(False)
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.set(True)
envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.set(False)
envs.SGLANG_EAGER_INPUT_NO_COPY.set(True)
elif model_arch in ["GptOssForCausalLM"]:
# Attention backend selection + XPU dtype validation moved to the
# override registry (arg_groups/overrides.py: _gpt_oss_overrides).
# Exempt MLX only: none of these backends exist on MPS, and MLX runs
# attention inside its own runner, so attention_backend is still
# unset here. Plain macOS stays on the list -- torch_native has
# neither sliding window nor attention sinks.
if not (is_mps() and use_mlx()):
supported_backends = [
"triton",
"trtllm_mha",
"fa3",
"fa4",
"ascend",
"intel_amx",
"intel_xpu",
"aiter",
]
prefill_attn_backend, decode_attn_backend = (
server_args._resolved_attention_backends()
)
assert (
prefill_attn_backend in supported_backends
and decode_attn_backend in supported_backends
), (
f"GptOssForCausalLM requires one of {supported_backends} attention backend, but got the following backends\n"
f"- Prefill: {prefill_attn_backend}\n"
f"- Decode: {decode_attn_backend}\n"
)
quant_method = get_quantization_config(hf_config)
is_mxfp4_quant_format = quant_method == "mxfp4"
if (
not server_args._resolved().enable_dp_attention
and cfg.nnodes == 1
and is_hip()
):
# TODO (Hubert): Put this back later
# server_args.enable_aiter_allreduce_fusion = True
logger.info("Enable Aiter AllReduce Fusion for GptOssForCausalLM")
quantization_config = getattr(hf_config, "quantization_config", None)
is_mxfp4_quant_format = (
quantization_config is not None
and quantization_config.get("quant_method") == "mxfp4"
)
# The mxfp4 dtype override moved to the override registry
# (arg_groups/overrides.py: _gpt_oss_overrides).
# The moe_runner_backend selection moved to the override registry
# (arg_groups/overrides.py: _gpt_oss_overrides).
if resolved_view(server_args).moe_runner_backend == "triton_kernel":
assert (
server_args._resolved().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()
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 (
expected_attn_tp_size is not None
and expected_attn_tp_size % effective_attn_tp_size != 0
):
raise ValueError(
"MiMoV2ForCausalLM requires effective attention TP "
f"size {expected_attn_tp_size} because its fused "
"qkv_proj weights are "
f"TP={expected_attn_tp_size}-interleaved; got "
f"{effective_attn_tp_size} "
f"(tp_size={cfg.tp_size}, dp_size={cfg.dp_size}, "
f"enable_dp_attention={view.enable_dp_attention}, "
f"attn_cp_size={view.attn_cp_size}). "
"Set --tp, --dp, --enable-dp-attention, and "
"--attention-context-parallel-size so the effective "
f"attention TP size is {expected_attn_tp_size}."
)
# enable_multi_layer_eagle for EAGLE moved to the override registry
# (arg_groups/overrides.py: _mimo_v2_overrides).
# MiMoV2 hierarchical cache runs on the unified radix tree, which
# is the default tree cache now. MiMoV2 has head_dim != v_head_dim,
# so the host KV pool uses asymmetric K/V allocation. Both
# kernel/page_first and direct/page_first_direct have split K/V
# transfer paths.
elif (
"Step3p5ForCausalLM" in model_arch
or "Step3p7ForConditionalGeneration" in model_arch
):
# Attention backend selection + EAGLE multi-layer +
# hierarchical-cache SWA writes moved to the override registry
# (arg_groups/overrides.py: _step3p_overrides).
pass
elif (
model_arch in ("Llama4ForConditionalGeneration", "Llama4ForCausalLM")
and cfg.device != "cpu"
):
# Attention backend auto-select moved to the override registry
# (arg_groups/overrides.py: _llama4_overrides).
attention_backend = resolved_view(server_args).attention_backend
assert attention_backend in {
"fa3",
"aiter",
"triton",
"ascend",
"trtllm_mha",
"intel_xpu",
}, f"fa3, aiter, triton, ascend, trtllm_mha or intel_xpu is required for Llama4 model but got {attention_backend}"
# The moe_runner_backend selection moved to the override registry
# (arg_groups/overrides.py: _llama4_overrides).
# Gemma2/Gemma3 (disable_hybrid_swa_memory) moved to the override registry
# (arg_groups/overrides.py: _gemma2_gemma3_overrides).
elif model_arch in (
"Gemma4ForConditionalGeneration",
"Gemma4ForCausalLM",
"Gemma4UnifiedForConditionalGeneration",
):
# Default attention backend selection moved to the override registry
# (arg_groups/overrides.py: _gemma4_overrides).
prefill_backend, decode_backend = server_args._resolved_attention_backends()
accepted_backends = (
"trtllm_mha",
"triton",
"ascend",
"intel_xpu",
"intel_amx",
)
assert (
prefill_backend in accepted_backends and decode_backend in accepted_backends
), (
"Gemma4 only supports trtllm_mha, triton, ascend, intel_xpu, or intel_amx "
f"attention backend, got prefill={prefill_backend}, decode={decode_backend}"
)
# The quantization/moe_runner_backend resolution moved to the override
# registry (arg_groups/overrides.py: _gemma4_overrides).
elif model_arch == "MossVLForConditionalGeneration":
# The prefill attention backend default + validation moved to the
# override registry (arg_groups/overrides.py: _moss_vl_overrides).
pass
elif model_arch in ["Exaone4ForCausalLM", "ExaoneMoEForCausalLM"]:
if hf_config.sliding_window_pattern is not None:
# disable_hybrid_swa_memory moved to the override registry
# (arg_groups/overrides.py: _exaone_overrides).
# https://docs.sglang.ai/advanced_features/attention_backend.html
accepted_backends = ["fa3", "triton", "trtllm_mha"]
attention_backend = resolved_view(server_args).attention_backend
assert (
attention_backend in accepted_backends
), f"One of the attention backends in {accepted_backends} is required for {model_arch}, but got {attention_backend}"
elif model_arch in ["Olmo2ForCausalLM"]:
# disable_hybrid_swa_memory + attention backend selection moved to
# the override registry (arg_groups/overrides.py: _olmo2_overrides).
# Flashinfer appears to degrade performance when sliding window attention
# is used for the Olmo2 architecture. Olmo2 does not use sliding window attention
# but Olmo3 does.
attention_backend = resolved_view(server_args).attention_backend
assert (
attention_backend != "flashinfer"
), "FlashInfer backend can significantly degrade the performance of Olmo3 models."
logger.info(f"Using {attention_backend} as attention backend for {model_arch}.")
elif model_arch in [
"Qwen3MoeForCausalLM",
"Qwen3VLMoeForConditionalGeneration",
"Qwen3NextForCausalLM",
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
]:
# The quantization/moe_runner_backend resolution moved to the
# override registry (arg_groups/overrides.py:
# _qwen3_moe_family_overrides); the hybrid sub-family's attention
# backend + page size defaults to _qwen3_5_hybrid_overrides.
pass
elif model_arch in ["Glm4MoeForCausalLM"]:
# The quantization/moe_runner_backend/enable_tf32_matmul resolution
# moved to the override registry (arg_groups/overrides.py:
# _glm4_moe_overrides).
pass
elif model_arch in ["Lfm2ForCausalLM", "Lfm2MoeForCausalLM"]:
# Attention backend selection moved to the override registry
# (arg_groups/overrides.py: _lfm2_overrides).
assert resolved_view(server_args).attention_backend != "triton", (
f"{model_arch} does not support triton attention backend, "
"as the first layer might not be an attention layer"
)
# MiniMaxM2ForCausalLM (enable_tf32_matmul) moved to the override registry
# (arg_groups/overrides.py: _minimax_m2_overrides).
# Qwen3VL aiter unified-attention page_size moved to the override registry
# (arg_groups/overrides.py: _qwen3vl_overrides).
# Hybrid-mamba radix cache handling for the per-arch branch call sites
# dissolved above: the resolution pass self-guards on the arch union
# (and the Granite layer_types probe), so one call covers them all.
# Hybrid-spec archs already resolved at the pre-dispatch call above;
# for them this re-invocation is an idempotent no-op plus validation.
# Kept ahead of the sparse-head pass: the legacy per-branch calls
# resolved before that tail write of disable_overlap_schedule.
server_args._handle_mamba_radix_cache(model_arch=model_arch)
from sglang.srt.arg_groups.overrides import (
_sparse_head_overlap_disable,
run_post_process_pass,
)
run_post_process_pass(server_args, _sparse_head_overlap_disable)
# The FlashInfer AllReduce Fusion auto-enable and the enforce-disable
# terminal moved to the resolution pipeline (arg_groups/overrides.py:
# _flashinfer_allreduce_fusion_auto_enable /
# _enforce_disable_allreduce_fusion), invoked here at their legacy
# slots.
from sglang.srt.arg_groups.overrides import (
_enforce_disable_allreduce_fusion,
_flashinfer_allreduce_fusion_auto_enable,
)
run_post_process_pass(server_args, _flashinfer_allreduce_fusion_auto_enable)
run_post_process_pass(server_args, _enforce_disable_allreduce_fusion)
def handle_model_capability_adjustments(server_args: Any):
cfg = resolving_view(server_args)
if parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE:
return
from sglang.srt.arg_groups.overrides import (
_hrm_text_attention_force,
run_post_process_pass,
)
model_config = server_args.get_model_config()
hf_config = model_config.hf_config
# HRM-Text needs bidirectional prompt attention (prefill), which only
# the Triton backend honors at the kernel level. Radix/prefix reuse is
# also unsafe: the recurrent forward writes direction-dependent KV
# across many slots.
is_hrm_text = getattr(
hf_config, "model_type", None
) == "hrm_text" or "HrmTextForCausalLM" in getattr(hf_config, "architectures", [])
# prefix_lm defaults to True upstream; defaulting False would skip the
# bidirectional-attention forcing and silently produce junk output.
if is_hrm_text and getattr(hf_config, "prefix_lm", True):
run_post_process_pass(server_args, _hrm_text_attention_force)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
chunked_prefill_size=-1,
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
disable_radix_cache=True,
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
disable_cuda_graph=True,
)
# cuda_graph_config was already parsed from the legacy boolean, so
# flipping the boolean alone would not stop graph capture.
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
logger.warning(
"HRM-Text (prefix_lm) detected: forcing --attention-backend "
"triton, --chunked-prefill-size -1, --disable-radix-cache, and "
"--disable-cuda-graph for correctness of the bidirectional "
"prompt attention."
)
# EmbeddingGemma is a Gemma3TextModel with bidirectional prompt
# attention. Prefix reuse and split prefills would reuse K/V states
# whose values depend on later prompt tokens, so both are invalid.
# Breakable CUDA Graph captures one complete prefill and is the graph
# mode validated for this encoder-style attention.
# Native encoder architectures declare a pooling-only task and do not
# need the legacy --is-embedding intent flag. Decoder checkpoints still
# require that explicit opt-in because their architecture alone does
# not distinguish embedding from generation serving.
#
# ``_handle_model_capability_adjustments`` is also exercised directly
# by a few focused tests that use a small ModelConfig stand-in. Keep
# the old predicate as a compatibility fallback while production
# ModelConfig instances use the central capability contract.
embedding_model_spec = getattr(model_config, "embedding_model_spec", None)
if (
embedding_model_spec is not None
and embedding_model_spec.auto_enable_embedding
and not cfg.is_embedding
):
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
is_embedding=True,
)
logger.info(
"Embedding architecture detected: enabling embedding mode automatically."
)
is_embedding_gemma = (
embedding_model_spec is not None
and embedding_model_spec.bcg_prefill_policy == BCGPrefillPolicy.FULL_ENCODER
)
if embedding_model_spec is None:
is_embedding_gemma = getattr(model_config, "is_embedding_gemma", False)
if is_embedding_gemma:
# This is an encoder-only model even though its HF architecture is
# named Gemma3TextModel. Marking it as embedding mode enables the
# FlashAttention raw-K/V fast path, which does not write or read
# the paged KV cache during its single prefill forward.
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
is_embedding=True,
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
disable_radix_cache=True,
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
chunked_prefill_size=-1,
)
# Submit a list-valued embeddings request atomically so BCG can
# replay its full prefill batch instead of starting item zero
# while the remaining texts are still being tokenized.
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
enable_tokenizer_batch_encode=True,
)
requested_prefill_backend = (
cfg.prefill_attention_backend or cfg.attention_backend
)
if (
is_cuda()
and (is_sm90_supported() or is_sm100_supported())
and requested_prefill_backend in (None, "fa3", "fa4")
):
# Hopper/Blackwell's default FA backend can consume raw K/V
# tensors for a single embedding prefill. Enable its no-KV
# pool path before memory-pool sizing; an explicit non-FA
# backend retains the existing paged-KV behavior.
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
prefill_only_disable_kv_cache=True,
)
server_args._validate_prefill_only_disable_kv_cache_args()
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
if is_cuda() and cfg.cuda_graph_config.prefill.backend != Backend.DISABLED:
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.BREAKABLE
),
)
# CUDA-graph sizing has already run by this point and derives
# its generic maximum from the 8K chunked-prefill default.
# On the Hopper/Blackwell FA raw-K/V path, raise the unlocked
# default to a full eight-way 2K embedding batch; callers can
# still override this for larger aggregate prefills.
prefill_config = cfg.cuda_graph_config.prefill
# Unit-level capability tests may invoke this hook without
# running the full CUDA-graph configuration parser, which is
# where this internal lock set is normally initialized.
# Treat that minimal construction as having no user-locked
# graph settings.
cuda_graph_config_locked = getattr(
server_args, "_cuda_graph_config_locked", set()
)
if (Phase.PREFILL, "max_bs") not in cuda_graph_config_locked:
sizing = {
"max_bs": max(
prefill_config.max_bs or 0,
model_config.context_len,
16384,
)
}
if (Phase.PREFILL, "bs") not in cuda_graph_config_locked:
sizing["bs"] = server_args._generate_prefill_cuda_graph_batch_sizes(
sizing["max_bs"]
)
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, **sizing
),
)
elif not is_cuda():
# BCG is CUDA-only. Other graph backends do not support this
# encoder-style prefill, so retain the eager Triton path.
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
logger.info(
"EmbeddingGemma detected: disabling radix cache and chunked "
"prefill; using breakable CUDA graph for CUDA prefill."
)
if (
model_config.is_multimodal
and not model_config.is_multimodal_chunked_prefill_supported
):
declare_resolution(
server_args,
"_handle_model_capability_adjustments",
chunked_prefill_size=-1,
)
logger.info(
f"Automatically turn off --chunked-prefill-size as it is not supported for "
f"{hf_config.model_type}"
)
def handle_mamba_radix_cache(server_args: Any, model_arch: str):
# Resolution moved to the resolution pipeline (arg_groups/overrides.py:
# _mamba_radix_cache_resolution), invoked here at each legacy call
# slot; this handler keeps the validation.
from sglang.srt.arg_groups.overrides import (
_mamba_radix_cache_resolution,
mamba_extra_buffer_of,
run_post_process_pass,
)
run_post_process_pass(server_args, _mamba_radix_cache_resolution)
view = resolved_view(server_args)
if not view.uses_mamba_radix_cache:
return
if mamba_extra_buffer_of(view):
server_args._validate_mamba_extra_buffer(view, model_arch)
else:
server_args._validate_mamba_no_buffer(view, model_arch)
def handle_language_model_only(server_args: Any):
cfg = resolving_view(server_args)
if not cfg.language_model_only:
return
for flag, name in (
(cfg.encoder_only, "--encoder-only"),
(cfg.language_only, "--language-only"),
(cfg.enable_prefix_mm_cache, "--enable-prefix-mm-cache"),
(
cfg.enable_broadcast_mm_inputs_process,
"--enable-broadcast-mm-inputs-process",
),
(cfg.mm_enable_dp_encoder, "--mm-enable-dp-encoder"),
):
if flag:
raise ValueError(f"--language-model-only cannot be combined with {name}")
if cfg.disaggregation_mode != "null":
raise ValueError(
"--language-model-only is incompatible with --disaggregation-mode "
"prefill/decode"
)
architectures = server_args.get_model_config().hf_config.architectures
if not any(
a in server_args.LANGUAGE_MODEL_ONLY_ARCHITECTURES for a in architectures
):
raise ValueError(
f"--language-model-only does not support {architectures}. "
f"Supported: {list(server_args.LANGUAGE_MODEL_ONLY_ARCHITECTURES)}."
)
@@ -0,0 +1,306 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the model source paths."""
from __future__ import annotations
import importlib
import logging
import os
from typing import Any, Optional
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolving_view,
)
from sglang.srt.utils.common import is_remote_url
from sglang.srt.utils.hf_transformers_utils import check_gguf_file
from sglang.srt.utils.runai_utils import ObjectStorageModel, is_runai_obj_uri
logger = logging.getLogger(__name__)
def handle_model_source_paths(server_args: Any):
"""Prepare metadata for model paths backed by remote object stores."""
cfg = resolving_view(server_args)
server_args._resolve_hf_gguf_model_path()
seen_paths = set()
for model_path in (
cfg.model_path,
cfg.tokenizer_path,
cfg.speculative_draft_model_path,
):
if (
model_path is not None
and model_path not in seen_paths
and is_runai_obj_uri(model_path)
):
ObjectStorageModel.download_and_get_path(model_path)
seen_paths.add(model_path)
def resolve_hf_gguf_model_path(server_args: Any):
"""Turn a Hub reference to a .gguf into a local file path."""
cfg = resolving_view(server_args)
from sglang.srt.utils.hf_transformers_utils import resolve_hf_gguf_reference
resolved = resolve_hf_gguf_reference(cfg.model_path, revision=cfg.revision)
if resolved is not None:
logger.info("Resolved GGUF %s -> %s", cfg.model_path, resolved)
if cfg.tokenizer_path == cfg.model_path:
declare_resolution(
server_args,
"_resolve_hf_gguf_model_path",
tokenizer_path=resolved,
)
declare_resolution(
server_args,
"_resolve_hf_gguf_model_path",
model_path=resolved,
)
# A speculative draft can be a .gguf too, and it is loaded by path, so it
# needs the same Hub-reference resolution as the target.
if cfg.speculative_draft_model_path:
resolved_draft = resolve_hf_gguf_reference(
cfg.speculative_draft_model_path,
revision=cfg.speculative_draft_model_revision,
)
if resolved_draft is not None:
logger.info(
"Resolved draft GGUF %s -> %s",
cfg.speculative_draft_model_path,
resolved_draft,
)
declare_resolution(
server_args,
"_resolve_hf_gguf_model_path",
speculative_draft_model_path=resolved_draft,
)
def handle_modelscope_paths(server_args: Any):
"""Resolve model / tokenizer / speculative-draft paths from the local
ModelScope cache when possible, falling back to snapshot_download
for any path that is not already present on disk.
Note: speculative_token_map is intentionally NOT handled here
because its value uses repo_id/filename semantics rather than a
plain repo ID. That resolution lives in
:func:`sglang.srt.speculative.spec_utils.load_token_map`.
"""
cfg = resolving_view(server_args)
ms_root = None
ms_snapshot_download = None
def _resolve_or_download(
path: Optional[str],
ignore_patterns: Optional[list] = None,
revision: Optional[str] = None,
) -> Optional[str]:
nonlocal ms_root, ms_snapshot_download
if path is None:
return None
if not path or os.path.exists(path):
return path
if ms_snapshot_download is None:
from modelscope.hub.snapshot_download import (
snapshot_download as _ms_snapshot_download,
)
from modelscope.utils.file_utils import get_model_cache_root
ms_snapshot_download = _ms_snapshot_download
ms_root = get_model_cache_root()
# Check ModelScope default cache
cached = os.path.join(ms_root, path)
if os.path.exists(cached):
return cached
# Check user-specified download dir
if cfg.download_dir:
alt = os.path.join(cfg.download_dir, path)
if os.path.exists(alt):
return alt
# Cache miss — download from ModelScope hub
return ms_snapshot_download(
path,
cache_dir=cfg.download_dir,
revision=revision,
**({"ignore_patterns": ignore_patterns} if ignore_patterns else {}),
)
declare_resolution(
server_args,
"_handle_modelscope_paths",
model_path=_resolve_or_download(cfg.model_path, revision=cfg.revision),
)
declare_resolution(
server_args,
"_handle_modelscope_paths",
tokenizer_path=_resolve_or_download(
cfg.tokenizer_path,
ignore_patterns=["*.bin", "*.safetensors"],
revision=cfg.revision,
),
)
if cfg.speculative_draft_model_path:
declare_resolution(
server_args,
"_handle_modelscope_paths",
speculative_draft_model_path=_resolve_or_download(
cfg.speculative_draft_model_path,
revision=cfg.speculative_draft_model_revision or "main",
),
)
def handle_load_format(server_args: Any):
# The quantization side of the gguf coupling moved to the pipeline
# (arg_groups/overrides.py: _gguf_quantization); load_format itself is
# genuine config (runtime user updates write it) and stays imperative.
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_gguf_quantization,
run_post_process_pass,
)
run_post_process_pass(server_args, _gguf_quantization)
if (cfg.load_format == "auto" or cfg.load_format == "gguf") and check_gguf_file(
cfg.model_path
):
declare_resolution(
server_args,
"_handle_load_format",
load_format="gguf",
)
if cfg.load_format == "auto" and server_args._is_mistral_native_format():
declare_resolution(
server_args,
"_handle_load_format",
load_format="mistral",
)
logger.info(
"Detected Mistral native format checkpoint, setting load_format='mistral'"
)
if is_runai_obj_uri(cfg.model_path):
declare_resolution(
server_args,
"_handle_load_format",
load_format="runai_streamer",
)
elif is_remote_url(cfg.model_path):
declare_resolution(
server_args,
"_handle_load_format",
load_format="remote",
)
if (
cfg.speculative_draft_model_path is not None
and is_runai_obj_uri(cfg.speculative_draft_model_path)
and cfg.speculative_draft_load_format is None
):
declare_resolution(
server_args,
"_handle_load_format",
speculative_draft_load_format="runai_streamer",
)
if cfg.custom_weight_loader is None:
declare_resolution(server_args, "_handle_load_format", custom_weight_loader=[])
if cfg.load_format == "remote_instance":
if cfg.remote_instance_weight_loader_backend != "modelexpress" and (
cfg.remote_instance_weight_loader_seed_instance_ip is None
or cfg.remote_instance_weight_loader_seed_instance_service_port is None
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader settings."
)
declare_resolution(
server_args,
"_handle_load_format",
load_format="auto",
)
elif (
cfg.remote_instance_weight_loader_send_weights_group_ports is None
and cfg.remote_instance_weight_loader_backend == "nccl"
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader NCCL group ports settings."
)
declare_resolution(
server_args,
"_handle_load_format",
load_format="auto",
)
elif (
cfg.remote_instance_weight_loader_backend == "transfer_engine"
and not server_args.validate_transfer_engine()
):
logger.warning(
"Fallback load_format to 'auto' due to 'transfer_engine' backend is not supported."
)
declare_resolution(
server_args,
"_handle_load_format",
load_format="auto",
)
# Check whether TransferEngine can be used when users want to start seed service that supports TransferEngine backend.
if cfg.remote_instance_weight_loader_start_seed_via_transfer_engine:
declare_resolution(
server_args,
"_handle_load_format",
remote_instance_weight_loader_start_seed_via_transfer_engine=server_args.validate_transfer_engine(),
)
# "ipc_cache" is an internal-only load format: ModelRunner sets it
# automatically when the weight cache is enabled, and it is not a public
# --load-format choice. Setting it directly is always wrong (no daemon is
# launched, and fallback_load_format inherits a nonsensical format), so
# reject it and point at the knob (defense-in-depth; the CLI already
# rejects it via LOAD_FORMAT_CHOICES).
if cfg.load_format == "ipc_cache":
raise ValueError(
"load_format='ipc_cache' is an internal-only format and must not "
"be set directly. Enable the weight cache via --weight-cache-mode "
"client (connect to an existing daemon) or daemon (launch one); "
"that selects IPC loading automatically."
)
# Speculative decoding loads an extra draft model whose weights the
# daemon does not export, so refuse the combination up front instead of
# failing deep inside draft-worker load (draft-model daemon TBD).
if cfg.weight_cache_mode != "off" and cfg.speculative_algorithm is not None:
raise ValueError(
"--weight-cache-mode is not supported together with speculative "
"decoding (--speculative-algorithm): the weight cache daemon does "
"not export the draft model's weights. Disable one of them "
"(--weight-cache-mode off) for this configuration."
)
def validate_transfer_engine(server_args: Any):
cfg = resolving_view(server_args)
try:
mooncake_available = importlib.util.find_spec("mooncake.engine") is not None
except (ModuleNotFoundError, ValueError):
mooncake_available = False
if not mooncake_available:
logger.warning(
"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
)
return False
elif cfg.enable_memory_saver:
logger.warning(
"Memory saver is enabled, which is not compatible with TransferEngine. Does not support using TransferEngine as remote instance weight loader backend."
)
return False
else:
return True
+477
View File
@@ -0,0 +1,477 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the MoE kernel configuration."""
from __future__ import annotations
import logging
import os
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
from sglang.srt.utils.common import is_npu, parse_connector_type
logger = logging.getLogger(__name__)
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.
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_moe_runner_backend_quant_constraints,
_moe_runner_fusion_disable,
run_post_process_pass,
)
run_post_process_pass(server_args, _moe_runner_backend_quant_constraints)
view = resolved_view(server_args)
if view.moe_runner_backend == "flashinfer_cutlass":
assert view.quantization in [
"modelopt_fp4",
"modelopt_fp8",
"modelopt_mixed",
None,
], f"Invalid quantization '{view.quantization}'. \nFlashInfer Cutlass MOE supports only: 'modelopt_fp4', 'modelopt_fp8', 'modelopt_mixed', or bfloat16 (None)."
assert view.ep_size in [
1,
cfg.tp_size,
], "The expert parallel size must be 1 or the same as the tensor parallel size"
if view.moe_runner_backend == "flashinfer_cutedsl":
# 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
), 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,
cfg.tp_size,
], "The expert parallel size must be 1 or the same as the tensor parallel size"
assert view.moe_a2a_backend in [
"none",
"deepep",
"flashinfer",
], (
f"flashinfer_cutedsl supports moe_a2a_backend='none', 'deepep', or 'flashinfer', "
f"got '{view.moe_a2a_backend}'."
)
if view.moe_a2a_backend == "deepep" and (
view.quantization == "nvfp4_online"
or envs.SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION.get()
):
raise ValueError(
"flashinfer_cutedsl per-token NVFP4 activation requires "
"moe_a2a_backend='none' or 'flashinfer'."
)
if view.moe_runner_backend in ["flashinfer_trtllm", "experimental_sgl_trtllm"]:
assert view.quantization in [
"modelopt_fp4",
"nvfp4_online",
"fp8",
"mxfp8",
"modelopt_fp8",
"modelopt_mixed",
"compressed-tensors",
None,
], f"Invalid quantization '{view.quantization}'. \nFlashInfer TRTLLM MOE supports only: 'modelopt_fp4', 'nvfp4_online', 'fp8', 'modelopt_fp8', 'modelopt_mixed', 'compressed-tensors', or bfloat16 (None)."
if view.moe_runner_backend == "flashinfer_trtllm_routed":
assert view.quantization in [
"fp8",
"mxfp8",
"modelopt_fp4",
"modelopt_mixed",
"nvfp4_online",
None,
], f"Invalid quantization '{view.quantization}'. \nFlashInfer TRTLLM routed MOE supports only: 'fp8', 'mxfp8', 'modelopt_fp4', 'modelopt_mixed', 'nvfp4_online', or bfloat16 (None)."
# The runner-driven shared-experts fusion disables moved to the
# pipeline (arg_groups/overrides.py: _moe_runner_fusion_disable),
# invoked here at the legacy write slots.
run_post_process_pass(server_args, _moe_runner_fusion_disable)
if resolved_view(server_args).moe_runner_backend == "cutlass" and resolved_view(
server_args
).quantization in [
"fp8",
"mxfp8",
]:
assert (
resolved_view(server_args).ep_size == 1
), "FP8/MXFP8 Cutlass MoE is only supported with ep_size == 1"
def handle_a2a_moe(server_args: Any):
# The backend overrides and the ep_size=tp_size adjustments moved to
# 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.
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_a2a_backend_overrides,
_a2a_ep_size,
_a2a_fusion_adjustments,
run_post_process_pass,
)
run_post_process_pass(server_args, _a2a_backend_overrides)
run_post_process_pass(server_args, _a2a_ep_size)
# The a2a-driven shared-experts fusion adjustments moved to the
# pipeline (arg_groups/overrides.py: _a2a_fusion_adjustments),
# invoked here at the legacy write slots.
run_post_process_pass(server_args, _a2a_fusion_adjustments)
a2a_backend = resolved_view(server_args).moe_a2a_backend
if cfg.enable_waterfill:
declare_resolution(
server_args, "_handle_a2a_moe", enforce_shared_experts_fusion=True
)
logger.info(f"Waterfill is enabled with moe_a2a_backend='{a2a_backend}'.")
if a2a_backend == "deepep":
if cfg.moe_runner_backend == "flashinfer_cutedsl":
if cfg.deepep_mode == "auto":
declare_resolution(
server_args,
"_handle_a2a_moe",
deepep_mode="low_latency",
)
logger.warning(
"Forcing --deepep-mode low_latency: flashinfer_cutedsl "
"FP4 MoE has no DeepEP normal-dispatch handler, so "
"deepep auto mode would crash during prefill. "
"low_latency covers both prefill and decode."
)
elif cfg.deepep_mode == "normal":
raise ValueError(
"flashinfer_cutedsl FP4 MoE only supports DeepEP "
"low_latency dispatch (masked layout). DeepEP normal "
"(prefill) dispatch has no CuteDSL FP4 handler. Pass "
"--deepep-mode low_latency or auto."
)
if cfg.deepep_mode == "normal":
logger.warning("Cuda graph is disabled because deepep_mode=`normal`")
declare_resolution(
server_args,
"_handle_a2a_moe",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_a2a_moe",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
if a2a_backend == "deepep_v2":
server_args._validate_deepep_v2_model_architecture()
if resolved_view(server_args).enable_deterministic_inference:
raise ValueError(
"DeepEP v2 does not forward deterministic=True to "
"ElasticBuffer, so deterministic sorting remains disabled. "
"Disable --enable-deterministic-inference or use "
"--moe-a2a-backend deepep."
)
# ElasticBuffer requires CUMEM, but not NVLS or its preallocation.
os.environ.setdefault("NCCL_CUMEM_ENABLE", "1")
# Respect model-level runner declarations before resolving auto.
resolved_runner = resolved_view(server_args).moe_runner_backend
if resolved_runner == "auto":
declare_resolution(
server_args, "_handle_a2a_moe", moe_runner_backend="deep_gemm"
)
logger.warning(
"DeepEP v2 MoE: resolved --moe-runner-backend auto -> deep_gemm."
)
elif resolved_runner != "deep_gemm":
raise ValueError(
"DeepEP v2 MoE currently supports only "
f"--moe-runner-backend deep_gemm. Got {resolved_runner!r}. "
"Add a runner adapter before enabling DeepEP v2 with other "
"MoE runners."
)
if cfg.enable_two_batch_overlap or cfg.enable_single_batch_overlap:
raise ValueError(
"DeepEP v2 MoE has not implemented the TBO/SBO overlap hooks yet. "
"Disable --enable-two-batch-overlap and "
"--enable-single-batch-overlap when using --moe-a2a-backend deepep_v2."
)
if cfg.enforce_shared_experts_fusion:
raise ValueError(
"DeepEP v2 MoE has not validated fused shared experts yet. "
"Remove --enforce-shared-experts-fusion when using "
"--moe-a2a-backend deepep_v2."
)
# Prefill reads host counts and is not graph-capturable.
declare_resolution(
server_args,
"_handle_a2a_moe",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
logger.warning(
f"DeepEP v2 MoE is enabled. The expert parallel size is adjusted to be the same as the tensor parallel size[{cfg.tp_size}]."
)
logger.warning(
"DeepEP v2 MoE is using deepep_v2_mode=%s. This controls "
"ElasticBuffer direct/hybrid mode and is independent from "
"--deepep-mode normal/low_latency. DeepEP v2 MoE enables the "
"decode CUDA graph on the masked decode path (any comm mode) "
"and disables shared expert fusion. "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK is a "
"per-rank communication buffer capacity, not a model limit; "
"increase it for large prefill/chunked-prefill workloads.",
cfg.deepep_v2_mode,
)
# The resolving view, not the field: `_a2a_backend_overrides` may have
# moved this already (waterfill forces `deepep`).
a2a_now = resolved_view(server_args).moe_a2a_backend
if (a2a_now == "none" and is_npu()) or a2a_now == "ascend_tp":
# FIXME (OrangeRedeng): for some reasons if pass "ascend_tp" accuracy drops to zero
declare_resolution(
server_args,
"_handle_a2a_moe",
moe_a2a_backend="none",
)
if cfg.moe_a2a_backend == "flashinfer":
assert (
resolved_view(server_args).enable_dp_attention
and cfg.dp_size == cfg.tp_size
), "Flashinfer MoE A2A is only supported with dp_size == tp_size and --enable-dp-attention"
if cfg.deepep_mode != "auto":
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
):
envs.SGLANG_MOE_NVFP4_DISPATCH.set(True)
logger.warning(
"SGLANG_MOE_NVFP4_DISPATCH is set to True for Flashinfer MoE A2A"
)
assert resolved_view(server_args).moe_runner_backend in [
"flashinfer_cutlass",
"flashinfer_cutedsl",
"flashinfer_trtllm_routed",
], "Flashinfer MoE A2A is only supported with flashinfer_cutlass, flashinfer_cutedsl or flashinfer_trtllm_routed moe runner backend"
if a2a_backend == "mori":
if cfg.deepep_mode == "auto":
declare_resolution(
server_args,
"_handle_a2a_moe",
deepep_mode="normal",
)
logger.warning("auto set deepep_mode=`normal` for MORI EP")
# Check chunked prefill for mori
# Skip validation if chunked prefill is disabled (i.e., size <= 0).
# 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()
) <= 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 "
"(chunked_prefill_size by default)"
)
if a2a_backend == "pplx":
if cfg.deepep_mode == "normal":
raise ValueError(
"moe_a2a_backend='pplx' only supports low-latency mode; "
"set --deepep-mode to 'low_latency' or 'auto'."
)
if cfg.deepep_mode == "auto":
declare_resolution(
server_args,
"_handle_a2a_moe",
deepep_mode="low_latency",
)
logger.warning("auto set deepep_mode=`low_latency` for PPLX EP")
# pplx-kernels' AllToAll needs numDPGroups (== attention dp_size) > 1;
# without DP attention numDPGroups == 1 and construction fails deep in
# the kernel. This also implies ep_size >= 2.
assert resolved_view(server_args).enable_dp_attention and cfg.dp_size >= 2, (
"moe_a2a_backend='pplx' requires --enable-dp-attention with at "
"least 2 DP groups (--dp-size >= 2)."
)
# pplx runs the masked DeepGEMM expert path (sm_90a): reject other
# runners and resolve auto -> deep_gemm. Unquantized bf16 pplx needs
# an explicit deep_gemm backend, otherwise the expert layer falls
# through to the deprecated masked path and asserts at runtime.
assert resolved_view(server_args).moe_runner_backend in ("deep_gemm", "auto"), (
"moe_a2a_backend='pplx' is only supported with --moe-runner-backend "
"deep_gemm (or auto)."
)
if cfg.moe_runner_backend == "auto":
declare_resolution(
server_args,
"_handle_a2a_moe",
moe_runner_backend="deep_gemm",
)
logger.warning("auto set moe_runner_backend=`deep_gemm` for PPLX EP")
# Check per-rank dispatch tokens for pplx
# Skip validation if chunked prefill is disabled (i.e., size <= 0)
# 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()
) <= 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 "
"(chunked_prefill_size, or the decode cuda-graph batch size)"
)
def validate_deepep_v2_speculative_draft(server_args: Any) -> None:
"""Reject an explicit or inherited DeepEP v2 draft backend."""
view = resolved_view(server_args)
draft_backend = view.speculative_moe_a2a_backend
if draft_backend is None and view.speculative_algorithm:
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
algorithm = SpeculativeAlgorithm.from_string(view.speculative_algorithm)
if not algorithm.is_ngram():
draft_backend = view.moe_a2a_backend
if draft_backend == "deepep_v2":
raise ValueError(
"DeepEP v2 MoE is not validated as a speculative draft backend. "
"Select another --speculative-moe-a2a-backend."
)
def validate_deepep_v2_dispatch_token_budget(server_args: Any) -> None:
"""Check the configured prefill and decode-graph buffer bounds."""
view = resolved_view(server_args)
if view.moe_a2a_backend != "deepep_v2":
return
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 (
view.max_prefill_tokens or 0
)
if prefill_tokens > capacity:
raise ValueError(
"DeepEP v2 per-rank prefill budget exceeds "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK: "
f"required={prefill_tokens}, capacity={capacity}. Raise the "
"environment value or lower --chunked-prefill-size/"
"--max-prefill-tokens."
)
if view.disaggregation_mode == "prefill":
return
decode_config = getattr(view.cuda_graph_config, "decode", None)
if decode_config is None or decode_config.backend == Backend.DISABLED:
return
graph_bs = decode_config.max_bs or 0
if view.max_running_requests is not None:
attn_dp_size = view.dp_size if view.enable_dp_attention else 1
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
if view.speculative_algorithm
else 1
)
graph_tokens = graph_bs * tokens_per_req
if graph_tokens > capacity:
raise ValueError(
"DeepEP v2 per-rank decode CUDA graph exceeds "
"SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK: "
f"required={graph_tokens}, capacity={capacity} "
f"(requests={graph_bs}, tokens/request={tokens_per_req}). Raise "
"the environment value or lower --cuda-graph-max-bs."
)
def validate_deepep_v2_model_architecture(server_args: Any) -> None:
"""Allow DeepEP v2 only where its model workflow is validated."""
if (
parse_connector_type(resolved_view(server_args).model_path)
== ConnectorType.INSTANCE
):
raise ValueError(
"DeepEP v2 MoE cannot validate a model loaded through an instance "
"connector. Load it from a model path or use "
"--moe-a2a-backend deepep."
)
architectures = (
getattr(server_args.get_model_config().hf_config, "architectures", None) or []
)
architecture = architectures[0] if architectures else None
# These architectures take the A2A MoE path and skip post-expert
# all-reduce.
validated_architectures = (
"DeepseekV3ForCausalLM",
"DeepseekV4ForCausalLM",
"Qwen3MoeForCausalLM",
)
if architecture not in validated_architectures:
raise ValueError(
f"DeepEP v2 MoE is not validated for {architecture!r}; supported "
f"architectures are {sorted(validated_architectures)}. "
"Other model workflows may require an all-reduce after A2A "
"combine. Use --moe-a2a-backend deepep."
)
def validate_cutedsl_a2a_token_budget(server_args: Any):
"""Fail fast if the FlashInfer A2A dispatcher workspace cannot cover the
largest CuteDSL MoE forward. Runs after speculative decoding is resolved
so cutedsl_moe_max_num_tokens() sees the final num_tokens_per_req."""
cfg = resolving_view(server_args)
view = resolved_view(server_args)
if not (
view.moe_a2a_backend == "flashinfer"
and view.moe_runner_backend == "flashinfer_cutedsl"
and cfg.max_prefill_tokens > 0
and cfg.disaggregation_mode != "decode"
):
return
required_tokens = server_args.cutedsl_moe_max_num_tokens()
max_dispatch_tokens_per_rank = (
envs.SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() or 1024
)
max_cutedsl_tokens = max_dispatch_tokens_per_rank * view.ep_size
if max_cutedsl_tokens < required_tokens:
required_per_rank = (required_tokens + view.ep_size - 1) // view.ep_size
raise ValueError(
"FlashInfer MoE A2A with flashinfer_cutedsl requires "
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK * "
"ep_size to cover the largest CuteDSL MoE forward "
f"({required_tokens} tokens). Otherwise the FlashInfer "
"dispatcher can crash at runtime with "
"`ValueError: num_tokens (...) exceeds max_num_tokens (...)`. "
"Current values: "
f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK="
f"{max_dispatch_tokens_per_rank}, ep_size={view.ep_size}, "
f"capacity={max_cutedsl_tokens}, required={required_tokens}. "
f"Set `export "
f"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK="
f"{required_per_rank}` or lower the relevant limit "
f"(e.g. --max-prefill-tokens) to <= {max_cutedsl_tokens}."
)
+26 -14
View File
@@ -227,23 +227,18 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
A slot that runs after resolution -- ``check_server_args`` hosts one -- lands A slot that runs after resolution -- ``check_server_args`` hosts one -- lands
in the same stash, which publish projects from later, so it needs no field in the same stash, which publish projects from later, so it needs no field
write either. After *publish* there is no such later projection: the stash write either. After *publish* there is no such later projection: the stash
would grow an entry nothing reads. So, like ``declare_late_resolution``, would grow an entry nothing reads.
this refuses the published record -- post-publish changes go to the bags
through ``get_context().override(...)``. So what is refused is the *declaration*, not the record. A pass that returns
an empty dict is a validation, and it may run on the published instance --
it has to, because ``Engine(server_args=sa)`` after ``Engine.shutdown()``
re-runs ``check_server_args`` on the very instance the context still holds.
A pass that returns a non-empty dict there is refused, as
``declare_late_resolution`` is -- post-publish changes go to the bags through
``get_context().override(...)``.
""" """
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context
try:
published = get_context().server_args
except ValueError:
published = None
if published is server_args:
raise ValueError(
f"run_post_process_pass({fn.__qualname__!r}) called on the published "
"config; the stash is projected at publish and never again, so a "
"declaration made here would be a silent no-op -- post-publish "
"changes go to the bags via get_context().override(...)"
)
declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args))) declared = fn(ResolvedView(server_args, overlay=_declaration_overlay(server_args)))
if not isinstance(declared, dict): if not isinstance(declared, dict):
raise TypeError( raise TypeError(
@@ -251,6 +246,23 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
f"got {type(declared).__name__}" f"got {type(declared).__name__}"
) )
if declared: if declared:
# Refused only once there is something to record. A pass that declares
# nothing is a validation, and `check_server_args` runs those again on
# a rebuild: `Engine(server_args=sa)` after `Engine.shutdown()` hands
# back the same instance while the context still holds it, and
# refusing on identity alone would fail that launch.
try:
published = get_context().server_args
except ValueError:
published = None
if published is server_args:
raise ValueError(
f"run_post_process_pass({fn.__qualname__!r}) declared "
f"{sorted(declared)} on the published config; the stash is "
"projected at publish and never again, so this would be a "
"silent no-op -- post-publish changes go to the bags via "
"get_context().override(...)"
)
entry = (fn.__qualname__, dict(declared)) entry = (fn.__qualname__, dict(declared))
stash = getattr(server_args, "_resolved_overrides", None) stash = getattr(server_args, "_resolved_overrides", None)
if stash is None: if stash is None:
@@ -0,0 +1,658 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for context- and decode-context parallelism."""
from __future__ import annotations
import logging
import os
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolved_view,
resolving_view,
)
from sglang.srt.connector import ConnectorType
from sglang.srt.environ import envs
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
from sglang.srt.utils.common import is_cuda, parse_connector_type
logger = logging.getLogger(__name__)
def handle_context_parallelism(server_args: Any):
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()
hf_config = model_config.hf_config
model_arch = hf_config.architectures[0]
if model_arch in CP_V2_DEFAULT_MODEL_CLASSES:
is_dsa_default_model = is_deepseek_dsa(hf_config)
# DSA CP-v2 currently supports only the interleave strategy.
enable_default_cp_v2 = not is_dsa_default_model or (
cfg.enable_prefill_cp and cfg.cp_strategy == "interleave"
)
if enable_default_cp_v2 and not envs.SGLANG_ENABLE_CP_V2.is_set():
envs.SGLANG_ENABLE_CP_V2.set(True)
if (
cfg.enable_prefill_cp
and model_arch in ("MiMoV2ForCausalLM", "MiMoV2FlashForCausalLM")
and envs.SGLANG_ENABLE_CP_V2.get()
):
if cfg.cp_strategy != "zigzag":
raise ValueError("MiMo V2 CP-v2 only supports --cp-strategy zigzag.")
if (
model_config.is_multimodal
and not cfg.language_only
and not cfg.language_model_only
):
raise ValueError(
"MiMo V2 CP-v2 only supports text inference; add "
"--language-only."
)
if cfg.enable_prefill_cp and cfg.cp_strategy is None:
raise ValueError(
"--cp-strategy must be set when --enable-prefill-cp is enabled."
)
if cfg.enable_prefill_context_parallel and cfg.enable_dsa_prefill_context_parallel:
raise ValueError(
"--enable-prefill-context-parallel and "
"--enable-nsa-prefill-context-parallel are mutually "
"exclusive. Use --enable-nsa-prefill-context-parallel for "
"DeepSeek V3.2 (NSA) models and "
"--enable-prefill-context-parallel for MLA-based models "
"(DeepSeek V3/R1, Kimi K2.5) or MHA/GQA-based models."
)
view = resolved_view(server_args)
if view.attn_cp_size > 1:
# The tp_size is the world size, not the real tensor parallel size
assert (
cfg.tp_size % view.attn_cp_size == 0
), "tp_size must be divisible by attn_cp_size"
assert (
cfg.tp_size % (cfg.dp_size * view.attn_cp_size) == 0
), "tp_size must be divisible by dp_size * attn_cp_size"
assert (
not cfg.enable_aiter_allreduce_fusion
), "Aiter allreduce fusion is not supported with context parallelism"
if cfg.moe_dp_size > 1:
# The tp_size is the world size, not the real tensor parallel size
assert (
cfg.tp_size % cfg.moe_dp_size == 0
), "tp_size must be divisible by moe_dp_size"
assert (
view.ep_size * cfg.moe_dp_size <= cfg.tp_size
), "ep_size * moe_dp_size must be less than or equal to tp_size"
assert cfg.pp_size == 1, "PP is not supported with context parallelism"
if view.ep_size > 1:
assert (
view.ep_size * cfg.moe_dp_size == cfg.tp_size
), "ep_size * moe_dp_size must be equal to tp_size"
assert (
not cfg.enable_aiter_allreduce_fusion
), "Aiter allreduce fusion is not supported with context parallelism"
if view.attn_cp_size != cfg.moe_dp_size:
assert (
cfg.moe_dp_size == 1
), "attn_cp_size != moe_dp_size is only supported when moe_dp_size == 1"
from sglang.srt.layers.cp.base import init_cp_strategy
init_cp_strategy(
enable_prefill_cp=bool(cfg.enable_prefill_cp),
cp_size=cfg.attn_cp_size,
cp_strategy=cfg.cp_strategy,
)
def handle_dcp_validation(server_args: Any):
cfg = resolving_view(server_args)
if cfg.dcp_size < 1:
raise ValueError(
"Decode context parallel size (--dcp-size / "
"--decode-context-parallel-size) must be >= 1, but got "
f"dcp_size={cfg.dcp_size}."
)
if cfg.dcp_comm_backend in ("a2a", "fi_a2a") and cfg.dcp_size <= 1:
raise ValueError(
f"--dcp-comm-backend {cfg.dcp_comm_backend} only affects the "
"decode context-parallel attention reduction and therefore "
"requires --dcp-size / --decode-context-parallel-size > 1, but "
f"got dcp_size={cfg.dcp_size}."
)
if cfg.dcp_comm_backend == "fi_a2a" and not is_cuda():
raise ValueError(
"--dcp-comm-backend fi_a2a delegates the exchange to FlashInfer's "
"MNNVL All-to-All kernel, which requires an NVIDIA CUDA platform "
"with SM90+ and MNNVL fabric memory (e.g. GB200 NVL72). The "
"authoritative fabric probe runs at model-runner init; use 'a2a' "
"or 'ag_rs' on clusters without MNNVL."
)
if cfg.dcp_replicate_q_proj:
if cfg.dcp_size <= 1:
raise ValueError("--dcp-replicate-q-proj requires --dcp-size > 1.")
if cfg.dcp_comm_backend not in ("a2a", "fi_a2a"):
raise ValueError(
"--dcp-replicate-q-proj only applies to the a2a/fi_a2a DCP "
"communication backend (it removes the head-dim Q all-gather); "
f"got --dcp-comm-backend={cfg.dcp_comm_backend}."
)
def handle_data_parallelism(server_args: Any):
# The dp_size==1 resets moved to the resolution pipeline
# (arg_groups/overrides.py: _data_parallelism_defaults).
cfg = resolving_view(server_args)
from sglang.srt.arg_groups.overrides import (
_data_parallelism_defaults,
run_post_process_pass,
)
run_post_process_pass(server_args, _data_parallelism_defaults)
if cfg.mm_enable_dp_encoder:
if cfg.tp_size == 1:
logger.warning(
"--mm-enable-dp-encoder is enabled with TP=1, so the encoder "
"has no data-parallel work to distribute. Disable it unless "
"you need to validate this configuration."
)
else:
logger.info(
"--mm-enable-dp-encoder is enabled across TP=%d. It replicates "
"the vision encoder and distributes image work across ranks; "
"this is most useful when high-resolution or multi-image ViT "
"prefill is a material part of TTFT. Measure against the default "
"for small-image workloads because replication and aggregation "
"can increase memory use and overhead.",
cfg.tp_size,
)
if resolved_view(server_args).enable_dp_attention:
declare_resolution(
server_args,
"_handle_data_parallelism",
schedule_conservativeness=cfg.schedule_conservativeness * 0.3,
)
assert cfg.tp_size % cfg.dp_size == 0
original_chunked_prefill_size = cfg.chunked_prefill_size
declare_resolution(
server_args,
"_handle_data_parallelism",
chunked_prefill_size=cfg.chunked_prefill_size // cfg.dp_size,
)
logger.warning(
f"DP attention is enabled. chunked prefill size is adjusted "
f"from {original_chunked_prefill_size} to {cfg.chunked_prefill_size}."
)
# The prefill CUDA graph max_bs was derived from the pre-DP-division
# chunked_prefill_size in _handle_gpu_memory_settings (which runs
# before this handler). Re-clamp it (and the captured shape list) to
# the per-DP-rank chunked_prefill_size so breakable CUDA graph
# capture never exceeds the MoE all-to-all's max_num_tokens budget,
# which is also sized from the DP-adjusted chunked_prefill_size.
prefill_cfg = cfg.cuda_graph_config.prefill
if (
prefill_cfg.backend != Backend.DISABLED
and prefill_cfg.max_bs is not None
and prefill_cfg.max_bs > cfg.chunked_prefill_size
and (Phase.PREFILL, "max_bs") not in server_args._cuda_graph_config_locked
):
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"]
)
declare_resolution(
server_args,
"_handle_data_parallelism",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, **clamped
),
)
# Resolve the phase-aware TP LM-head default before validating the
# resulting DP/TP LM-head configuration.
from sglang.srt.arg_groups.overrides import (
_dp_lm_head_validation,
_tp_lm_head_all_to_all_default,
)
run_post_process_pass(server_args, _tp_lm_head_all_to_all_default)
run_post_process_pass(server_args, _dp_lm_head_validation)
def handle_dwdp(server_args: Any):
cfg = resolving_view(server_args)
if cfg.dwdp_size <= 1:
return
assert (
cfg.dwdp_size >= 2
), f"dwdp_size must be >= 2 when enabled, got {cfg.dwdp_size}"
assert (
cfg.dwdp_size == cfg.tp_size
), f"dwdp_size ({cfg.dwdp_size}) must equal tp_size ({cfg.tp_size})"
assert cfg.disaggregation_mode in (
"null",
"prefill",
), "DWDP requires --disaggregation-mode null or prefill"
assert (
not cfg.enable_eplb
), "EPLB dynamic migration conflicts with static DWDP partitioning"
assert (
cfg.speculative_algorithm is None
), "DWDP does not support speculative decoding (MTP/draft workers)"
assert cfg.pp_size == 1, "DWDP requires pp_size == 1"
assert (
not cfg.enable_two_batch_overlap
), "DWDP's prefetch event protocol does not support two-batch overlap"
if cfg.disaggregation_mode == "null":
logger.warning(
"DWDP with --disaggregation-mode null: decode steps re-fetch all "
"remote expert weights every step, which is slow. DWDP is "
"recommended only with --disaggregation-mode prefill."
)
declare_resolution(
server_args,
"_handle_dwdp",
dp_size=cfg.dwdp_size,
)
declare_resolution(
server_args,
"_handle_dwdp",
enable_dp_attention=True,
)
declare_resolution(
server_args, "_handle_dwdp", enable_dp_attention_local_control_broadcast=True
)
declare_resolution(
server_args,
"_handle_dwdp",
enable_dp_lm_head=True,
)
declare_resolution(
server_args,
"_handle_dwdp",
moe_dense_tp_size=1,
)
declare_resolution(
server_args,
"_handle_dwdp",
ep_size=cfg.dwdp_size,
)
declare_resolution(
server_args,
"_handle_dwdp",
moe_dp_size=1,
)
declare_resolution(
server_args,
"_handle_dwdp",
moe_a2a_backend="none",
)
envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.set(True)
declare_resolution(
server_args,
"_handle_dwdp",
disable_cuda_graph=True,
)
logger.info(
f"DWDP enabled: dwdp_size={cfg.dwdp_size}, "
f"auto-forced dp_size={cfg.dp_size}, ep_size={cfg.dwdp_size}, "
f"moe_dense_tp_size=1, moe_a2a_backend=none, "
f"dp_attention_local_control_broadcast=True, "
f"enable_dp_lm_head=True, SCHEDULER_SKIP_ALL_GATHER=True, "
f"disable_cuda_graph=True"
)
def handle_elastic_ep(server_args: Any):
cfg = resolving_view(server_args)
if cfg.elastic_ep_rejoin:
if cfg.ep_join_mode is None:
logger.warning(
"--elastic-ep-rejoin is deprecated, use --elastic-ep-join-mode recover instead."
)
declare_resolution(
server_args,
"_handle_elastic_ep",
ep_join_mode="recover",
)
else:
assert cfg.ep_join_mode == "recover", (
"--elastic-ep-rejoin (deprecated) conflicts with "
f"--elastic-ep-join-mode {cfg.ep_join_mode}."
)
if cfg.elastic_ep_backend is not None:
if cfg.enable_eplb:
if cfg.eplb_algorithm == "auto":
declare_resolution(
server_args,
"_handle_elastic_ep",
eplb_algorithm="elasticity_aware",
)
assert cfg.eplb_algorithm in [
"elasticity_aware",
"elasticity_aware_hierarchical",
], "Elastic EP requires eplb_algorithm to be set to 'auto' or 'elasticity_aware(_hierarchical)'."
assert cfg.pp_size == 1, "PP size should be set to 1 under elastic EP"
if cfg.elastic_ep_backend == "mooncake":
declare_resolution(
server_args,
"_handle_elastic_ep",
mooncake_ib_device=server_args._validate_ib_devices(
cfg.mooncake_ib_device
),
)
if cfg.ep_join_mode is not None:
assert (
cfg.elastic_ep_backend is not None
), "--elastic-ep-join-mode requires --elastic-ep-backend to be set."
if cfg.ep_join_mode == "scale":
assert cfg.node_rank == 1, (
"Elastic EP scale-up requires one joining TP group at "
f"--node-rank 1 (got {cfg.node_rank})."
)
assert cfg.ep_join_rank_offset > 0, (
"Elastic EP scale joiners require "
"--elastic-ep-join-rank-offset set to the current "
"effective EP size."
)
if cfg.ep_join_rank_offset != 0:
assert cfg.ep_join_mode == "scale", (
"--elastic-ep-join-rank-offset is only valid with "
"--elastic-ep-join-mode scale."
)
assert cfg.ep_join_rank_offset >= 0, "elastic EP join rank offset must be >= 0."
if cfg.max_ep_size is not None:
assert (
cfg.elastic_ep_backend is not None
), "--max-ep-size requires --elastic-ep-backend to be set."
assert cfg.max_ep_size > 0, "--max-ep-size must be a positive integer."
scaling_active = (
cfg.elastic_ep_backend is not None
and cfg.max_ep_size is not None
and cfg.max_ep_size > cfg.tp_size
)
if cfg.elastic_ep_initial_size is not None:
assert scaling_active, (
"--elastic-ep-initial-size is only valid for an Elastic EP "
"deployment with --max-ep-size larger than its local TP size."
)
if scaling_active:
resolved = resolved_view(server_args)
assert (
cfg.elastic_ep_scale_timeout > 0
), "--elastic-ep-scale-timeout must be greater than zero."
assert cfg.tokenizer_worker_num == 1, (
"Elastic EP runtime scale-up currently requires "
"--tokenizer-worker-num 1."
)
assert (
not cfg.use_ray
), "Elastic EP runtime scale-up does not support --use-ray."
assert not cfg.enable_elastic_expert_backup, (
"Elastic EP runtime scale-up does not support "
"--enable-elastic-expert-backup."
)
declare_resolution(
server_args,
"_handle_elastic_ep",
enable_dp_attention_local_control_broadcast=True,
)
if cfg.ep_join_mode == "scale":
assert cfg.elastic_ep_initial_size is not None, (
"Elastic EP scale joiners require --elastic-ep-initial-size "
"set to the primary deployment's launch-time EP size."
)
assert cfg.elastic_ep_initial_size <= cfg.ep_join_rank_offset, (
"--elastic-ep-initial-size cannot exceed the current EP size "
f"(initial={cfg.elastic_ep_initial_size}, "
f"current={cfg.ep_join_rank_offset})."
)
join_target = cfg.ep_join_rank_offset + cfg.tp_size
assert join_target <= cfg.max_ep_size, (
"Elastic EP joining group exceeds --max-ep-size "
f"(join_target={join_target}, max_ep_size={cfg.max_ep_size})."
)
if cfg.tp_size == 1:
assert cfg.moe_dense_tp_size == 1, (
"A single-rank Elastic EP joining group requires "
"--moe-dense-tp-size 1."
)
else:
if cfg.elastic_ep_initial_size is None:
declare_resolution(
server_args,
"_handle_elastic_ep",
elastic_ep_initial_size=cfg.tp_size,
)
assert cfg.elastic_ep_initial_size == cfg.tp_size, (
"The primary --elastic-ep-initial-size must equal its "
f"launch-time TP size ({cfg.tp_size})."
)
assert cfg.elastic_ep_initial_size > 0
assert cfg.load_balance_method == "round_robin", (
"Elastic EP scale-up requires --load-balance-method round_robin; "
"load-aware methods "
"require global-rank load snapshots after scale "
f"(got {cfg.load_balance_method})."
)
assert cfg.elastic_ep_backend == "mooncake", (
"Elastic EP runtime scale-up requires --elastic-ep-backend "
f"mooncake (got elastic_ep_backend={cfg.elastic_ep_backend})."
)
assert cfg.pp_size == 1, (
"Elastic EP scale-up requires --pp-size 1 "
f"(got pp_size={cfg.pp_size}); WORLD must not span PP stages."
)
decode_cuda_graph_disabled = (
cfg.cuda_graph_config.decode.backend == Backend.DISABLED
)
prefill_cuda_graph_disabled = (
cfg.cuda_graph_config.prefill.backend == Backend.DISABLED
)
assert decode_cuda_graph_disabled and prefill_cuda_graph_disabled, (
"Elastic EP runtime scale-up requires decode and prefill CUDA "
"graphs to be disabled."
)
assert resolved.enable_dp_attention, (
"Elastic EP scale-up requires --enable-dp-attention; without it "
"the TP group is not equivalent to WORLD and the post-scale "
"collective path is invalid."
)
assert resolved.enable_dp_lm_head, (
"Elastic EP scale-up requires --enable-dp-lm-head so output "
"projection does not depend on the joining group's TP size."
)
assert resolved.attn_cp_size == 1, (
"Elastic EP scale-up requires --attn-cp-size 1 "
f"(got attn_cp_size={resolved.attn_cp_size})."
)
assert cfg.moe_dp_size == 1, (
"Elastic EP scale-up requires --moe-dp-size 1 "
f"(got moe_dp_size={cfg.moe_dp_size})."
)
assert resolved.ep_size == cfg.tp_size, (
"Elastic EP scale-up requires ep_size == tp_size "
f"(got ep_size={resolved.ep_size}, tp_size={cfg.tp_size}); EP, TP "
"and the attention DP group must all coincide with WORLD."
)
assert cfg.dp_size == cfg.tp_size, (
"Elastic EP scale-up requires dp_size == tp_size "
f"(got dp_size={cfg.dp_size}, tp_size={cfg.tp_size})."
)
assert resolved.moe_a2a_backend == "nixl", (
"Elastic EP scale-up requires --moe-a2a-backend nixl "
f"(got moe_a2a_backend={resolved.moe_a2a_backend})."
)
def handle_eplb_and_dispatch(server_args: Any):
cfg = resolving_view(server_args)
if cfg.enable_eplb and (cfg.expert_distribution_recorder_mode is None):
declare_resolution(
server_args,
"_handle_eplb_and_dispatch",
expert_distribution_recorder_mode="stat",
)
logger.warning(
"EPLB is enabled. The expert_distribution_recorder_mode is automatically set."
)
# Without an a2a backend all EP ranks run the MoE over the same tokens and
# sum their partial outputs, so the pick has to agree across ranks.
needs_rank_invariant_dispatch = resolved_view(server_args).moe_a2a_backend == "none"
if (cfg.enable_eplb or (cfg.init_expert_location != "trivial")) and (
cfg.ep_dispatch_algorithm is None
):
declare_resolution(
server_args,
"_handle_eplb_and_dispatch",
ep_dispatch_algorithm=(
"dynamic" if needs_rank_invariant_dispatch else "static"
),
)
# `dynamic` / `fake` switch to the row-index pick; `static` reads a
# per-rank table and `lp` samples inside its kernel.
if needs_rank_invariant_dispatch and cfg.ep_dispatch_algorithm in (
"static",
"lp",
):
raise ValueError(
f"--ep-dispatch-algorithm {cfg.ep_dispatch_algorithm} picks a "
"different physical replica per rank, which only holds up when an "
"a2a backend routes each token to a single rank. Use "
"--ep-dispatch-algorithm dynamic with --moe-a2a-backend none."
)
if cfg.enable_eplb and cfg.ep_join_mode != "scale":
assert resolved_view(server_args).ep_size > 1
def handle_legacy_cp_arguments(server_args: Any):
cfg = resolving_view(server_args)
legacy_mode_to_strategy = {
"in-seq-split": "zigzag",
"round-robin-split": "interleave",
}
strategy_to_legacy_mode = {
"zigzag": "in-seq-split",
"interleave": "round-robin-split",
}
if cfg.enable_prefill_context_parallel or cfg.enable_dsa_prefill_context_parallel:
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
enable_prefill_cp=True,
)
if cfg.enable_prefill_context_parallel and cfg.cp_strategy is None:
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
cp_strategy=legacy_mode_to_strategy[cfg.prefill_cp_mode],
)
if cfg.enable_dsa_prefill_context_parallel and cfg.cp_strategy is None:
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
cp_strategy=legacy_mode_to_strategy[cfg.dsa_prefill_cp_mode],
)
if cfg.enable_prefill_context_parallel and cfg.enable_dsa_prefill_context_parallel:
return
if not cfg.enable_prefill_cp or cfg.cp_strategy is None:
return
mode = strategy_to_legacy_mode[cfg.cp_strategy]
use_dsa_legacy_aliases = cfg.enable_dsa_prefill_context_parallel or getattr(
resolved_view(server_args), "attention_backend", None
) in ("dsa", "dsv4")
if use_dsa_legacy_aliases:
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
enable_dsa_prefill_context_parallel=True,
)
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
enable_prefill_context_parallel=False,
)
else:
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
enable_prefill_context_parallel=True,
)
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
dsa_prefill_cp_mode=mode,
)
declare_resolution(
server_args,
"_handle_legacy_cp_arguments",
prefill_cp_mode=mode,
)
def handle_expert_distribution_metrics(server_args: Any):
cfg = resolving_view(server_args)
if "SGLANG_ENABLE_EPLB_BALANCEDNESS_METRIC" in os.environ:
raise ValueError(
"SGLANG_ENABLE_EPLB_BALANCEDNESS_METRIC is no longer supported. Use "
"--expert-balancedness-report-mode with one of: off, server_log, "
"prometheus, both."
)
if server_args.should_report_expert_balancedness() and (
cfg.expert_distribution_recorder_mode is None
):
declare_resolution(
server_args,
"_handle_expert_distribution_metrics",
expert_distribution_recorder_mode="stat",
)
if cfg.expert_distribution_recorder_buffer_size is None:
if (x := cfg.eplb_rebalance_num_iterations) is not None:
declare_resolution(
server_args,
"_handle_expert_distribution_metrics",
expert_distribution_recorder_buffer_size=x,
)
elif cfg.expert_distribution_recorder_mode is not None:
declare_resolution(
server_args,
"_handle_expert_distribution_metrics",
expert_distribution_recorder_buffer_size=1000,
)
@@ -3,7 +3,7 @@ from __future__ import annotations
import dataclasses import dataclasses
import logging import logging
import os import os
from typing import TYPE_CHECKING from typing import TYPE_CHECKING, Any
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
declare_resolution, declare_resolution,
@@ -182,3 +182,81 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None:
"_alias_bootstrap_port_to_api_port", "_alias_bootstrap_port_to_api_port",
disaggregation_bootstrap_port=cfg.port, disaggregation_bootstrap_port=cfg.port,
) )
def handle_encoder_disaggregation(server_args: Any):
from sglang.srt.server_args import resolve_encoder_transfer_backend
cfg = resolving_view(server_args)
server_args._handle_language_model_only()
if cfg.enable_prefix_mm_cache and not cfg.encoder_only:
raise ValueError(
"--enable-prefix-mm-cache requires --encoder-only to be enabled"
)
if cfg.encoder_only and cfg.language_only:
raise ValueError("Cannot set --encoder-only and --language-only together")
if cfg.encoder_only and not cfg.disaggregation_mode == "null":
raise ValueError(
"Cannot set --encoder-only and --disaggregation-mode prefill/decode together"
)
if cfg.language_only and len(cfg.encoder_urls) == 0:
logger.info(
"--language-only is set without --encoder-urls. Encoders are "
"expected to register dynamically via the "
"EncoderBootstrapServer."
)
# Validate IB devices when mooncake backend is used
if (
cfg.disaggregation_transfer_backend == "mooncake"
and cfg.disaggregation_mode in ("prefill", "decode")
) or cfg.encoder_transfer_backend == "mooncake":
declare_resolution(
server_args,
"_handle_encoder_disaggregation",
disaggregation_ib_device=server_args._validate_ib_devices(
cfg.disaggregation_ib_device
),
)
# Validate model type for encoder disaggregation
hf_config = server_args.get_model_config().hf_config
model_arch = hf_config.architectures[0]
if cfg.encoder_transfer_backend == "auto":
declare_resolution(
server_args,
"_handle_encoder_disaggregation",
encoder_transfer_backend=resolve_encoder_transfer_backend(
cfg.encoder_transfer_backend, model_arch, cfg.tp_size
),
)
if cfg.encoder_only or cfg.language_only:
logger.info(
"Encoder transfer backend auto-resolved to %s for %s at TP%d.",
cfg.encoder_transfer_backend,
model_arch,
cfg.tp_size,
)
if (cfg.encoder_only or cfg.language_only) and model_arch not in [
"Qwen2VLForConditionalGeneration",
"Qwen3VLForConditionalGeneration",
"Qwen2_5_VLForConditionalGeneration",
"Qwen3VLMoeForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
"Qwen3OmniMoeForConditionalGeneration",
"Qwen2AudioForConditionalGeneration",
"Qwen2_5OmniForConditionalGeneration",
"Dots3NoteForCausalLM",
"KimiVLForConditionalGeneration",
"KimiK25ForConditionalGeneration",
"KimiK3ForConditionalGeneration",
"MiMoV2ForCausalLM",
]:
raise ValueError(
f"Model type {model_arch} is not supported for encoder disaggregation. "
f"Supported architectures: Qwen2VL, Qwen3VL, Qwen3.5, InternS2, "
f"Qwen2Audio, Qwen2.5Omni, Dots3-Note, Kimi, MiMoV2."
)
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for the per-platform backend defaults."""
from __future__ import annotations
import logging
from typing import Any
from sglang.srt.arg_groups.overrides import (
declare_resolution,
resolving_view,
)
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
from sglang.srt.utils.common import is_cuda, is_hip, is_host_cpu_arm64, is_npu
logger = logging.getLogger(__name__)
def handle_npu_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "npu":
from sglang.srt.hardware_backend.npu.utils import set_default_server_args
set_default_server_args(server_args)
current = cfg.cuda_graph_config.prefill.tc_compiler
if current is not None and current != "eager":
logger.warning(
"At this moment Ascend platform only support prefill graph compilation with "
"cuda_graph_config[prefill].tc_compiler='eager'."
)
declare_resolution(
server_args,
"_handle_npu_backends",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, tc_compiler="eager"
),
)
def handle_mps_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "mps":
if not use_mlx():
declare_resolution(
server_args,
"_handle_mps_backends",
disable_overlap_schedule=True,
)
def handle_amd_specifics(server_args: Any):
if is_hip():
declare_resolution(
server_args, "_handle_amd_specifics", triton_attention_num_kv_splits=16
)
def handle_nccl_pre_warm(server_args: Any):
# pre_warm_nccl is only used with CUDA or HIP hardware or NPU hardware
cfg = resolving_view(server_args)
if cfg.pre_warm_nccl and not (is_cuda() or is_hip() or is_npu()):
logger.warning(
"pre_warm_nccl is only applicable for CUDA or HIP hardware or NPU hardware. "
"Ignoring pre_warm_nccl setting on current hardware."
)
declare_resolution(server_args, "_handle_nccl_pre_warm", pre_warm_nccl=False)
def handle_xpu_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "xpu":
# Decode graph is opt-in on XPU: unless the user explicitly set
# --cuda-graph-backend-decode (or --cuda-graph-config), keep it
# disabled so the default startup doesn't require graph capture.
if (Phase.DECODE, "backend") not in server_args._cuda_graph_config_locked:
declare_resolution(
server_args,
"_handle_xpu_backends",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
elif cfg.cuda_graph_config.decode.backend not in (
Backend.DISABLED,
Backend.FULL,
):
logger.warning(
"XPU platform only supports decode backend 'full'; "
"disabling unsupported decode backend '%s'.",
cfg.cuda_graph_config.decode.backend,
)
declare_resolution(
server_args,
"_handle_xpu_backends",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
def handle_cpu_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "cpu":
if cfg.attention_backend is None:
declare_resolution(
server_args,
"_handle_cpu_backends",
attention_backend=(
"torch_native" if is_host_cpu_arm64() else "intel_amx"
),
)
declare_resolution(
server_args,
"_handle_cpu_backends",
sampling_backend="pytorch",
)
def handle_hpu_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "hpu":
declare_resolution(
server_args,
"_handle_hpu_backends",
attention_backend="torch_native",
)
declare_resolution(
server_args,
"_handle_hpu_backends",
sampling_backend="pytorch",
)
@@ -0,0 +1,906 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument resolution for serving-surface and multimodal entry validation."""
from __future__ import annotations
import json
import logging
import os
import random
import socket
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, Phase, with_phase
from sglang.srt.utils.common import (
configure_media_url_security,
get_device,
get_device_sm,
is_cuda,
is_hip,
is_mnnvl_fabric_device,
is_sm90_supported,
is_sm100_supported,
is_sm120_supported,
)
from sglang.utils import is_in_ci
logger = logging.getLogger(__name__)
def handle_ssl_validation(server_args: Any):
"""Ensure SSL arguments are consistent and referenced files exist."""
cfg = resolving_view(server_args)
if cfg.ssl_keyfile and not cfg.ssl_certfile:
raise ValueError(
"--ssl-keyfile requires --ssl-certfile to be specified as well."
)
if cfg.ssl_certfile and not cfg.ssl_keyfile:
raise ValueError(
"--ssl-certfile requires --ssl-keyfile to be specified as well."
)
if not cfg.ssl_certfile and not cfg.ssl_keyfile:
if cfg.ssl_ca_certs:
raise ValueError(
"--ssl-ca-certs has no effect without --ssl-certfile and --ssl-keyfile."
)
if cfg.ssl_keyfile_password:
raise ValueError(
"--ssl-keyfile-password has no effect without --ssl-certfile and --ssl-keyfile."
)
# Validate files exist early to avoid late failures after model loading.
if cfg.ssl_keyfile and not os.path.isfile(cfg.ssl_keyfile):
raise ValueError(
f"SSL key file not found: '{cfg.ssl_keyfile}'. "
f"Please check the --ssl-keyfile path."
)
if cfg.ssl_certfile and not os.path.isfile(cfg.ssl_certfile):
raise ValueError(
f"SSL certificate file not found: '{cfg.ssl_certfile}'. "
f"Please check the --ssl-certfile path."
)
if cfg.ssl_ca_certs and not os.path.isfile(cfg.ssl_ca_certs):
raise ValueError(
f"SSL CA certificates file not found: '{cfg.ssl_ca_certs}'. "
f"Please check the --ssl-ca-certs path."
)
if cfg.enable_ssl_refresh and not (cfg.ssl_certfile and cfg.ssl_keyfile):
raise ValueError(
"--enable-ssl-refresh requires --ssl-certfile and --ssl-keyfile "
"to be specified."
)
if cfg.enable_http2:
if not 0 < cfg.http2_max_concurrent_streams < 2**32:
raise ValueError(
"--http2-max-concurrent-streams must be between 1 and " "4294967295."
)
try:
import granian # noqa: F401
except ImportError:
raise ValueError(
"--enable-http2 requires the 'granian' package. "
'Install it with: pip install "sglang[http2]"'
)
if cfg.enable_ssl_refresh:
raise ValueError(
"--enable-ssl-refresh is not supported with --enable-http2. "
"Granian does not support SSL certificate hot-reloading. "
"Use Uvicorn (the default) or handle certificate rotation externally."
)
def handle_asr_validation(server_args: Any):
"""Validate transcription/ASR-specific server args."""
cfg = resolving_view(server_args)
if cfg.asr_max_buffer_seconds <= 0:
raise ValueError(
f"--asr-max-buffer-seconds must be positive "
f"(got {cfg.asr_max_buffer_seconds})."
)
if cfg.asr_max_concurrent_sessions <= 0:
raise ValueError(
f"--asr-max-concurrent-sessions must be positive "
f"(got {cfg.asr_max_concurrent_sessions})."
)
def handle_multimodal(server_args: Any):
"""Validate mm_process_config structure before model loading."""
cfg = resolving_view(server_args)
if (
cfg.mm_preprocess_cache_size_mb is not None
and cfg.mm_preprocess_cache_size_mb < 0
):
raise ValueError("mm_preprocess_cache_size_mb must be non-negative")
if cfg.mm_process_config is not None:
if not isinstance(cfg.mm_process_config, dict):
raise TypeError(
f"mm_process_config must be a dict, "
f"but got {type(cfg.mm_process_config)}"
)
for key in ("image", "video", "audio"):
if key in cfg.mm_process_config and not isinstance(
cfg.mm_process_config[key], dict
):
raise TypeError(
f"mm_process_config['{key}'] must be a dict, "
f"but got {type(cfg.mm_process_config[key])}"
)
def handle_crash_dump_env(server_args: Any):
cfg = resolving_view(server_args)
if not cfg.crash_dump_folder:
return
_CUDA_COREDUMP_DEFAULTS = {
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION": "1",
"CUDA_ENABLE_USER_TRIGGERED_COREDUMP": "1",
"CUDA_COREDUMP_SHOW_PROGRESS": "1",
"CUDA_COREDUMP_GENERATION_FLAGS": (
"skip_nonrelocated_elf_images,skip_global_memory,"
"skip_shared_memory,skip_local_memory,skip_constbank_memory"
),
"CUDA_COREDUMP_FILE": f"{cfg.crash_dump_folder}/%h/core.cuda.%t.%p",
"CUDA_COREDUMP_PIPE": "/tmp/corepipe.cuda.%h.%p",
}
for key, value in _CUDA_COREDUMP_DEFAULTS.items():
if key not in os.environ:
os.environ[key] = value
logger.info("Auto-set %s=%s (from --crash-dump-folder)", key, value)
coredump_dir = os.path.dirname(
os.environ["CUDA_COREDUMP_FILE"].replace("%h", socket.gethostname())
)
if "%" in coredump_dir:
logger.warning(
"Cannot pre-create CUDA coredump directory %s: only %%h is "
"supported in the directory part of CUDA_COREDUMP_FILE; "
"coredumps may fail to write.",
coredump_dir,
)
elif coredump_dir:
try:
os.makedirs(coredump_dir, exist_ok=True)
except OSError as e:
logger.warning(
"Failed to create CUDA coredump directory %s: %s; "
"coredumps may fail to write.",
coredump_dir,
e,
)
def handle_media_url_security(server_args: Any):
"""Normalize and publish the media URL policy before workers start."""
cfg = resolving_view(server_args)
declare_resolution(
server_args,
"_handle_media_url_security",
allowed_media_domains=configure_media_url_security(
cfg.allowed_media_domains,
cfg.media_url_max_file_size_mb,
),
)
def handle_load_balance_method(server_args: Any):
cfg = resolving_view(server_args)
if cfg.disaggregation_mode not in ("null", "prefill", "decode"):
raise ValueError(f"Invalid disaggregation_mode={cfg.disaggregation_mode!r}")
if cfg.load_balance_method == "auto":
# Default behavior:
# - non-PD: round_robin
# - PD prefill: follow_bootstrap_room
# - PD decode: round_robin
declare_resolution(
server_args,
"_handle_load_balance_method",
load_balance_method=(
"follow_bootstrap_room"
if cfg.disaggregation_mode == "prefill"
else "round_robin"
),
)
return
def handle_grammar_backend(server_args: Any):
cfg = resolving_view(server_args)
if cfg.grammar_backend is None:
declare_resolution(
server_args, "_handle_grammar_backend", grammar_backend="xgrammar"
)
def handle_debug_utils(server_args: Any):
cfg = resolving_view(server_args)
if is_in_ci() and cfg.soft_watchdog_timeout is None:
logger.info("Set soft_watchdog_timeout since in CI")
declare_resolution(
server_args, "_handle_debug_utils", soft_watchdog_timeout=300
)
def handle_deprecated_args(server_args: Any):
cfg = resolving_view(server_args)
if cfg.disable_fast_image_processor:
if cfg.image_processor_backend not in {"auto", "pil"}:
raise ValueError(
"--disable-fast-image-processor conflicts with "
f"--image-processor-backend={cfg.image_processor_backend}."
)
logger.warning(
"--disable-fast-image-processor is deprecated; use "
"--image-processor-backend=pil instead."
)
declare_resolution(
server_args, "_handle_deprecated_args", image_processor_backend="pil"
)
# Handle deprecated tool call parsers
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
if cfg.tool_call_parser in deprecated_tool_call_parsers:
logger.warning(
f"The tool_call_parser '{cfg.tool_call_parser}' is deprecated. Please use '{deprecated_tool_call_parsers[cfg.tool_call_parser]}' instead."
)
declare_resolution(
server_args,
"_handle_deprecated_args",
tool_call_parser=deprecated_tool_call_parsers[cfg.tool_call_parser],
)
# When user passes --enable-flashinfer-allreduce-fusion, enable with auto backend
if (
cfg.enable_flashinfer_allreduce_fusion
and cfg.flashinfer_allreduce_fusion_backend is None
):
logger.warning(
"--enable-flashinfer-allreduce-fusion is deprecated. "
"Please use --flashinfer-allreduce-fusion-backend=auto instead."
)
declare_resolution(
server_args,
"_handle_deprecated_args",
flashinfer_allreduce_fusion_backend="auto",
)
declare_resolution(
server_args,
"_handle_deprecated_args",
enable_flashinfer_allreduce_fusion=False,
)
# Deprecated attention-backend alias: "compressed" -> "dsv4".
renamed = {}
for attr in (
"attention_backend",
"decode_attention_backend",
"prefill_attention_backend",
"speculative_draft_attention_backend",
):
if getattr(server_args, attr, None) == "compressed":
logger.warning(
"--%s=compressed is deprecated; use 'dsv4' instead.",
attr.replace("_", "-"),
)
renamed[attr] = "dsv4"
if renamed:
declare_resolution(server_args, "_handle_deprecated_args", **renamed)
# --grpc-mode is a deprecated alias for --smg-grpc-mode.
if cfg.grpc_mode and not cfg.smg_grpc_mode:
logger.warning(
"--grpc-mode is deprecated and will be removed in a future "
"version. Use --smg-grpc-mode for the legacy SMG gRPC server, "
"or --grpc-port for the native gRPC server."
)
declare_resolution(
server_args,
"_handle_deprecated_args",
smg_grpc_mode=True,
)
# Native gRPC tuning knob is env-only; --grpc-port (CLI) enables the
# native server, falling back to SGLANG_GRPC_PORT.
declare_resolution(
server_args,
"_handle_deprecated_args",
grpc_worker_threads=envs.SGLANG_GRPC_WORKER_THREADS.get(),
)
grpc_port_env = envs.SGLANG_GRPC_PORT.get()
if cfg.grpc_port is None and grpc_port_env is not None:
declare_resolution(
server_args,
"_handle_deprecated_args",
grpc_port=grpc_port_env,
)
# Legacy SMG defaults its port to --port + 10000. Derive/validate only
# when gRPC is in use, so HTTP-only high ports don't fail validation.
legacy_grpc = cfg.smg_grpc_mode or cfg.grpc_mode
if legacy_grpc and cfg.grpc_port is None:
declare_resolution(
server_args,
"_handle_deprecated_args",
grpc_port=cfg.port + 10000,
)
if cfg.grpc_port is not None:
if not (1 <= cfg.grpc_port <= 65535):
raise ValueError(
"--grpc-port / SGLANG_GRPC_PORT "
f"({cfg.grpc_port}) must be between 1 and 65535"
)
if cfg.grpc_worker_threads is not None and cfg.grpc_worker_threads < 1:
raise ValueError(
"SGLANG_GRPC_WORKER_THREADS "
f"({cfg.grpc_worker_threads}) must be >= 1"
)
# Native gRPC is incompatible with launch paths it doesn't wire into.
# Legacy takes precedence over grpc_port, keeping re-runs idempotent.
native_grpc = cfg.grpc_port is not None and not legacy_grpc
if cfg.sidecar_args is not None:
if cfg.sidecar is None:
raise ValueError("--sidecar-args requires --sidecar.")
if not isinstance(cfg.sidecar_args, list) or not all(
isinstance(arg, str) for arg in cfg.sidecar_args
):
raise ValueError("--sidecar-args must be a JSON array of strings.")
if cfg.sidecar is not None:
if not cfg.sidecar.strip():
raise ValueError("--sidecar must not be empty.")
if legacy_grpc:
raise ValueError(
"--sidecar requires SGLang's native gRPC server; "
"it cannot be combined with --smg-grpc-mode/--grpc-mode."
)
if cfg.grpc_port is None:
raise ValueError("--sidecar requires --grpc-port or SGLANG_GRPC_PORT.")
if native_grpc:
if cfg.use_ray:
raise ValueError(
"--grpc-port is not supported with --use-ray: the Ray "
"serve launch path does not start the native gRPC server."
)
if cfg.encoder_only:
raise ValueError(
"--grpc-port is not supported with --encoder-only: "
"encoder disaggregation uses its own server."
)
if cfg.tokenizer_worker_num > 1:
raise ValueError(
"Native gRPC does not yet support --tokenizer-worker-num > 1. "
"Unset --grpc-port or set --tokenizer-worker-num 1."
)
if cfg.api_key or cfg.admin_api_key:
raise ValueError(
"--grpc-port is incompatible with --api-key/--admin-api-key: "
"the native gRPC listener bypasses HTTP auth middleware."
)
def handle_environment_variables(server_args: Any):
cfg = resolving_view(server_args)
server_args._handle_multimodal_feature_transport()
envs.SGLANG_ENABLE_TORCH_COMPILE.set("1" if cfg.enable_torch_compile else "0")
if cfg.mamba_ssm_dtype is not None:
envs.SGLANG_MAMBA_SSM_DTYPE.set(cfg.mamba_ssm_dtype)
envs.SGLANG_DISABLE_OUTLINES_DISK_CACHE.set(
"1" if cfg.disable_outlines_disk_cache else "0"
)
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.set(
"1" if cfg.enable_deterministic_inference else "0"
)
if cfg.enable_deterministic_inference:
envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.set("0")
if cfg.debug_cuda_graph:
if not (is_cuda() or is_hip()):
logger.warning(
"--debug-cuda-graph is not supported on non CUDA/HIP devices. "
"Disabling breakable CUDA graph."
)
declare_resolution(
server_args, "_handle_environment_variables", debug_cuda_graph=False
)
else:
envs.SGLANG_USE_BREAKABLE_CUDA_GRAPH.set("1")
logger.warning(
"Debug mode for CUDA graph is enabled via breakable CUDA graph. "
"All operations will run eagerly through the graph capture/replay path."
)
if cfg.enable_deepseek_v4_fp4_indexer and not (
is_sm100_supported() or is_sm120_supported()
):
raise ValueError(
"--enable-deepseek-v4-fp4-indexer requires SM100 or SM120 GPUs with "
"DeepGEMM FP4 indexer support."
)
# FP8 W_o GEMM needs DeepGEMM JIT. Enable exactly where the runtime can run
# it, mirroring the forward scale split: the ue8m0 path
# (DEEPGEMM_SCALE_UE8M0, true sm100, default on) or an sm90 opt-in
# fp32-scale path (use FP4 expert ckpt). Disable in every other case.
if is_cuda() and envs.SGLANG_OPT_FP8_WO_A_GEMM.get():
from sglang.srt.layers import deep_gemm_wrapper
sm = get_device_sm()
explicit = envs.SGLANG_OPT_FP8_WO_A_GEMM.is_set()
supported = deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 or (
deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and is_sm90_supported() and explicit
)
if not supported and explicit:
logger.warning(
"Disabling SGLANG_OPT_FP8_WO_A_GEMM: requires DeepGEMM JIT "
"and sm100+ (Blackwell), or explicit opt-in on sm90; "
"detected sm%d.",
sm,
)
if not supported:
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
def handle_other_validations(server_args: Any):
cfg = resolving_view(server_args)
if cfg.default_chat_template_kwargs is not None and not isinstance(
cfg.default_chat_template_kwargs, dict
):
raise ValueError("--default-chat-template-kwargs must decode to a JSON object")
# Handle optimistic prefill validation
if cfg.optimistic_prefill_attempts > 0 and cfg.disaggregation_mode == "prefill":
if cfg.pp_size > 1:
logger.warning("Optimistic prefill does not support pp_size > 1")
declare_resolution(
server_args,
"_handle_other_validations",
optimistic_prefill_attempts=0,
)
elif cfg.enable_hierarchical_cache and (
cfg.hicache_storage_backend is not None
or cfg.hicache_write_policy != "write_back"
):
logger.warning(
"Optimistic prefill only supports L2 hierarchical cache "
"with write-back policy"
)
declare_resolution(
server_args,
"_handle_other_validations",
optimistic_prefill_attempts=0,
)
elif resolved_view(server_args).uses_mamba_radix_cache:
logger.warning(
"Optimistic prefill does not support models that use "
"mamba radix cache."
)
declare_resolution(
server_args,
"_handle_other_validations",
optimistic_prefill_attempts=0,
)
# Handle model inference tensor dump.
if cfg.debug_tensor_dump_output_folder is not None:
logger.warning(
"Cuda graph and server warmup are disabled because of using tensor dump mode"
)
declare_resolution(
server_args,
"_handle_other_validations",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_other_validations",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
declare_resolution(
server_args, "_handle_other_validations", skip_server_warmup=True
)
if cfg.msprobe_dump_config is not None:
logger.warning(
"When msProbe is enabled, "
"cuda graph is disabled because msProbe only supports dump in eager mode, "
"warmup is disabled(skip_server_warmup=True) because there is no need to dump data for this stage."
)
declare_resolution(
server_args,
"_handle_other_validations",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED
),
)
declare_resolution(
server_args,
"_handle_other_validations",
cuda_graph_config=with_phase(
cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED
),
)
declare_resolution(
server_args, "_handle_other_validations", skip_server_warmup=True
)
# Validate limit_mm_per_prompt modalities
if cfg.limit_mm_data_per_request:
if isinstance(cfg.limit_mm_data_per_request, str):
declare_resolution(
server_args,
"_handle_other_validations",
limit_mm_data_per_request=json.loads(cfg.limit_mm_data_per_request),
)
if isinstance(cfg.limit_mm_data_per_request, dict):
allowed_modalities = {"image", "video", "audio"}
for modality in cfg.limit_mm_data_per_request.keys():
if modality not in allowed_modalities:
raise ValueError(
f"Invalid modality '{modality}' in --limit-mm-data-per-request."
f"Allowed modalities are: {list(allowed_modalities)}"
)
# Validate preferred_sampling_params
if cfg.preferred_sampling_params:
if isinstance(cfg.preferred_sampling_params, str):
declare_resolution(
server_args,
"_handle_other_validations",
preferred_sampling_params=json.loads(cfg.preferred_sampling_params),
)
# Validate preferred_sampling_params doesn't use tokenizer-dependent features
if cfg.skip_tokenizer_init:
from sglang.srt.sampling.sampling_params import SamplingParams
test_params = SamplingParams(**cfg.preferred_sampling_params)
# raises if tokenizer-dependent features used
test_params.normalize(None)
def handle_missing_default_values(server_args: Any):
cfg = resolving_view(server_args)
if cfg.tokenizer_path is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
tokenizer_path=cfg.model_path,
)
if cfg.served_model_name is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
served_model_name=cfg.model_path,
)
if cfg.device is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
device=get_device(),
)
# strip device index from user if any (e.g. "cuda:0" -> "cuda")
declare_resolution(
server_args,
"_handle_missing_default_values",
device=cfg.device.split(":")[0],
)
if cfg.random_seed is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
random_seed=random.randint(0, 1 << 30),
)
if cfg.mm_process_config is None:
declare_resolution(
server_args, "_handle_missing_default_values", mm_process_config={}
)
# Handle ModelScope model downloads
if envs.SGLANG_USE_MODELSCOPE.get():
server_args._handle_modelscope_paths()
# In speculative scenario:
# - If `speculative_draft_model_quantization` is specified, the draft model uses this quantization method.
# - Otherwise, the draft model defaults to the same quantization as the target model.
if cfg._speculative_draft_quantization_explicitly_set is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
_speculative_draft_quantization_explicitly_set=cfg.speculative_draft_model_quantization
is not None,
)
if cfg.speculative_draft_model_quantization is None:
declare_resolution(
server_args,
"_handle_missing_default_values",
speculative_draft_model_quantization=cfg.quantization,
)
# Resolve --quantization unquant before model config validation. Record
# the explicit opt-out so later auto-detection does not re-enable
# quantization.
if cfg.quantization == "unquant":
declare_resolution(
server_args,
"_handle_missing_default_values",
quantization=None,
)
server_args._quantization_explicitly_unset = True
else:
server_args._quantization_explicitly_unset = False
if cfg.speculative_draft_model_quantization == "unquant":
declare_resolution(
server_args,
"_handle_missing_default_values",
speculative_draft_model_quantization=None,
)
def handle_return_hidden_states_mode(server_args: Any):
cfg = resolving_view(server_args)
if cfg.return_hidden_states_mode not in (None, "last", "full"):
raise ValueError(
"return_hidden_states_mode must be one of: None, 'last', or 'full'."
)
if cfg.return_hidden_states_mode is None:
if cfg.enable_return_hidden_states:
declare_resolution(
server_args,
"_handle_return_hidden_states_mode",
return_hidden_states_mode="full",
)
else:
declare_resolution(
server_args,
"_handle_return_hidden_states_mode",
enable_return_hidden_states=True,
)
def handle_prefill_delayer_env_compat(server_args: Any):
if envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.get():
declare_resolution(
server_args,
"_handle_prefill_delayer_env_compat",
enable_prefill_delayer=True,
)
if x := envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get():
declare_resolution(
server_args,
"_handle_prefill_delayer_env_compat",
prefill_delayer_max_delay_passes=x,
)
if x := envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.get():
declare_resolution(
server_args,
"_handle_prefill_delayer_env_compat",
prefill_delayer_token_usage_low_watermark=x,
)
def handle_tokenizer_batching(server_args: Any):
cfg = resolving_view(server_args)
if cfg.enable_tokenizer_batch_encode and cfg.enable_dynamic_batch_tokenizer:
raise ValueError(
"Cannot enable both --enable-tokenizer-batch-encode and --enable-dynamic-batch-tokenizer. "
"Please choose one tokenizer batching approach."
)
if cfg.skip_tokenizer_init and not envs.SGLANG_RUST_SERVER.get():
# Tokenizer workers still serve HTTP / state / output work, so
# their fanout is preserved; detokenizer workers only decode.
if cfg.detokenizer_worker_num != 1:
logger.warning(
"skip_tokenizer_init=True leaves no decode work for detokenizer workers; "
f"forcing detokenizer_worker_num=1 (requested {cfg.detokenizer_worker_num})."
)
declare_resolution(
server_args, "_handle_tokenizer_batching", detokenizer_worker_num=1
)
if cfg.enable_tokenizer_batch_encode:
logger.warning(
"skip_tokenizer_init=True ignores --enable-tokenizer-batch-encode; disabling it."
)
declare_resolution(
server_args,
"_handle_tokenizer_batching",
enable_tokenizer_batch_encode=False,
)
if cfg.enable_dynamic_batch_tokenizer:
logger.warning(
"skip_tokenizer_init=True ignores --enable-dynamic-batch-tokenizer; disabling it."
)
declare_resolution(
server_args,
"_handle_tokenizer_batching",
enable_dynamic_batch_tokenizer=False,
)
logger.info(
"skip_tokenizer_init=True: string-based stop conditions (stop, stop_regex) "
"and min_new_tokens are unavailable."
)
def handle_multimodal_feature_transport(server_args: Any):
"""Resolve multimodal feature transport before tokenizer workers start.
CUDA IPC is opt-in because its fixed pool on ``base_gpu_id`` reduces the
memory left for model/KV-cache allocations. Multi-node MNNVL deployments
may still auto-select CUDA VMM. The legacy CUDA IPC flag and environment
variable remain supported so existing deployments map to this policy.
"""
cfg = resolving_view(server_args)
requested_transport = cfg.mm_feature_transport
legacy_ipc_is_set = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.is_set()
legacy_ipc_enabled = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
if cfg.keep_mm_feature_on_device:
if requested_transport not in (None, "cuda_ipc"):
raise ValueError(
"--keep-mm-feature-on-device conflicts with "
f"--mm-feature-transport={requested_transport}. Use only "
"--mm-feature-transport=cuda_ipc."
)
requested_transport = "cuda_ipc"
logger.warning(
"--keep-mm-feature-on-device is deprecated; using "
"--mm-feature-transport=cuda_ipc instead."
)
if requested_transport is None:
if legacy_ipc_is_set:
requested_transport = "cuda_ipc" if legacy_ipc_enabled else "cpu"
logger.warning(
"SGLANG_USE_CUDA_IPC_TRANSPORT is deprecated; use "
"--mm-feature-transport=%s instead.",
requested_transport,
)
elif cfg.encoder_only:
requested_transport = "cpu"
logger.info(
"Multimodal feature transport auto-resolved to cpu for "
"encoder-only serving; encoder outputs use "
"--encoder-transfer-backend instead."
)
elif (
server_args.get_model_config().is_multimodal
and is_cuda()
and cfg.disaggregation_mode == "null"
):
# A full GPU pool always degrades to CPU transport per tensor.
# Keep CUDA IPC opt-in because even an idle pool consumes HBM
# that would otherwise back the KV cache. Multi-node
# auto-selection is limited to GB200/GB300 systems where the
# runtime already enables the MNNVL/IMEX communication stack.
if cfg.nnodes == 1:
requested_transport = "cpu"
elif is_mnnvl_fabric_device() and os.path.exists(
"/dev/nvidia-caps-imex-channels/channel0"
):
from sglang.srt.model_loader.utils import (
supports_cuda_vmm_feature_transport,
)
if supports_cuda_vmm_feature_transport(server_args.get_model_config()):
requested_transport = "cuda_vmm"
logger.info(
"Multimodal feature transport auto-resolved to "
"cuda_vmm (multi-node GB200/GB300 MNNVL). Pass "
"--mm-feature-transport=cpu to opt out."
)
else:
requested_transport = "cpu"
logger.info(
"Multimodal feature transport auto-resolved to cpu: "
"the model has not opted into CUDA VMM transport."
)
else:
requested_transport = "cpu"
if is_mnnvl_fabric_device():
logger.info(
"Multimodal feature transport auto-resolved to cpu: "
"GB200/GB300 was detected but no IMEX channel is "
"mounted. Configure the MNNVL compute domain or pass "
"--mm-feature-transport=cuda_vmm after doing so."
)
else:
requested_transport = "cpu"
elif legacy_ipc_is_set and legacy_ipc_enabled != (
requested_transport == "cuda_ipc"
):
logger.warning(
"--mm-feature-transport=%s overrides the conflicting legacy "
"SGLANG_USE_CUDA_IPC_TRANSPORT=%s setting.",
requested_transport,
int(legacy_ipc_enabled),
)
if cfg.encoder_only and requested_transport in ("cuda_ipc", "cuda_vmm"):
logger.warning(
"--mm-feature-transport=%s does not control encoder-only "
"output transfer; using cpu for this inactive transport. Select "
"--encoder-transfer-backend for encoder outputs.",
requested_transport,
)
requested_transport = "cpu"
if requested_transport == "cuda_vmm":
if not is_cuda():
raise ValueError("--mm-feature-transport=cuda_vmm requires NVIDIA CUDA.")
if cfg.pp_size != 1:
raise ValueError(
"--mm-feature-transport=cuda_vmm does not support pipeline "
"parallelism."
)
if envs.SGLANG_RUST_SERVER.get():
raise ValueError(
"--mm-feature-transport=cuda_vmm is not supported with "
"SGLANG_RUST_SERVER."
)
pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get()
handle_kind = "CUDA FABRIC" if cfg.nnodes > 1 else "POSIX FD"
logger.info(
"Using CUDA VMM for multimodal features with %s sharing: "
"reserving up to %d MiB on base GPU %d across %d tokenizer "
"worker(s). This reduces KV cache headroom; a full pool falls "
"back to inline CPU transport.",
handle_kind,
pool_budget_mb,
cfg.base_gpu_id,
cfg.tokenizer_worker_num,
)
if requested_transport == "cuda_ipc":
if not is_cuda():
raise ValueError("--mm-feature-transport=cuda_ipc requires NVIDIA CUDA.")
if cfg.nnodes != 1:
raise ValueError(
"--mm-feature-transport=cuda_ipc only supports a single node."
)
pool_budget_mb = envs.SGLANG_MM_FEATURE_CACHE_MB.get()
logger.info(
"Using CUDA IPC for multimodal features: reserving up to %d MiB "
"on base GPU %d across %d tokenizer worker(s). This reduces KV "
"cache headroom; a full pool falls back to CPU transport.",
pool_budget_mb,
cfg.base_gpu_id,
cfg.tokenizer_worker_num,
)
logger.info(
"CUDA IPC pool-handle caching is %s. It reuses mappings to the "
"existing bounded pool without reserving another pool; set "
"SGLANG_USE_IPC_POOL_HANDLE_CACHE=0 to disable it.",
("enabled" if envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get() else "disabled"),
)
declare_resolution(
server_args,
"_handle_multimodal_feature_transport",
mm_feature_transport=requested_transport,
)
# The bounded IPC pool owns device residency. Do not retain unpooled
# tensors after a pool miss, which would make HBM use request-dependent.
declare_resolution(
server_args,
"_handle_multimodal_feature_transport",
keep_mm_feature_on_device=False,
)
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.set(
"1" if requested_transport == "cuda_ipc" else "0"
)
@@ -0,0 +1,430 @@
# SPDX-License-Identifier: Apache-2.0
"""Server-argument validation that spans no single family."""
from __future__ import annotations
import json
import logging
import os
from typing import Any, Dict, List, Optional
from sglang.srt.arg_groups.overrides import (
resolving_view,
)
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
parse_ib_device_config,
)
from sglang.srt.utils.common import is_hip, is_npu, torch_release
from sglang.srt.utils.runai_utils import is_runai_obj_uri
logger = logging.getLogger(__name__)
def check_server_args(server_args: Any):
cfg = resolving_view(server_args)
# Check parallel size constraints
if cfg.ep_join_mode != "scale":
assert (
cfg.tp_size * cfg.pp_size
) % cfg.nnodes == 0, "tp_size must be divisible by number of nodes"
assert cfg.pp_max_micro_batch_size is None or cfg.pp_max_micro_batch_size >= 1, (
"pp_max_micro_batch_size must be a positive integer or None (for auto-compute). "
f"Got: {cfg.pp_max_micro_batch_size}"
)
assert not (cfg.disable_cuda_graph_padding and cfg.enable_torch_compile), (
"--disable-cuda-graph-padding is incompatible with --enable-torch-compile. "
"With padding disabled, every distinct batch size gets its own torch.compile + "
"Triton autotune cycle (O(max_batch_size) compilations) instead of the small fixed "
"set of padded bucket sizes, causing engine initialisation to stall for many minutes. "
"Remove --disable-cuda-graph-padding or --enable-torch-compile."
)
if cfg.pp_size > 1:
assert (
cfg.disable_overlap_schedule and cfg.speculative_algorithm is None
), "Pipeline parallelism is not compatible with overlap schedule, speculative decoding"
assert cfg.min_free_slots_delay is None, (
"--min-free-slots-delay is not supported with pipeline "
"parallelism: allocatable slots per microbatch are bounded by "
"pp-max-micro-batch-size, so the threshold may never be reached"
)
assert not (
cfg.dp_size > 1 and cfg.nnodes != 1 and not cfg.enable_dp_attention
), "multi-node data parallel is not supported unless dp attention!"
assert cfg.base_gpu_id >= 0, "base_gpu_id must be non-negative"
assert cfg.gpu_id_step >= 1, "gpu_id_step must be positive"
assert cfg.moe_dense_tp_size in (
None,
1,
cfg.tp_size,
), "moe_dense_tp_size only supports None, 1, or tp_size currently"
# Check served model name to not have colon as it is reserved for LoRA adapter syntax
if not is_runai_obj_uri(cfg.served_model_name):
assert ":" not in cfg.served_model_name, (
"served_model_name cannot contain a colon (':') character. "
"The colon is reserved for the 'model:adapter' syntax used in LoRA adapter specification. "
f"Invalid value: '{cfg.served_model_name}'"
)
# Check LoRA
server_args.check_lora_server_args()
# Check speculative decoding
if cfg.speculative_algorithm is not None:
assert (
not cfg.enable_mixed_chunk
), "enable_mixed_chunk is required for speculative decoding"
# Check chunked prefill
# Skip validation if chunked prefill is disabled (i.e., size <= 0).
# Skip validation if disaggregation mode is decode.
if cfg.chunked_prefill_size > 0 and cfg.disaggregation_mode != "decode":
assert (
cfg.chunked_prefill_size % cfg.page_size == 0
), "chunked_prefill_size must be divisible by page_size"
# Check pdmux
if cfg.enable_pdmux:
assert (
cfg.pp_size == 1
), "PD-Multiplexing is only supported with pipeline parallelism disabled (pp_size=1)."
assert (
cfg.chunked_prefill_size == -1
), "PD-Multiplexing is not compatible with chunked prefill."
assert (
cfg.disaggregation_mode == "null"
), "PD-Multiplexing is not compatible with disaggregation mode."
assert (
cfg.disable_overlap_schedule
), "PD-Multiplexing is not compatible with overlap schedule."
# NOTE: CUDA Green Context may encounter potential issues with CudaGraph on torch 2.7.x – 2.8.x, leading to performance degradation.
import torch
if torch_release >= (2, 7):
logger.warning(
"WARNING: PD-Multiplexing may experience performance degradation with torch versions > 2.6.x.\n"
f" Current torch version is {torch.__version__}.\n"
" Please manually install torch 2.6.x."
)
assert cfg.tokenizer_worker_num > 0, "Tokenizer worker num must >= 1"
assert cfg.detokenizer_worker_num > 0, "Detokenizer worker num must >= 1"
assert cfg.mm_processor_worker_num >= 0, "Multimodal processor worker num must >= 0"
assert cfg.mm_io_worker_num >= 0, "Multimodal I/O worker num must >= 0"
server_args.validate_buckets_rule(
"--prompt-tokens-buckets", cfg.prompt_tokens_buckets
)
server_args.validate_buckets_rule(
"--generation-tokens-buckets", cfg.generation_tokens_buckets
)
# Check scheduling policy
if cfg.enable_priority_scheduling:
assert cfg.schedule_policy in [
"fcfs",
"lof",
], f"To use priority scheduling, schedule_policy must be 'fcfs' or 'lof'. '{cfg.schedule_policy}' is not supported."
if cfg.default_priority_value is None:
logger.warning(
"--default-priority-value is not set while --enable-priority-scheduling is enabled. "
"Requests without explicit priority will have priority=None, "
"resulting in priority='None' string labels in Prometheus metrics."
)
else:
if cfg.disable_priority_preemption:
logger.warning(
"--disable-priority-preemption has no effect without --enable-priority-scheduling"
)
if cfg.default_priority_value is not None:
logger.warning(
"--default-priority-value has no effect without --enable-priority-scheduling"
)
if cfg.retraction_policy == "priority" and not cfg.enable_priority_scheduling:
raise ValueError(
"--retraction-policy priority requires --enable-priority-scheduling"
)
# Check hisparse
# Moved to the resolution pipeline (arg_groups/overrides.py:
# _hisparse_validation), invoked here at its legacy slot.
from sglang.srt.arg_groups.overrides import (
_hisparse_validation,
run_post_process_pass,
)
run_post_process_pass(server_args, _hisparse_validation)
assert (
cfg.schedule_conservativeness >= 0
), "schedule_conservativeness must be non-negative"
if cfg.model_impl == "mindspore":
assert is_npu(), "MindSpore model impl is only supported on Ascend npu."
# Check metrics labels
if (
not cfg.tokenizer_metrics_custom_labels_header
and cfg.tokenizer_metrics_allowed_custom_labels
):
raise ValueError(
"Please set --tokenizer-metrics-custom-labels-header when setting --tokenizer-metrics-allowed-custom-labels."
)
# Check metrics exporters
if cfg.export_metrics_to_file and cfg.export_metrics_to_file_dir is None:
raise ValueError(
"--export-metrics-to-file-dir is required when --export-metrics-to-file is enabled"
)
# Check two batch overlap backend requirement.
server_args._check_two_batch_overlap()
# Check communications compression
if cfg.enable_quant_communications and cfg.tp_size == 1:
raise ValueError("Communications quantization is only used with tp_size != 1")
if cfg.enable_quant_communications and cfg.device != "npu":
raise ValueError("Communications quantization is only supported for NPU device")
# grpc_port is None for HTTP-only launches, so the == comparison is
# already False there; no explicit None check needed.
if not (cfg.smg_grpc_mode or cfg.grpc_mode) and cfg.grpc_port == cfg.port:
raise ValueError(
f"--grpc-port ({cfg.grpc_port}) must differ from --port ({cfg.port})"
)
# TODO: Also validate grpc_port != metrics_http_port and grpc_port != nccl_port
# to avoid opaque bind errors at runtime. Deferred because metrics_http_port
# and nccl_port have dynamic defaults that may not be resolved yet here.
if cfg.gc_threshold:
if not (1 <= len(cfg.gc_threshold) <= 3):
raise ValueError(
"When setting gc_threshold, it must contain 1 to 3 integers."
)
if cfg.kv_canary_sweep_interval > 0 and cfg.kv_canary == "none":
raise ValueError(
"--kv-canary-sweep-interval requires --kv-canary in {log, raise}"
)
server_args.check_load_publish_args()
def validate_buckets_rule(server_args: Any, arg_name: str, buckets_rule: List[str]):
if not buckets_rule:
return
assert len(buckets_rule) > 0, f"{arg_name} cannot be empty list"
rule = buckets_rule[0]
assert rule in [
"tse",
"default",
"custom",
], f"Unsupported {arg_name} rule type: '{rule}'. Must be one of: 'tse', 'default', 'custom'"
if rule == "tse":
assert (
len(buckets_rule) == 4
), f"{arg_name} TSE rule requires exactly 4 parameters: ['tse', middle, base, count], got {len(buckets_rule)}"
try:
middle = float(buckets_rule[1])
base = float(buckets_rule[2])
count = int(buckets_rule[3])
except (ValueError, IndexError):
assert (
False
), f"{arg_name} TSE rule parameters must be: ['tse', <float:middle>, <float:base>, <int:count>]"
assert base > 1, f"{arg_name} TSE base must be larger than 1, got: {base}"
assert count > 0, f"{arg_name} TSE count must be positive, got: {count}"
assert middle > 0, f"{arg_name} TSE middle must be positive, got: {middle}"
elif rule == "default":
assert (
len(buckets_rule) == 1
), f"{arg_name} default rule should only have one parameter: ['default'], got {len(buckets_rule)}"
elif rule == "custom":
assert (
len(buckets_rule) >= 2
), f"{arg_name} custom rule requires at least one bucket value: ['custom', value1, ...]"
try:
bucket_values = [float(x) for x in buckets_rule[1:]]
except ValueError:
assert False, f"{arg_name} custom rule bucket values must be numeric"
assert len(set(bucket_values)) == len(
bucket_values
), f"{arg_name} custom rule bucket values should not contain duplicates"
assert all(
val >= 0 for val in bucket_values
), f"{arg_name} custom rule bucket values should be non-negative"
def check_load_publish_args(server_args: Any):
"""Fail fast at the entrypoint on a --load-publish-endpoint the
scheduler would decline (no active kv-events publisher to advertise
through, unbindable, overlapping the KV range, u16 overflow) rather
than only warning — or silently doing nothing — from a scheduler
subprocess. Routes through the same resolver the scheduler binds and
/server_info advertises with."""
server_cfg = resolving_view(server_args)
mode = (server_cfg.load_publish_endpoint or "").strip()
if not mode or mode.lower() == "off":
return # disabled; nothing to validate
from sglang.srt.disaggregation.kv_events import (
KVEventsConfig,
resolve_load_pub_range,
)
if not server_cfg.kv_events_config:
raise ValueError(
"--load-publish-endpoint requires --kv-events-config: routers"
" discover the load range through /server_info's kv_events"
" block, absent without a publisher."
)
try:
cfg = KVEventsConfig.from_cli(server_cfg.kv_events_config)
except Exception as e:
raise ValueError(f"--kv-events-config is not parseable: {e}")
if cfg.publisher == "null" or not cfg.endpoint:
raise ValueError(
"--load-publish-endpoint needs an active --kv-events-config"
" publisher; got publisher='null' or an empty endpoint."
)
_, reason = resolve_load_pub_range(
kv_endpoint=cfg.endpoint,
replay_endpoint=cfg.replay_endpoint,
dp_size=server_cfg.dp_size,
load_publish_endpoint=mode,
)
if reason:
raise ValueError(reason)
def validate_ib_devices(server_args: Any, device_str: Optional[str]) -> Optional[str]:
"""
Validate IB devices before passing to mooncake.
Args:
device_str: Comma-separated IB device names, a per-GPU JSON mapping,
or a path to a JSON file containing that mapping.
Returns:
A normalized comma-separated string or per-GPU JSON mapping string, or None if input is None.
"""
if device_str is None:
logger.warning(
"No IB devices specified for Mooncake backend, falling back to auto discovery."
)
return None
def _normalize_device_group(raw_value: str, context: str) -> str:
if not isinstance(raw_value, str):
raise ValueError(
f"Invalid IB device format for {context}: expected a string. "
f"Got {type(raw_value)}"
)
devices = [d.strip() for d in raw_value.split(",") if d.strip()]
if not devices:
raise ValueError(f"No valid IB devices specified for {context}")
unique_devices = list(dict.fromkeys(devices))
if len(unique_devices) != len(devices):
logger.warning(
"Duplicate IB devices specified for %s: %s. Deduplicating to: %s",
context,
raw_value,
",".join(unique_devices),
)
invalid_devices = [d for d in unique_devices if d not in available_devices]
if len(invalid_devices) != 0:
raise ValueError(
f"Invalid IB devices specified for {context}: {invalid_devices}. "
f"Available devices: {sorted(available_devices)}"
)
return ",".join(unique_devices)
normalized_input = device_str.strip()
if not normalized_input:
raise ValueError("No valid IB devices specified")
# Get available IB devices from sysfs
ib_sysfs_path = "/sys/class/infiniband"
if not os.path.isdir(ib_sysfs_path):
raise RuntimeError(
f"InfiniBand sysfs path not found: {ib_sysfs_path}. "
"Please ensure InfiniBand drivers are installed."
)
available_devices = set(os.listdir(ib_sysfs_path))
if len(available_devices) == 0:
raise RuntimeError(f"No IB devices found in {ib_sysfs_path}")
parsed_config = parse_ib_device_config(normalized_input)
if isinstance(parsed_config, str):
return _normalize_device_group(normalized_input, "all GPUs")
assert parsed_config is not None
normalized_mapping: Dict[str, str] = {}
for gpu_key, gpu_devices in parsed_config.items():
normalized_key = str(gpu_key)
normalized_mapping[normalized_key] = _normalize_device_group(
gpu_devices, f"GPU {normalized_key}"
)
if not normalized_mapping:
raise ValueError("No valid GPU mappings found in IB device JSON")
return json.dumps(normalized_mapping, separators=(",", ":"))
def validate_experimental_sgl_marlin(server_args: Any):
view = server_args._resolved()
if view.moe_runner_backend != "experimental_sgl_marlin":
return
# ===== TO BE REFACTORED ====
from sglang.srt.lora.marlin_lora_temp.policy import (
validate_experimental_sgl_marlin_server_args,
)
validate_experimental_sgl_marlin_server_args(server_args, view)
def validate_prefill_decode_interval(server_args: Any):
cfg = resolving_view(server_args)
if cfg.prefill_decode_interval < 0:
raise ValueError("--prefill-decode-interval must be non-negative.")
def check_two_batch_overlap(server_args: Any):
# With no EP a2a backend, two-batch-overlap is only valid on the non-EP
# DP TP-MoE path (overlapping the DP all_gatherv / reduce_scatterv with
# the other ubatch's compute), which requires DP attention. Enabling it
# there needs no extra opt-in env flag.
cfg = resolving_view(server_args)
cp_tbo = (
is_hip()
and cfg.enable_dsa_prefill_context_parallel
and cfg.dsa_prefill_cp_mode == "round-robin-split"
)
if (
cfg.enable_two_batch_overlap
and cfg.moe_a2a_backend == "none"
and not cfg.enable_dp_attention
and not cp_tbo
):
raise ValueError(
"When enabling two batch overlap without an EP a2a backend "
"(moe_a2a_backend='none'), --enable-dp-attention is required "
"(DeepSeek-V4 non-EP DP TBO path)."
)
File diff suppressed because it is too large Load Diff
@@ -20,7 +20,7 @@ class TestServerArgsCPUBackend(unittest.TestCase):
server_args.sampling_backend = None server_args.sampling_backend = None
return server_args return server_args
@patch("sglang.srt.server_args.is_host_cpu_arm64", return_value=True) @patch("sglang.srt.arg_groups.platform_hook.is_host_cpu_arm64", return_value=True)
def test_arm_cpu_defaults_to_torch_native(self, _mock_is_arm64): def test_arm_cpu_defaults_to_torch_native(self, _mock_is_arm64):
server_args = self._make_server_args() server_args = self._make_server_args()
@@ -31,7 +31,7 @@ class TestServerArgsCPUBackend(unittest.TestCase):
) )
self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch") self.assertEqual(resolution_result(server_args, "sampling_backend"), "pytorch")
@patch("sglang.srt.server_args.is_host_cpu_arm64", return_value=False) @patch("sglang.srt.arg_groups.platform_hook.is_host_cpu_arm64", return_value=False)
def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64): def test_x86_cpu_defaults_to_intel_amx(self, _mock_is_arm64):
server_args = self._make_server_args() server_args = self._make_server_args()
@@ -184,7 +184,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
with ( with (
patch.object(args, "get_model_config", return_value=args._model_config), patch.object(args, "get_model_config", return_value=args._model_config),
patch("sglang.srt.server_args.is_cuda", return_value=True), patch("sglang.srt.arg_groups.model_hook.is_cuda", return_value=True),
): ):
args._handle_model_capability_adjustments() args._handle_model_capability_adjustments()
@@ -159,8 +159,11 @@ class TestHelionKDADispatcher(unittest.TestCase):
def test_replayssm_accepts_helion_and_rejects_other_backends(self): def test_replayssm_accepts_helion_and_rejects_other_backends(self):
with ( with (
patch("sglang.srt.server_args.is_sm100_supported", return_value=False), patch(
patch("sglang.srt.server_args.is_cuda", return_value=False), "sglang.srt.arg_groups.attention_hook.is_sm100_supported",
return_value=False,
),
patch("sglang.srt.arg_groups.attention_hook.is_cuda", return_value=False),
): ):
helion_args = ServerArgs( helion_args = ServerArgs(
model_path="dummy", model_path="dummy",
@@ -184,8 +187,11 @@ class TestHelionKDADispatcher(unittest.TestCase):
mamba_ssm_dtype="bfloat16", mamba_ssm_dtype="bfloat16",
) )
with ( with (
patch("sglang.srt.server_args.is_sm100_supported", return_value=True), patch(
patch("sglang.srt.server_args.is_cuda", return_value=False), "sglang.srt.arg_groups.attention_hook.is_sm100_supported",
return_value=True,
),
patch("sglang.srt.arg_groups.attention_hook.is_cuda", return_value=False),
): ):
args._handle_linear_attn_backend() args._handle_linear_attn_backend()
@@ -37,6 +37,20 @@ _READ_BEFORE_RESOLUTION = frozenset({"is_embedding"})
# has to be looked at. # has to be looked at.
_STALE_IN_THE_MODEL_CONFIG = frozenset({"speculative_algorithm"}) _STALE_IN_THE_MODEL_CONFIG = frozenset({"speculative_algorithm"})
# Behind the expert-pack build. `expert_pack_hook.handle_expert_pack` builds a
# model configuration, and it always did -- the walk stopped at the record's
# file and never saw it, so these three read as decided before the first build.
# The call sits behind `load_format != "expert_pack": return`, so it is the
# first build only on an expert-pack launch. Pre-existing; named rather than
# fixed, because fixing it means moving the build or the hook.
_STALE_BEHIND_THE_EXPERT_PACK_BUILD = frozenset(
{
"_speculative_draft_quantization_explicitly_set",
"model_path",
"speculative_draft_model_quantization",
}
)
# The same staleness through the registries: `_handle_model_specific_adjustments` # The same staleness through the registries: `_handle_model_specific_adjustments`
# builds the model configuration and *then* collects the override declarations, # builds the model configuration and *then* collects the override declarations,
# both inside one handler body. Named rather than fixed (that means moving the # both inside one handler body. Named rather than fixed (that means moving the
@@ -98,13 +112,27 @@ def _registry_collection_is_after_the_build():
collection above this handler's own `get_model_config()` call does not move collection above this handler's own `get_model_config()` call does not move
it above the configuration another handler already cached. it above the configuration another handler already cached.
""" """
tree = _parsed(_SRT / "server_args.py") handler = None
handler = next( for source, wanted in (
node (_SRT / "server_args.py", "_handle_model_specific_adjustments"),
for node in ast.walk(tree) *(
if isinstance(node, ast.FunctionDef) (path, "handle_model_specific_adjustments")
and node.name == "_handle_model_specific_adjustments" for path in sorted((_SRT / "arg_groups").glob("*.py"))
) ),
):
for node in ast.walk(_parsed(source)):
if isinstance(node, ast.FunctionDef) and node.name == wanted:
if any(
isinstance(child, ast.Call)
and getattr(child.func, "attr", getattr(child.func, "id", None))
== "collect_model_override_declarations"
for child in ast.walk(node)
):
handler = node
break
if handler is not None:
break
assert handler is not None, "the model-specific handler was not found"
build = collect = None build = collect = None
for node in ast.walk(handler): for node in ast.walk(handler):
if not isinstance(node, ast.Call): if not isinstance(node, ast.Call):
@@ -283,6 +311,21 @@ def _hook_declarations(dispatch, source_module):
return out return out
def _hook_functions():
"""Module-level resolution functions under `arg_groups/`.
A handler that moved out of the record leaves a slot behind that imports
one of these and calls it. Without following that hop the scan stops at
the slot and silently loses everything the handler does.
"""
functions = {}
for path in sorted((_SRT / "arg_groups").glob("*.py")):
for node in _parsed(path).body:
if isinstance(node, ast.FunctionDef):
functions.setdefault(node.name, node)
return functions
def _pipeline(): def _pipeline():
"""(ordered steps, {step: methods it reaches}) for the resolution dispatch.""" """(ordered steps, {step: methods it reaches}) for the resolution dispatch."""
tree = _parsed(_SRT / "server_args.py") tree = _parsed(_SRT / "server_args.py")
@@ -294,6 +337,28 @@ def _pipeline():
methods = { methods = {
node.name: node for node in record.body if isinstance(node, ast.FunctionDef) node.name: node for node in record.body if isinstance(node, ast.FunctionDef)
} }
hooks = _hook_functions()
# Follow exactly one edge: the slot's own `from arg_groups.X import f` /
# `f(self)`. Merging every hook function by bare name would let the walk
# wander into families the slot never calls.
slot_target = {}
for name, node in methods.items():
imported = {
alias.asname or alias.name
for child in ast.walk(node)
if isinstance(child, ast.ImportFrom)
and child.module
and child.module.startswith("sglang.srt.arg_groups")
for alias in child.names
}
called = {
child.func.id
for child in ast.walk(node)
if isinstance(child, ast.Call) and isinstance(child.func, ast.Name)
}
for target in sorted(imported & called & set(hooks)):
slot_target.setdefault(name, target)
methods.update({name: hooks[name] for name in slot_target.values()})
dispatch = methods["_run_resolution_pipeline"] dispatch = methods["_run_resolution_pipeline"]
steps = [ steps = [
name name
@@ -313,14 +378,18 @@ def _pipeline():
return seen return seen
seen.add(name) seen.add(name)
for node in ast.walk(methods[name]): for node in ast.walk(methods[name]):
if not isinstance(node, ast.Call):
continue
if ( if (
isinstance(node, ast.Call) isinstance(node.func, ast.Attribute)
and isinstance(node.func, ast.Attribute)
and isinstance(node.func.value, ast.Name) and isinstance(node.func.value, ast.Name)
and node.func.value.id == "self" and node.func.value.id == "self"
and node.func.attr in methods and node.func.attr in methods
): ):
reaches(node.func.attr, seen) reaches(node.func.attr, seen)
target = slot_target.get(name)
if target is not None:
reaches(target, seen)
return seen return seen
step_lines = {} step_lines = {}
@@ -512,6 +581,7 @@ class TestModelConfigReadsResolvedInput(CustomTestCase):
_READ_BEFORE_RESOLUTION _READ_BEFORE_RESOLUTION
| _STALE_IN_THE_MODEL_CONFIG | _STALE_IN_THE_MODEL_CONFIG
| _STALE_FROM_THE_REGISTRIES | _STALE_FROM_THE_REGISTRIES
| _STALE_BEHIND_THE_EXPERT_PACK_BUILD
) )
late = sorted( late = sorted(
field field
@@ -520,8 +520,11 @@ class TestProgramsResolveBeforeReadingResolution(CustomTestCase):
declarers = {"_declare", "declare_resolution", "declare_late_resolution"} declarers = {"_declare", "declare_resolution", "declare_late_resolution"}
fields = set() fields = set()
field_names = {field.name for field in _dataclasses.fields(_ServerArgs)} field_names = {field.name for field in _dataclasses.fields(_ServerArgs)}
for name in ("server_args.py", "arg_groups/overrides.py"): # The record plus every module under `arg_groups/`: a handler declares
tree = ast.parse((srt / name).read_text(encoding="utf-8-sig")) # from whichever of the two it lives in.
sources = [srt / "server_args.py", *sorted((srt / "arg_groups").rglob("*.py"))]
for source in sources:
tree = ast.parse(source.read_text(encoding="utf-8-sig"))
for node in ast.walk(tree): for node in ast.walk(tree):
# Registry data: provider dict keys are field names as # Registry data: provider dict keys are field names as
# *data*, invisible to the keyword scan below. Filtered # *data*, invisible to the keyword scan below. Filtered
@@ -72,6 +72,15 @@ _ATTRIBUTE_SPELLED = _BAG_ACCESSORS - {"get_device"}
_OWN = ("server_args.py", "runtime_context.py") _OWN = ("server_args.py", "runtime_context.py")
def _pipeline_sources():
"""The record plus every module under `arg_groups/`.
A handler that moved out of the record takes its imports with it, so
seeding the walk from two files would stop covering it.
"""
return [_SRT / "server_args.py", *sorted((_SRT / "arg_groups").rglob("*.py"))]
def _module_of(name): def _module_of(name):
"""`sglang.srt.a.b` -> the file, if it is one of ours.""" """`sglang.srt.a.b` -> the file, if it is one of ours."""
if not name or not name.startswith("sglang.srt."): if not name or not name.startswith("sglang.srt."):
@@ -196,9 +205,29 @@ def _functions_in(path):
} }
def _locally_shadowed_accessors(path):
"""Accessor names this file imports from somewhere that is not the context.
`get_device` is both the `device` bag accessor and the hardware probe in
`utils.common`. Matching the bare name would report the probe as a bag read,
so a name imported from elsewhere in this file is not the accessor.
"""
shadowed = set()
for node in ast.walk(ast.parse(path.read_text(encoding="utf-8-sig"))):
if isinstance(node, ast.ImportFrom) and node.module:
if node.module.endswith("runtime_context"):
continue
for alias in node.names:
name = alias.asname or alias.name
if name in _BAG_ACCESSORS:
shadowed.add(name)
return shadowed
def _reaches_a_bag(path, entry): def _reaches_a_bag(path, entry):
"""Does `entry` in `path` reach a bag accessor, following calls in-module?""" """Does `entry` in `path` reach a bag accessor, following calls in-module?"""
functions = _functions_in(path) functions = _functions_in(path)
shadowed = _locally_shadowed_accessors(path)
seen = set() seen = set()
def walk(name): def walk(name):
@@ -216,7 +245,7 @@ def _reaches_a_bag(path, entry):
continue continue
if not isinstance(node.func, ast.Name): if not isinstance(node.func, ast.Name):
continue continue
if node.func.id in _BAG_ACCESSORS: if node.func.id in _BAG_ACCESSORS and node.func.id not in shadowed:
return node.lineno return node.lineno
found = walk(node.func.id) found = walk(node.func.id)
if found is not None: if found is not None:
@@ -241,9 +270,7 @@ class TestResolutionReadsNoBag(CustomTestCase):
def test_the_walk_finds_something_to_walk(self): def test_the_walk_finds_something_to_walk(self):
"""A collapsed import map would make the pin vacuous.""" """A collapsed import map would make the pin vacuous."""
imported = _imported_symbols( imported = _imported_symbols(_pipeline_sources())
[_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"]
)
self.assertGreater( self.assertGreater(
len(imported), len(imported),
20, 20,
@@ -282,9 +309,7 @@ class TestResolutionReadsNoBag(CustomTestCase):
) )
def test_nothing_the_pipeline_calls_reads_a_bag(self): def test_nothing_the_pipeline_calls_reads_a_bag(self):
imported = _imported_symbols( imported = _imported_symbols(_pipeline_sources())
[_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"]
)
reachable = { reachable = {
(path, symbol) for path, symbols in imported.items() for symbol in symbols (path, symbol) for path, symbols in imported.items() for symbol in symbols
} | _registered_entries() } | _registered_entries()
@@ -9,7 +9,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import sglang.srt.server_args as server_args_module import sglang.srt.server_args as server_args_module
from sglang.srt.arg_groups import pd_disaggregation_hook from sglang.srt.arg_groups import parallel_hook, pd_disaggregation_hook, serving_hook
from sglang.srt.arg_groups.overrides import resolution_result from sglang.srt.arg_groups.overrides import resolution_result
from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding
from sglang.srt.entrypoints.sidecar import ( from sglang.srt.entrypoints.sidecar import (
@@ -40,7 +40,9 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
register_cpu_ci(est_time=11, suite="base-c-test-cpu") register_cpu_ci(est_time=11, suite="base-c-test-cpu")
# Mock get_device() so all tests run on CPU-only CI runners # Mock get_device() so all tests run on CPU-only CI runners
_mock_device = patch("sglang.srt.server_args.get_device", return_value="cuda") _mock_device = patch(
"sglang.srt.arg_groups.serving_hook.get_device", return_value="cuda"
)
_mock_device.start() _mock_device.start()
@@ -223,7 +225,7 @@ class TestMmEncoderDataParallelLogging(CustomTestCase):
model_path="dummy", mm_enable_dp_encoder=True, tp_size=1 model_path="dummy", mm_enable_dp_encoder=True, tp_size=1
) )
with self.assertLogs(server_args_module.logger, level="WARNING") as logs: with self.assertLogs(parallel_hook.logger, level="WARNING") as logs:
server_args._handle_data_parallelism() server_args._handle_data_parallelism()
self.assertIn("TP=1", logs.output[0]) self.assertIn("TP=1", logs.output[0])
@@ -234,7 +236,7 @@ class TestMmEncoderDataParallelLogging(CustomTestCase):
model_path="dummy", mm_enable_dp_encoder=True, tp_size=4 model_path="dummy", mm_enable_dp_encoder=True, tp_size=4
) )
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(parallel_hook.logger, level="INFO") as logs:
server_args._handle_data_parallelism() server_args._handle_data_parallelism()
self.assertIn("TP=4", logs.output[0]) self.assertIn("TP=4", logs.output[0])
@@ -255,7 +257,7 @@ class TestImageProcessorBackend(CustomTestCase):
def test_legacy_flag_maps_to_pil_with_one_warning(self): def test_legacy_flag_maps_to_pil_with_one_warning(self):
server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True) server_args = ServerArgs(model_path="dummy", disable_fast_image_processor=True)
with self.assertLogs(server_args_module.logger, level="WARNING") as logs: with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_deprecated_args() server_args._handle_deprecated_args()
self.assertEqual( self.assertEqual(
@@ -285,7 +287,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
def _set_model_type(server_args, *, is_multimodal): def _set_model_type(server_args, *, is_multimodal):
server_args._model_config = SimpleNamespace(is_multimodal=is_multimodal) server_args._model_config = SimpleNamespace(is_multimodal=is_multimodal)
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda): def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", model_path="dummy",
@@ -295,7 +297,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}): with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -307,12 +309,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("base GPU 2", output) self.assertIn("base GPU 2", output)
self.assertIn("4 tokenizer worker", output) self.assertIn("4 tokenizer worker", output)
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_legacy_keep_flag_maps_to_cuda_ipc(self, _mock_is_cuda): def test_legacy_keep_flag_maps_to_cuda_ipc(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", keep_mm_feature_on_device=True) server_args = ServerArgs(model_path="dummy", keep_mm_feature_on_device=True)
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}): with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "0"}):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs: with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -335,12 +337,12 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertRaisesRegex(ValueError, "conflicts.*cuda_vmm"): with self.assertRaisesRegex(ValueError, "conflicts.*cuda_vmm"):
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_explicit_cpu_overrides_legacy_environment(self, _mock_is_cuda): def test_explicit_cpu_overrides_legacy_environment(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cpu") server_args = ServerArgs(model_path="dummy", mm_feature_transport="cpu")
with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}): with patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}):
with self.assertLogs(server_args_module.logger, level="WARNING") as logs: with self.assertLogs(serving_hook.logger, level="WARNING") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -361,7 +363,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_text_only_model(self, _mock_is_cuda): def test_default_transport_is_cpu_for_text_only_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy") server_args = ServerArgs(model_path="dummy")
self._set_model_type(server_args, is_multimodal=False) self._set_model_type(server_args, is_multimodal=False)
@@ -376,7 +378,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_multimodal_model(self, _mock_is_cuda): def test_default_transport_is_cpu_for_multimodal_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy") server_args = ServerArgs(model_path="dummy")
self._set_model_type(server_args, is_multimodal=True) self._set_model_type(server_args, is_multimodal=True)
@@ -391,9 +393,11 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.os.path.exists", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True) @patch(
@patch("sglang.srt.server_args.is_cuda", return_value=True) "sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
@patch( @patch(
"sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport", "sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport",
return_value=True, return_value=True,
@@ -410,7 +414,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with patch.dict(os.environ, {}, clear=False): with patch.dict(os.environ, {}, clear=False):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear() envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -422,9 +426,11 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("auto-resolved to cuda_vmm", output) self.assertIn("auto-resolved to cuda_vmm", output)
self.assertIn("CUDA FABRIC", output) self.assertIn("CUDA FABRIC", output)
@patch("sglang.srt.server_args.os.path.exists", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True) @patch(
@patch("sglang.srt.server_args.is_cuda", return_value=True) "sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
@patch( @patch(
"sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport", "sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport",
return_value=False, return_value=False,
@@ -439,15 +445,17 @@ class TestMultimodalFeatureTransport(CustomTestCase):
server_args = ServerArgs(model_path="dummy", nnodes=2) server_args = ServerArgs(model_path="dummy", nnodes=2)
self._set_model_type(server_args, is_multimodal=True) self._set_model_type(server_args, is_multimodal=True)
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu") self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu")
self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output)) self.assertIn("has not opted into CUDA VMM", "\n".join(logs.output))
@patch("sglang.srt.server_args.os.path.exists", return_value=False) @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=False)
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True) @patch(
@patch("sglang.srt.server_args.is_cuda", return_value=True) "sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=True
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_without_imex_channel( def test_default_transport_is_cpu_without_imex_channel(
self, _mock_is_cuda, _mock_is_mnnvl, _mock_path_exists self, _mock_is_cuda, _mock_is_mnnvl, _mock_path_exists
): ):
@@ -456,7 +464,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with patch.dict(os.environ, {}, clear=False): with patch.dict(os.environ, {}, clear=False):
envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear() envs.SGLANG_USE_CUDA_IPC_TRANSPORT.clear()
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -465,8 +473,10 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("no IMEX channel", "\n".join(logs.output)) self.assertIn("no IMEX channel", "\n".join(logs.output))
@patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=False) @patch(
@patch("sglang.srt.server_args.is_cuda", return_value=True) "sglang.srt.arg_groups.serving_hook.is_mnnvl_fabric_device", return_value=False
)
@patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_multinode_non_mnnvl( def test_default_transport_is_cpu_for_multinode_non_mnnvl(
self, _mock_is_cuda, _mock_is_mnnvl self, _mock_is_cuda, _mock_is_mnnvl
): ):
@@ -482,7 +492,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_default_transport_is_cpu_for_language_only_model(self, _mock_is_cuda): def test_default_transport_is_cpu_for_language_only_model(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", language_only=True) server_args = ServerArgs(model_path="dummy", language_only=True)
self._set_model_type(server_args, is_multimodal=True) self._set_model_type(server_args, is_multimodal=True)
@@ -496,14 +506,14 @@ class TestMultimodalFeatureTransport(CustomTestCase):
) )
self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()) self.assertFalse(envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get())
@patch("sglang.srt.server_args.is_cuda", return_value=False) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=False)
def test_cuda_ipc_rejects_non_nvidia_platforms(self, _mock_is_cuda): def test_cuda_ipc_rejects_non_nvidia_platforms(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_ipc") server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_ipc")
with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"): with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"):
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_ipc_rejects_multi_node(self, _mock_is_cuda): def test_cuda_ipc_rejects_multi_node(self, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", mm_feature_transport="cuda_ipc", nnodes=2 model_path="dummy", mm_feature_transport="cuda_ipc", nnodes=2
@@ -512,7 +522,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
with self.assertRaisesRegex(ValueError, "single node"): with self.assertRaisesRegex(ValueError, "single node"):
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_vmm_is_explicit_and_uses_shared_budget(self, _mock_is_cuda): def test_cuda_vmm_is_explicit_and_uses_shared_budget(self, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", model_path="dummy",
@@ -525,7 +535,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}), patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}),
envs.SGLANG_MM_FEATURE_CACHE_MB.override(256), envs.SGLANG_MM_FEATURE_CACHE_MB.override(256),
): ):
with self.assertLogs(server_args_module.logger, level="INFO") as logs: with self.assertLogs(serving_hook.logger, level="INFO") as logs:
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
self.assertEqual( self.assertEqual(
@@ -539,14 +549,14 @@ class TestMultimodalFeatureTransport(CustomTestCase):
self.assertIn("2 tokenizer worker", output) self.assertIn("2 tokenizer worker", output)
self.assertIn("falls back to inline CPU", output) self.assertIn("falls back to inline CPU", output)
@patch("sglang.srt.server_args.is_cuda", return_value=False) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=False)
def test_cuda_vmm_rejects_non_nvidia_platforms(self, _mock_is_cuda): def test_cuda_vmm_rejects_non_nvidia_platforms(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm") server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"): with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"):
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_vmm_rejects_rust_server(self, _mock_is_cuda): def test_cuda_vmm_rejects_rust_server(self, _mock_is_cuda):
server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm") server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm")
@@ -556,7 +566,7 @@ class TestMultimodalFeatureTransport(CustomTestCase):
): ):
server_args._handle_multimodal_feature_transport() server_args._handle_multimodal_feature_transport()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.serving_hook.is_cuda", return_value=True)
def test_cuda_vmm_rejects_pipeline_parallelism(self, _mock_is_cuda): def test_cuda_vmm_rejects_pipeline_parallelism(self, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", mm_feature_transport="cuda_vmm", pp_size=2 model_path="dummy", mm_feature_transport="cuda_vmm", pp_size=2
@@ -577,7 +587,7 @@ class TestMambaCacheStochasticRounding(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "--mamba-ssm-dtype float16"): with self.assertRaisesRegex(ValueError, "--mamba-ssm-dtype float16"):
server_args._handle_mamba_backend() server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=False) @patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=False)
def test_rejects_non_cuda(self, _mock_is_cuda): def test_rejects_non_cuda(self, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", model_path="dummy",
@@ -588,8 +598,8 @@ class TestMambaCacheStochasticRounding(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "NVIDIA CUDA"): with self.assertRaisesRegex(ValueError, "NVIDIA CUDA"):
server_args._handle_mamba_backend() server_args._handle_mamba_backend()
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=True)
@patch("sglang.srt.server_args.is_sm100_supported", return_value=False) @patch("sglang.srt.arg_groups.mamba_hook.is_sm100_supported", return_value=False)
def test_rejects_triton_without_sm100(self, _mock_sm100, _mock_is_cuda): def test_rejects_triton_without_sm100(self, _mock_sm100, _mock_is_cuda):
server_args = ServerArgs( server_args = ServerArgs(
model_path="dummy", model_path="dummy",
@@ -1925,11 +1935,11 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
prefill=PhaseConfig(backend=Backend.TC_PIECEWISE) prefill=PhaseConfig(backend=Backend.TC_PIECEWISE)
) )
with ( with (
patch("sglang.srt.server_args.is_hip", return_value=False), patch("sglang.srt.arg_groups.cuda_graph_hook.is_hip", return_value=False),
patch("sglang.srt.server_args.is_npu", return_value=False), patch("sglang.srt.arg_groups.cuda_graph_hook.is_npu", return_value=False),
patch("sglang.srt.server_args.is_cpu", return_value=False), patch("sglang.srt.arg_groups.cuda_graph_hook.is_cpu", return_value=False),
patch("sglang.srt.server_args.is_mps", return_value=False), patch("sglang.srt.arg_groups.cuda_graph_hook.is_mps", return_value=False),
patch("sglang.srt.server_args.is_xpu", return_value=False), patch("sglang.srt.arg_groups.cuda_graph_hook.is_xpu", return_value=False),
): ):
args._disable_tc_piecewise_cudagraph_if_incompatible() args._disable_tc_piecewise_cudagraph_if_incompatible()
@@ -2661,7 +2671,7 @@ class TestGrpcServerArgs(CustomTestCase):
def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self): def test_grpc_mode_is_deprecated_alias_for_smg_grpc_mode(self):
sa = self._args(grpc_mode=True) sa = self._args(grpc_mode=True)
with self.assertLogs(server_args_module.logger, level="WARNING") as cm: with self.assertLogs(serving_hook.logger, level="WARNING") as cm:
sa._handle_deprecated_args() sa._handle_deprecated_args()
self.assertTrue(resolution_result(sa, "smg_grpc_mode")) self.assertTrue(resolution_result(sa, "smg_grpc_mode"))
self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output)) self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output))
@@ -46,7 +46,7 @@ class TestValidateMambaExtraBufferLazyDflash(CustomTestCase):
), mock.patch( ), mock.patch(
# Keep the test runnable on CPU-only hosts: the platform assert is # Keep the test runnable on CPU-only hosts: the platform assert is
# not what is under test here. # not what is under test here.
"sglang.srt.server_args.is_cuda", "sglang.srt.arg_groups.mamba_hook.is_cuda",
return_value=True, return_value=True,
): ):
ServerArgs._validate_mamba_extra_buffer( ServerArgs._validate_mamba_extra_buffer(
@@ -69,6 +69,27 @@ class TestValidateMambaExtraBufferLazyDflash(CustomTestCase):
_lazy_view(speculative_num_draft_tokens=512, mamba_track_interval=256) _lazy_view(speculative_num_draft_tokens=512, mamba_track_interval=256)
) )
def test_the_chunk_size_is_not_read_before_the_page_size_resolves(self):
"""`mamba_cache_chunk_size` is derived from `page_size`, which the
pipeline writes *after* `_handle_model_specific_adjustments` runs this
validator. The read has to stay inside the `page_size is not None`
guard: evaluating it at the call site raises `TypeError` on the
unresolved `None` (hit by Qwen3-Next under PD disaggregation)."""
from sglang.srt.arg_groups.mamba_hook import validate_mamba_extra_buffer
def _must_not_be_read():
raise AssertionError("the chunk size was read before page_size resolved")
with mock.patch(
"sglang.srt.arg_groups.overrides.supports_mamba_cache_extra_buffer",
return_value=True,
), mock.patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=True):
validate_mamba_extra_buffer(
_lazy_view(page_size=None),
"Qwen3NextForCausalLM",
mamba_cache_chunk_size_of=_must_not_be_read,
)
class TestDflashVerifyRunsMambaTrackHook(CustomTestCase): class TestDflashVerifyRunsMambaTrackHook(CustomTestCase):
"""prepare_for_verify calls prepare_mamba_track_for_verify after the batch """prepare_for_verify calls prepare_mamba_track_for_verify after the batch
@@ -195,9 +195,12 @@ def _declared_by_late_resolution():
It forwards `**fields` to `declare_late_resolution`, so the keywords sit at 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. its call sites and a scan for the declarer's own name finds none of them.
""" """
tree = ast.parse((_SRT / "server_args.py").read_text(encoding="utf-8-sig")) # The record plus `arg_groups/`: a hook calls it on the record it was
# handed, so scanning the record's file alone finds nothing.
sources = [_SRT / "server_args.py", *sorted((_SRT / "arg_groups").rglob("*.py"))]
fields = set() fields = set()
for node in ast.walk(tree): for source in sources:
for node in ast.walk(ast.parse(source.read_text(encoding="utf-8-sig"))):
if ( if (
isinstance(node, ast.Call) isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute) and isinstance(node.func, ast.Attribute)
@@ -1236,6 +1236,33 @@ class TestGoldenModelOverrides(_IsolatedPublish):
"operator's input", "operator's input",
) )
def test_a_pass_that_declares_nothing_runs_on_the_published_record(self):
"""A validation slot has to survive a rebuild on the same record.
`Engine.shutdown()` leaves the launch published, and `Engine(server_args=sa)`
with the same instance calls `check_server_args()` again before
republishing. `_hisparse_validation` reaches the pass runner from there
and returns nothing, so refusing on identity alone would fail the
second launch.
"""
from sglang.srt.arg_groups.overrides import run_post_process_pass
from sglang.srt.runtime_context import publish, reset_context
sa = self._construct("LlamaForCausalLM", "llama")
self.addCleanup(reset_context)
publish(sa, role="scheduler")
def _declares_nothing(view):
return {}
run_post_process_pass(sa, _declares_nothing) # must not raise
def _declares_something(view):
return {"attention_backend": "triton"}
with self.assertRaisesRegex(ValueError, r"on the published config"):
run_post_process_pass(sa, _declares_something)
def test_attention_backend_user_choice_declares_nothing_extra(self): def test_attention_backend_user_choice_declares_nothing_extra(self):
sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton") sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton")
self.assertEqual(self._resolved(sa, "attention_backend"), "triton") self.assertEqual(self._resolved(sa, "attention_backend"), "triton")
@@ -511,12 +511,28 @@ class TestSuppliedInstanceExposure(CustomTestCase):
"prefill_attention_backend", "prefill_attention_backend",
"speculative_draft_attention_backend", "speculative_draft_attention_backend",
} }
deprecated = next(
node # The handler lives in `arg_groups/serving_hook.py`, reached either as a
for node in ast.walk(sa_class) # record method or as a bare-name call, so look the loop up by both.
if isinstance(node, ast.FunctionDef) def _deprecated_alias_handler():
for node in ast.walk(sa_class):
if (
isinstance(node, ast.FunctionDef)
and node.name == "_handle_deprecated_args" and node.name == "_handle_deprecated_args"
) and any(isinstance(n, ast.For) for n in ast.walk(node))
):
return node
for path in sorted((_PACKAGE_ROOT / "arg_groups").glob("*.py")):
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
for node in tree.body:
if (
isinstance(node, ast.FunctionDef)
and node.name == "handle_deprecated_args"
):
return node
raise AssertionError("the deprecated-alias handler was not found")
deprecated = _deprecated_alias_handler()
found_tuples = [ found_tuples = [
{elt.value for elt in node.iter.elts if isinstance(elt, ast.Constant)} {elt.value for elt in node.iter.elts if isinstance(elt, ast.Constant)}
for node in ast.walk(deprecated) for node in ast.walk(deprecated)
@@ -10,7 +10,9 @@ register_cpu_ci(est_time=7, suite="base-a-test-cpu")
register_cpu_ci(est_time=5, suite="base-c-test-cpu") register_cpu_ci(est_time=5, suite="base-c-test-cpu")
# Mock get_device() so ServerArgs tests run on CPU-only CI runners # Mock get_device() so ServerArgs tests run on CPU-only CI runners
_mock_device = patch("sglang.srt.server_args.get_device", return_value="cuda") _mock_device = patch(
"sglang.srt.arg_groups.serving_hook.get_device", return_value="cuda"
)
_mock_device.start() _mock_device.start()