config: the resolution pipeline moves out of the record (#36789)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
726665e08e
commit
c2928e86d7
@@ -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]}"
|
||||||
|
)
|
||||||
@@ -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."
|
||||||
|
)
|
||||||
@@ -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}."
|
||||||
|
)
|
||||||
@@ -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."
|
||||||
@@ -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."
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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}."
|
||||||
|
)
|
||||||
@@ -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)."
|
||||||
|
)
|
||||||
+232
-5529
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()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user