From c2928e86d78ee99a295c313409b7e8918d0c1b03 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 28 Aug 2026 10:17:24 -0700 Subject: [PATCH] config: the resolution pipeline moves out of the record (#36789) Co-authored-by: Claude Opus 5 --- .../sglang/srt/arg_groups/attention_hook.py | 627 ++ .../sglang/srt/arg_groups/cuda_graph_hook.py | 455 ++ python/sglang/srt/arg_groups/dllm_hook.py | 124 + python/sglang/srt/arg_groups/hicache_hook.py | 209 + python/sglang/srt/arg_groups/kv_cache_hook.py | 425 ++ python/sglang/srt/arg_groups/lora_hook.py | 220 + python/sglang/srt/arg_groups/mamba_hook.py | 154 + python/sglang/srt/arg_groups/memory_hook.py | 268 + python/sglang/srt/arg_groups/model_hook.py | 856 +++ .../sglang/srt/arg_groups/model_path_hook.py | 306 + python/sglang/srt/arg_groups/moe_hook.py | 477 ++ python/sglang/srt/arg_groups/overrides.py | 40 +- python/sglang/srt/arg_groups/parallel_hook.py | 658 ++ .../srt/arg_groups/pd_disaggregation_hook.py | 80 +- python/sglang/srt/arg_groups/platform_hook.py | 133 + python/sglang/srt/arg_groups/serving_hook.py | 906 +++ .../sglang/srt/arg_groups/validation_hook.py | 430 ++ python/sglang/srt/server_args.py | 5781 +---------------- .../cpu/test_server_args_backend.py | 4 +- .../test_multimodal_piecewise_cuda_graph.py | 2 +- .../attention/test_kda_helion_dispatcher.py | 14 +- .../test_model_config_reads_resolved_input.py | 88 +- .../test_resolution_is_reproducible.py | 7 +- .../test_resolution_reads_no_bag.py | 39 +- .../unit/server_args/test_server_args.py | 98 +- .../spec/test_dflash_extra_buffer_lazy.py | 23 +- .../unit/test_chain_read_ratchet.py | 19 +- test/registered/unit/test_model_overrides.py | 27 + ...test_supplied_instance_exposure_ratchet.py | 28 +- test/registered/utils/test_network_address.py | 4 +- 30 files changed, 6863 insertions(+), 5639 deletions(-) create mode 100644 python/sglang/srt/arg_groups/attention_hook.py create mode 100644 python/sglang/srt/arg_groups/cuda_graph_hook.py create mode 100644 python/sglang/srt/arg_groups/dllm_hook.py create mode 100644 python/sglang/srt/arg_groups/hicache_hook.py create mode 100644 python/sglang/srt/arg_groups/kv_cache_hook.py create mode 100644 python/sglang/srt/arg_groups/lora_hook.py create mode 100644 python/sglang/srt/arg_groups/mamba_hook.py create mode 100644 python/sglang/srt/arg_groups/memory_hook.py create mode 100644 python/sglang/srt/arg_groups/model_hook.py create mode 100644 python/sglang/srt/arg_groups/model_path_hook.py create mode 100644 python/sglang/srt/arg_groups/moe_hook.py create mode 100644 python/sglang/srt/arg_groups/parallel_hook.py create mode 100644 python/sglang/srt/arg_groups/platform_hook.py create mode 100644 python/sglang/srt/arg_groups/serving_hook.py create mode 100644 python/sglang/srt/arg_groups/validation_hook.py diff --git a/python/sglang/srt/arg_groups/attention_hook.py b/python/sglang/srt/arg_groups/attention_hook.py new file mode 100644 index 000000000..87cecd947 --- /dev/null +++ b/python/sglang/srt/arg_groups/attention_hook.py @@ -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." + ) diff --git a/python/sglang/srt/arg_groups/cuda_graph_hook.py b/python/sglang/srt/arg_groups/cuda_graph_hook.py new file mode 100644 index 000000000..d37e48121 --- /dev/null +++ b/python/sglang/srt/arg_groups/cuda_graph_hook.py @@ -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]}" + ) diff --git a/python/sglang/srt/arg_groups/dllm_hook.py b/python/sglang/srt/arg_groups/dllm_hook.py new file mode 100644 index 000000000..c59063602 --- /dev/null +++ b/python/sglang/srt/arg_groups/dllm_hook.py @@ -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, + ) diff --git a/python/sglang/srt/arg_groups/hicache_hook.py b/python/sglang/srt/arg_groups/hicache_hook.py new file mode 100644 index 000000000..850100f19 --- /dev/null +++ b/python/sglang/srt/arg_groups/hicache_hook.py @@ -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." + ) diff --git a/python/sglang/srt/arg_groups/kv_cache_hook.py b/python/sglang/srt/arg_groups/kv_cache_hook.py new file mode 100644 index 000000000..710a0831f --- /dev/null +++ b/python/sglang/srt/arg_groups/kv_cache_hook.py @@ -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." + ) diff --git a/python/sglang/srt/arg_groups/lora_hook.py b/python/sglang/srt/arg_groups/lora_hook.py new file mode 100644 index 000000000..801cc873e --- /dev/null +++ b/python/sglang/srt/arg_groups/lora_hook.py @@ -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}." + ) diff --git a/python/sglang/srt/arg_groups/mamba_hook.py b/python/sglang/srt/arg_groups/mamba_hook.py new file mode 100644 index 000000000..3bb704dde --- /dev/null +++ b/python/sglang/srt/arg_groups/mamba_hook.py @@ -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." diff --git a/python/sglang/srt/arg_groups/memory_hook.py b/python/sglang/srt/arg_groups/memory_hook.py new file mode 100644 index 000000000..876a77426 --- /dev/null +++ b/python/sglang/srt/arg_groups/memory_hook.py @@ -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." + ) diff --git a/python/sglang/srt/arg_groups/model_hook.py b/python/sglang/srt/arg_groups/model_hook.py new file mode 100644 index 000000000..f20f7babd --- /dev/null +++ b/python/sglang/srt/arg_groups/model_hook.py @@ -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 + # 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)}." + ) diff --git a/python/sglang/srt/arg_groups/model_path_hook.py b/python/sglang/srt/arg_groups/model_path_hook.py new file mode 100644 index 000000000..5ed2efb60 --- /dev/null +++ b/python/sglang/srt/arg_groups/model_path_hook.py @@ -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 diff --git a/python/sglang/srt/arg_groups/moe_hook.py b/python/sglang/srt/arg_groups/moe_hook.py new file mode 100644 index 000000000..2af9ed1ac --- /dev/null +++ b/python/sglang/srt/arg_groups/moe_hook.py @@ -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}." + ) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 5500381d5..1dd9274bd 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -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 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 - would grow an entry nothing reads. So, like ``declare_late_resolution``, - this refuses the published record -- post-publish changes go to the bags - through ``get_context().override(...)``. + would grow an entry nothing reads. + + 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 - 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))) if not isinstance(declared, dict): raise TypeError( @@ -251,6 +246,23 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None: f"got {type(declared).__name__}" ) 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)) stash = getattr(server_args, "_resolved_overrides", None) if stash is None: diff --git a/python/sglang/srt/arg_groups/parallel_hook.py b/python/sglang/srt/arg_groups/parallel_hook.py new file mode 100644 index 000000000..e245ae809 --- /dev/null +++ b/python/sglang/srt/arg_groups/parallel_hook.py @@ -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, + ) diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index 7d45061ba..e7e78ac2f 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -3,7 +3,7 @@ from __future__ import annotations import dataclasses import logging import os -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from sglang.srt.arg_groups.overrides import ( declare_resolution, @@ -182,3 +182,81 @@ def _alias_bootstrap_port_to_api_port(server_args: ServerArgs) -> None: "_alias_bootstrap_port_to_api_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." + ) diff --git a/python/sglang/srt/arg_groups/platform_hook.py b/python/sglang/srt/arg_groups/platform_hook.py new file mode 100644 index 000000000..d0c5a607a --- /dev/null +++ b/python/sglang/srt/arg_groups/platform_hook.py @@ -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", + ) diff --git a/python/sglang/srt/arg_groups/serving_hook.py b/python/sglang/srt/arg_groups/serving_hook.py new file mode 100644 index 000000000..5afaff2d5 --- /dev/null +++ b/python/sglang/srt/arg_groups/serving_hook.py @@ -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" + ) diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py new file mode 100644 index 000000000..881f9ce83 --- /dev/null +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -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', , , ]" + 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)." + ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 124aec7ed..1e5d4193c 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -38,14 +38,10 @@ import copy import dataclasses import functools import glob -import importlib -import importlib.util import json import logging import math import os -import random -import socket import tempfile import uuid from typing import Any, Callable, Dict, List, Literal, Optional, Union @@ -68,22 +64,14 @@ from sglang.srt.arg_groups.overrides import ( 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.distributed.device_communicators.mooncake_transfer_engine import ( - parse_ib_device_config, -) from sglang.srt.environ import envs from sglang.srt.function_call.function_call_parser import FunctionCallParser from sglang.srt.hardware_backend.mlx.runtime import use_mlx from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.model_executor.cuda_graph_config import ( - ALLOWED_BACKENDS_PER_PHASE, Backend, CudaGraphConfig, Phase, - default_cuda_graph_config, parse_cuda_graph_config_arg, with_phase, ) @@ -93,39 +81,18 @@ from sglang.srt.speculative.decoupled_spec_io import DecoupledSpecIpcConfig from sglang.srt.utils.common import ( LORA_TARGET_ALL_MODULES, SUPPORTED_LORA_TARGET_MODULES, - configure_media_url_security, - get_device, get_device_memory_capacity, - get_device_sm, - get_quantization_config, human_readable_int, - is_blackwell_supported, - is_cpu, - is_cuda, is_flashinfer_available, is_hip, is_hopper_with_cuda_12_3, - is_host_cpu_arm64, - is_mnnvl_fabric_device, is_mps, - is_musa, is_no_spec_infer_or_topk_one, - is_npu, - is_remote_url, - is_sm90_supported, - is_sm100_or_sm110_supported, is_sm100_supported, - is_sm120_supported, - is_xpu, json_list_type, nullable_str, - parse_connector_type, - torch_release, ) -from sglang.srt.utils.hf_transformers_utils import check_gguf_file from sglang.srt.utils.network import NetworkAddress, get_free_port, wait_port_available -from sglang.srt.utils.runai_utils import ObjectStorageModel, is_runai_obj_uri -from sglang.utils import is_in_ci logger = logging.getLogger(__name__) @@ -4019,246 +3986,19 @@ class ServerArgs: self._resolution_finished = True def _handle_return_hidden_states_mode(self): - cfg = resolving_view(self) - 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: - self._declare( - "_handle_return_hidden_states_mode", - return_hidden_states_mode="full", - ) - else: - self._declare( - "_handle_return_hidden_states_mode", - enable_return_hidden_states=True, - ) + from sglang.srt.arg_groups.serving_hook import handle_return_hidden_states_mode + + handle_return_hidden_states_mode(self) def _handle_model_capability_adjustments(self): - cfg = resolving_view(self) - 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, - ) + from sglang.srt.arg_groups.model_hook import handle_model_capability_adjustments - model_config = self.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(self, _hrm_text_attention_force) - self._declare( - "_handle_model_capability_adjustments", - chunked_prefill_size=-1, - ) - self._declare( - "_handle_model_capability_adjustments", - disable_radix_cache=True, - ) - self._declare( - "_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. - self._declare( - "_handle_model_capability_adjustments", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_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 - ): - self._declare( - "_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. - self._declare( - "_handle_model_capability_adjustments", - is_embedding=True, - ) - self._declare( - "_handle_model_capability_adjustments", - disable_radix_cache=True, - ) - self._declare( - "_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. - self._declare( - "_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. - self._declare( - "_handle_model_capability_adjustments", - prefill_only_disable_kv_cache=True, - ) - self._validate_prefill_only_disable_kv_cache_args() - self._declare( - "_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: - self._declare( - "_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( - self, "_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"] = self._generate_prefill_cuda_graph_batch_sizes( - sizing["max_bs"] - ) - self._declare( - "_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. - self._declare( - "_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 - ): - self._declare( - "_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}" - ) + handle_model_capability_adjustments(self) def _handle_model_source_paths(self): - """Prepare metadata for model paths backed by remote object stores.""" - cfg = resolving_view(self) - self._resolve_hf_gguf_model_path() + from sglang.srt.arg_groups.model_path_hook import handle_model_source_paths - 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) + handle_model_source_paths(self) def _handle_pd_disaggregation(self): from sglang.srt.arg_groups.pd_disaggregation_hook import ( @@ -4268,487 +4008,59 @@ class ServerArgs: handle_pd_disaggregation(self) def _handle_dcp_validation(self): - cfg = resolving_view(self) - 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}." - ) + from sglang.srt.arg_groups.parallel_hook import handle_dcp_validation + + handle_dcp_validation(self) def _handle_load_balance_method(self): - cfg = resolving_view(self) - if cfg.disaggregation_mode not in ("null", "prefill", "decode"): - raise ValueError(f"Invalid disaggregation_mode={cfg.disaggregation_mode!r}") + from sglang.srt.arg_groups.serving_hook import handle_load_balance_method - if cfg.load_balance_method == "auto": - # Default behavior: - # - non-PD: round_robin - # - PD prefill: follow_bootstrap_room - # - PD decode: round_robin - self._declare( - "_handle_load_balance_method", - load_balance_method=( - "follow_bootstrap_room" - if cfg.disaggregation_mode == "prefill" - else "round_robin" - ), - ) - return + handle_load_balance_method(self) def _handle_ssl_validation(self): - """Ensure SSL arguments are consistent and referenced files exist.""" - cfg = resolving_view(self) - 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." - ) + from sglang.srt.arg_groups.serving_hook import handle_ssl_validation - 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." - ) + handle_ssl_validation(self) def _handle_multimodal(self): - """Validate mm_process_config structure before model loading.""" - cfg = resolving_view(self) - 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])}" - ) + from sglang.srt.arg_groups.serving_hook import handle_multimodal + + handle_multimodal(self) def _handle_media_url_security(self): - """Normalize and publish the media URL policy before workers start.""" - cfg = resolving_view(self) - self._declare( - "_handle_media_url_security", - allowed_media_domains=configure_media_url_security( - cfg.allowed_media_domains, - cfg.media_url_max_file_size_mb, - ), - ) + from sglang.srt.arg_groups.serving_hook import handle_media_url_security + + handle_media_url_security(self) def _handle_deprecated_args(self): - cfg = resolving_view(self) - 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." - ) - self._declare("_handle_deprecated_args", image_processor_backend="pil") + from sglang.srt.arg_groups.serving_hook import handle_deprecated_args - # 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." - ) - self._declare( - "_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." - ) - self._declare( - "_handle_deprecated_args", - flashinfer_allreduce_fusion_backend="auto", - ) - self._declare( - "_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(self, attr, None) == "compressed": - logger.warning( - "--%s=compressed is deprecated; use 'dsv4' instead.", - attr.replace("_", "-"), - ) - renamed[attr] = "dsv4" - if renamed: - self._declare("_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." - ) - self._declare( - "_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. - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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." - ) + handle_deprecated_args(self) def _handle_prefill_delayer_env_compat(self): - if envs.SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE.get(): - self._declare( - "_handle_prefill_delayer_env_compat", - enable_prefill_delayer=True, - ) - if x := envs.SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES.get(): - self._declare( - "_handle_prefill_delayer_env_compat", - prefill_delayer_max_delay_passes=x, - ) - if x := envs.SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK.get(): - self._declare( - "_handle_prefill_delayer_env_compat", - prefill_delayer_token_usage_low_watermark=x, - ) + from sglang.srt.arg_groups.serving_hook import handle_prefill_delayer_env_compat + + handle_prefill_delayer_env_compat(self) def _handle_missing_default_values(self): - cfg = resolving_view(self) - if cfg.tokenizer_path is None: - self._declare( - "_handle_missing_default_values", - tokenizer_path=cfg.model_path, - ) - if cfg.served_model_name is None: - self._declare( - "_handle_missing_default_values", - served_model_name=cfg.model_path, - ) - if cfg.device is None: - self._declare( - "_handle_missing_default_values", - device=get_device(), - ) - # strip device index from user if any (e.g. "cuda:0" -> "cuda") - self._declare( - "_handle_missing_default_values", - device=cfg.device.split(":")[0], - ) - if cfg.random_seed is None: - self._declare( - "_handle_missing_default_values", - random_seed=random.randint(0, 1 << 30), - ) - if cfg.mm_process_config is None: - self._declare("_handle_missing_default_values", mm_process_config={}) + from sglang.srt.arg_groups.serving_hook import handle_missing_default_values - # Handle ModelScope model downloads - if envs.SGLANG_USE_MODELSCOPE.get(): - self._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: - self._declare( - "_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: - self._declare( - "_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": - self._declare( - "_handle_missing_default_values", - quantization=None, - ) - self._quantization_explicitly_unset = True - else: - self._quantization_explicitly_unset = False - if cfg.speculative_draft_model_quantization == "unquant": - self._declare( - "_handle_missing_default_values", - speculative_draft_model_quantization=None, - ) + handle_missing_default_values(self) def _handle_modelscope_paths(self): - """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. + from sglang.srt.arg_groups.model_path_hook import handle_modelscope_paths - 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(self) - - 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 {}), - ) - - self._declare( - "_handle_modelscope_paths", - model_path=_resolve_or_download(cfg.model_path, revision=cfg.revision), - ) - self._declare( - "_handle_modelscope_paths", - tokenizer_path=_resolve_or_download( - cfg.tokenizer_path, - ignore_patterns=["*.bin", "*.safetensors"], - revision=cfg.revision, - ), - ) - if cfg.speculative_draft_model_path: - self._declare( - "_handle_modelscope_paths", - speculative_draft_model_path=_resolve_or_download( - cfg.speculative_draft_model_path, - revision=cfg.speculative_draft_model_revision or "main", - ), - ) + handle_modelscope_paths(self) def _handle_hpu_backends(self): - cfg = resolving_view(self) - if cfg.device == "hpu": - self._declare( - "_handle_hpu_backends", - attention_backend="torch_native", - ) - self._declare( - "_handle_hpu_backends", - sampling_backend="pytorch", - ) + from sglang.srt.arg_groups.platform_hook import handle_hpu_backends + + handle_hpu_backends(self) def _handle_cpu_backends(self): - cfg = resolving_view(self) - if cfg.device == "cpu": - if cfg.attention_backend is None: - self._declare( - "_handle_cpu_backends", - attention_backend=( - "torch_native" if is_host_cpu_arm64() else "intel_amx" - ), - ) - self._declare( - "_handle_cpu_backends", - sampling_backend="pytorch", - ) + from sglang.srt.arg_groups.platform_hook import handle_cpu_backends + + handle_cpu_backends(self) def _handle_hardware_runtime_validation(self): # This is intentionally independent of self.device: setting @@ -4758,259 +4070,56 @@ class ServerArgs: use_mlx() def _handle_npu_backends(self): - cfg = resolving_view(self) - if cfg.device == "npu": - from sglang.srt.hardware_backend.npu.utils import set_default_server_args + from sglang.srt.arg_groups.platform_hook import handle_npu_backends - set_default_server_args(self) - - 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'." - ) - self._declare( - "_handle_npu_backends", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, tc_compiler="eager" - ), - ) + handle_npu_backends(self) def _handle_mps_backends(self): - cfg = resolving_view(self) - if cfg.device == "mps": - if not use_mlx(): - self._declare( - "_handle_mps_backends", - disable_overlap_schedule=True, - ) + from sglang.srt.arg_groups.platform_hook import handle_mps_backends + + handle_mps_backends(self) def _handle_xpu_backends(self): - cfg = resolving_view(self) - 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 self._cuda_graph_config_locked: - self._declare( - "_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, - ) - self._declare( - "_handle_xpu_backends", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) + from sglang.srt.arg_groups.platform_hook import handle_xpu_backends + + handle_xpu_backends(self) # ------------------------------------------------------------------ # CUDA graph configuration resolution # ------------------------------------------------------------------ def _apply_inkling_prefill_cuda_graph_default(self): - """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(self) - 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 = self.get_model_config().hf_config.architectures[0] - if arch in ( - "InklingForConditionalGeneration", - "InklingForConditionalGenerationMTP", - ): - self._declare( - "_apply_inkling_prefill_cuda_graph_default", - cuda_graph_backend_prefill=Backend.FULL, - ) + from sglang.srt.arg_groups.cuda_graph_hook import ( + apply_inkling_prefill_cuda_graph_default, + ) + + apply_inkling_prefill_cuda_graph_default(self) def _apply_muse_glimmer_prefill_cuda_graph_max_bs_default(self): - cfg = resolving_view(self) - if ( - cfg.cuda_graph_max_bs_prefill is not None - or parse_connector_type(cfg.model_path) == ConnectorType.INSTANCE - ): - return - arch = self.get_model_config().hf_config.architectures[0] - if arch in ("MuseGlimmerForCausalLM", "MuseGlimmerForConditionalGeneration"): - self._declare( - "_apply_muse_glimmer_prefill_cuda_graph_max_bs_default", - cuda_graph_max_bs_prefill=512, - ) + from sglang.srt.arg_groups.cuda_graph_hook import ( + apply_muse_glimmer_prefill_cuda_graph_max_bs_default, + ) + + apply_muse_glimmer_prefill_cuda_graph_max_bs_default(self) def _handle_cuda_graph_config(self): - cfg = resolving_view(self) + from sglang.srt.arg_groups.cuda_graph_hook import handle_cuda_graph_config - self._parse_cuda_graph_config() - self._apply_cuda_graph_compatibility() - self._apply_deepep_adjustments() - self._apply_cuda_graph_disaggregation_roles() - self._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." - ) + handle_cuda_graph_config(self) def _apply_deepep_adjustments(self): - """Config adjustments required by the DeepEP a2a backend.""" - cfg = resolving_view(self) - if resolved_view(self).moe_a2a_backend != "deepep": - return + from sglang.srt.arg_groups.cuda_graph_hook import apply_deepep_adjustments - # 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 = self._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, - ) - self._declare( - "_apply_deepep_adjustments", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, - Phase.PREFILL, - bs=aligned, - max_bs=aligned[-1], - ), - ) + apply_deepep_adjustments(self) def _parse_cuda_graph_config(self): - """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 self._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(self) - 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() + from sglang.srt.arg_groups.cuda_graph_hook import parse_cuda_graph_config - 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) - - self._declare( - "_parse_cuda_graph_config", - cuda_graph_config=config, - ) - self._cuda_graph_config_locked = locked + parse_cuda_graph_config(self) def _apply_cuda_graph_compatibility(self): - """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(self) - if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked: - return + from sglang.srt.arg_groups.cuda_graph_hook import apply_cuda_graph_compatibility - # 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 self.get_model_config().is_multimodal_piecewise_cuda_graph_supported - and not self.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 self._resolved_attention_backends()[0] != "trtllm_mla" - ): - logger.info( - "Using tc_piecewise CUDA graph for validated multimodal " - "decoder prefill." - ) - self._declare( - "_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: - self._disable_tc_piecewise_cudagraph_if_incompatible() - elif cfg.cuda_graph_config.prefill.backend == Backend.BREAKABLE: - self._disable_breakable_cudagraph_if_incompatible() - elif cfg.cuda_graph_config.prefill.backend == Backend.FULL: - self._disable_full_prefill_cudagraph_if_incompatible() + apply_cuda_graph_compatibility(self) def _apply_cuda_graph_disaggregation_roles(self): cfg = resolving_view(self) @@ -5032,499 +4141,47 @@ class ServerArgs: ) def _disable_tc_piecewise_cudagraph_if_incompatible(self): - """TcPiecewise (torch.compile + piecewise) is incompatible with - these configurations. Most are torch.compile / dynamo limitations. - """ - cfg = resolving_view(self) + from sglang.srt.arg_groups.cuda_graph_hook import ( + disable_tc_piecewise_cudagraph_if_incompatible, + ) - rules = [ - ( - "model-arch blacklist", - lambda: self.get_model_config().is_piecewise_cuda_graph_disabled_model, - ), - ("DP attention", lambda: self._resolved().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(self).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: self.get_model_config().is_multimodal - and not self.get_model_config().is_multimodal_piecewise_cuda_graph_supported, - ), - ( - "GGUF quantization", - lambda: cfg.load_format == "gguf" - or resolved_view(self).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: self._resolved().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(): - self._declare( - "_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 + disable_tc_piecewise_cudagraph_if_incompatible(self) def _disable_breakable_cudagraph_if_incompatible(self): - """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(self) - from sglang.srt.configs.model_config import is_deepseek_v4 - from sglang.srt.layers.cp.bcg import supports_prefill_cp_bcg + from sglang.srt.arg_groups.cuda_graph_hook import ( + disable_breakable_cudagraph_if_incompatible, + ) - 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(self.get_model_config().hf_config), - ), - # CP all_gather replay size mismatch under BCG. - ( - "context parallel (attn_cp_size > 1)", - lambda: self._resolved().attn_cp_size > 1 - and not supports_prefill_cp_bcg(self), - ), - # 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(self).moe_a2a_backend - not in ("none", "deepep", "megamoe", "flashinfer"), - ), - # Multimodal prefill replay faults under BCG; allowlisted archs opt back in. - ( - "multimodal model", - lambda: self.get_model_config().is_multimodal - and not self.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, - ) - self._declare( - "_disable_breakable_cudagraph_if_incompatible", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - return + disable_breakable_cudagraph_if_incompatible(self) def _disable_full_prefill_cudagraph_if_incompatible(self): - """Full prefill CG: empty rule list today; see the experimental warning.""" - cfg = resolving_view(self) - rules = [] - for name, predicate in rules: - if predicate(): - logger.warning( - "Full prefill CUDA graph is incompatible with %s; " - "disabling prefill CUDA graph.", - name, - ) - self._declare( - "_disable_full_prefill_cudagraph_if_incompatible", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - return + from sglang.srt.arg_groups.cuda_graph_hook import ( + disable_full_prefill_cudagraph_if_incompatible, + ) + + disable_full_prefill_cudagraph_if_incompatible(self) def _disable_prefill_cuda_graph_for_deepseek_trtllm_mla(self): - """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(self) + from sglang.srt.arg_groups.cuda_graph_hook import ( + disable_prefill_cuda_graph_for_deepseek_trtllm_mla, + ) - if (Phase.PREFILL, "backend") in self._cuda_graph_config_locked: - return - if cfg.cuda_graph_config.prefill.backend == Backend.DISABLED: - return - if ( - "DeepseekV3ForCausalLM" - not in self.get_model_config().hf_config.architectures - ): - return - prefill_attention_backend, _ = self._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, - ) - self._declare( - "_disable_prefill_cuda_graph_for_deepseek_trtllm_mla", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) + disable_prefill_cuda_graph_for_deepseek_trtllm_mla(self) def _validate_cuda_graph_config(self): - cfg = resolving_view(self) - 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]}" - ) + from sglang.srt.arg_groups.cuda_graph_hook import validate_cuda_graph_config + + validate_cuda_graph_config(self) def _handle_multi_item_scoring(self): - """Setup and validate multi-item scoring constraints. + from sglang.srt.arg_groups.attention_hook import handle_multi_item_scoring - 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(self) - 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.") - self._declare( - "_handle_multi_item_scoring", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_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.") - self._declare( - "_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.") - self._declare( - "_handle_multi_item_scoring", - chunked_prefill_size=-1, - ) - - prefill_backend, decode_backend = self._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}" - ) + handle_multi_item_scoring(self) def _handle_gpu_memory_settings(self, gpu_mem): - """ - Configure GPU memory-dependent settings including - chunked_prefill_size, cuda_graph_config[decode].max_bs, and mem_fraction_static. + from sglang.srt.arg_groups.memory_hook import handle_gpu_memory_settings - 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(self) - # 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: - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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: - self._declare( - "_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 = ( - self._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: - self._declare( - "_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 - self._declare( - "_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 = self._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 self.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 = ( - self._generate_prefill_cuda_graph_batch_sizes( - prefill_cuda_graph_config.max_bs - ) - ) - - if cuda_graph_config != cfg.cuda_graph_config: - self._declare( - "_handle_gpu_memory_settings", cuda_graph_config=cuda_graph_config - ) - - if cfg.mem_fraction_static is None: - if self.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 += self.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 += self.reserve_for_deepep_a2a_mb() - - self._declare( - "_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 = self.get_model_config() - if ( - model_config.is_multimodal - and not cfg.language_only - and not cfg.language_model_only - and cfg.disaggregation_mode != "decode" - ): - self.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." - ) + handle_gpu_memory_settings(self, gpu_mem) def post_capture_kv_sizing_planned(self) -> bool: """Whether the mem_fraction heuristic may skip the graph reserve; must be @@ -5762,559 +4419,9 @@ class ServerArgs: validate_hisparse_kv_cache_dtype(self) def _handle_model_specific_adjustments(self): - cfg = resolving_view(self) - from sglang.srt.configs.model_config import ( - get_mimo_v2_fused_qkv_expected_tp_size, - is_deepseek_dsa, - ) + from sglang.srt.arg_groups.model_hook import handle_model_specific_adjustments - if cfg.enable_deterministic_inference: - self._declare( - "_handle_model_specific_adjustments", - enforce_disable_flashinfer_allreduce_fusion=True, - ) - - self._declare( - "_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 = self.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: - self._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, self, hf_config - ) - validate_declarations(self, model_overrides) - self._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(self) - apply_kimi_k3_spec_backend_defaults(self) - - if model_arch in [ - "DeepseekV4ForCausalLM", - ]: - from sglang.srt.arg_groups.deepseek_v4_hook import ( - apply_deepseek_v4_defaults, - ) - - apply_deepseek_v4_defaults(self, 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). - self._declare( - "_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() - self._set_default_dsa_kv_cache_dtype( - major, resolved_view(self).quantization - ) - self._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 self.use_mla_backend(): - self._declare( - "_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(self, _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 - # 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 self._resolved().enable_dp_attention and cfg.nnodes == 1: - # TODO (Hubert): Put this back later - # self.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(self, _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(self) - validate_deepseek_v4_mega_moe_token_budget(self) - - 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 = ( - self._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 self._resolved().enable_dp_attention - and cfg.nnodes == 1 - and is_hip() - ): - # TODO (Hubert): Put this back later - # self.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(self).moe_runner_backend == "triton_kernel": - assert ( - self._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 = self._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(self).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 = self._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(self).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(self).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(self).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. - self._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(self, _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(self, _flashinfer_allreduce_fusion_auto_enable) - run_post_process_pass(self, _enforce_disable_allreduce_fusion) + handle_model_specific_adjustments(self) def _support_mamba_cache_extra_buffer(self, model_arch: str): from sglang.srt.arg_groups.overrides import supports_mamba_cache_extra_buffer @@ -6322,70 +4429,26 @@ class ServerArgs: return supports_mamba_cache_extra_buffer(self, model_arch) def _validate_mamba_no_buffer(self, 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." + from sglang.srt.arg_groups.mamba_hook import validate_mamba_no_buffer + + validate_mamba_no_buffer(view, model_arch) def _validate_mamba_extra_buffer(self, view, model_arch: str): - from sglang.srt.arg_groups.overrides import supports_mamba_cache_extra_buffer + from sglang.srt.arg_groups.mamba_hook import validate_mamba_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 - assert self.mamba_cache_chunk_size is not None - - if ( - view.chunked_prefill_size is not None - and 0 < view.chunked_prefill_size < self.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, - self.mamba_cache_chunk_size, - ) + validate_mamba_extra_buffer( + view, + model_arch, + mamba_cache_chunk_size_of=lambda: self.mamba_cache_chunk_size, + ) def _handle_mamba_radix_cache(self, 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, - ) + from sglang.srt.arg_groups.model_hook import handle_mamba_radix_cache - run_post_process_pass(self, _mamba_radix_cache_resolution) - view = resolved_view(self) - if not view.uses_mamba_radix_cache: - return - - if mamba_extra_buffer_of(view): - self._validate_mamba_extra_buffer(view, model_arch) - else: - self._validate_mamba_no_buffer(view, model_arch) + handle_mamba_radix_cache(self, model_arch) def _handle_sampling_backend(self): # Moved to the resolution pipeline (arg_groups/overrides.py: @@ -6473,263 +4536,23 @@ class ServerArgs: return "triton" def _handle_attention_backend_compatibility(self): - cfg = resolving_view(self) - model_config = self.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, + from sglang.srt.arg_groups.attention_hook import ( + handle_attention_backend_compatibility, ) - # Split-backend override + default fill. - run_post_process_pass(self, _attention_backend_default) - - # Torch native and flex attention backends - attention_backend = resolved_view(self).attention_backend - if attention_backend == "torch_native": - logger.warning( - "Cuda graph is disabled because of using torch native attention backend" - ) - self._declare( - "_handle_attention_backend_compatibility", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_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" - ) - self._declare( - "_handle_attention_backend_compatibility", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_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") - self._declare( - "_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(self, _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(self, _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(self, _cutedsl_prefill_backend_fill) - - prefill_backend, decode_backend = self._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(self, _attention_backend_fa3_fp8_fallback) - - run_post_process_pass(self, _fa4_page_constraint) - - # AMD platforms backends - if resolved_view(self).attention_backend == "aiter": - if model_config.context_len > 8192: - self._declare( - "_handle_attention_backend_compatibility", - mem_fraction_static=cfg.mem_fraction_static * 0.85, - ) - - # Other platforms backends - run_post_process_pass(self, _attention_backend_platform_fallbacks) - - prefill_backend, decode_backend = self._resolved_attention_backends() - if self.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(self, _intel_xpu_page_constraint) - - # Dual chunk flash attention backend - run_post_process_pass(self, _attention_backend_dual_chunk) - if resolved_view(self).attention_backend == "dual_chunk_flash_attn": - logger.warning( - "Mixed chunk and radix cache are disabled when using dual-chunk flash attention backend" - ) - self._declare( - "_handle_attention_backend_compatibility", - enable_mixed_chunk=False, - ) - self._declare( - "_handle_attention_backend_compatibility", - disable_radix_cache=True, - ) + handle_attention_backend_compatibility(self) def _handle_mxfp8_kv_cache_compatibility(self): - """MXFP8 KV cache uses operands available only on SM100+ (Blackwell).""" - cfg = resolving_view(self) - 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." - ) + from sglang.srt.arg_groups.kv_cache_hook import ( + handle_mxfp8_kv_cache_compatibility, + ) + + handle_mxfp8_kv_cache_compatibility(self) def _handle_kv4_compatibility(self): - """Check FP4 KV cache compatibility with the attention backend""" - cfg = resolving_view(self) + from sglang.srt.arg_groups.kv_cache_hook import handle_kv4_compatibility - if cfg.kv_cache_dtype not in ("nvfp4", "fp4_mx_block16"): - return - - use_mla_backend = self.use_mla_backend() - prefill_backend, decode_backend = self._resolved_attention_backends() - attention_backend = resolved_view(self).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.") + handle_kv4_compatibility(self) def _handle_page_size(self): # Moved to the resolution pipeline (arg_groups/overrides.py: @@ -6742,747 +4565,59 @@ class ServerArgs: run_post_process_pass(self, _page_size_default) def _handle_amd_specifics(self): - if is_hip(): - self._declare("_handle_amd_specifics", triton_attention_num_kv_splits=16) + from sglang.srt.arg_groups.platform_hook import handle_amd_specifics + + handle_amd_specifics(self) def _handle_nccl_pre_warm(self): - # pre_warm_nccl is only used with CUDA or HIP hardware or NPU hardware - cfg = resolving_view(self) - 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." - ) - self._declare("_handle_nccl_pre_warm", pre_warm_nccl=False) + from sglang.srt.arg_groups.platform_hook import handle_nccl_pre_warm + + handle_nccl_pre_warm(self) def _handle_grammar_backend(self): - cfg = resolving_view(self) - if cfg.grammar_backend is None: - self._declare("_handle_grammar_backend", grammar_backend="xgrammar") + from sglang.srt.arg_groups.serving_hook import handle_grammar_backend + + handle_grammar_backend(self) def _handle_mamba_backend(self): - cfg = resolving_view(self) - if cfg.mamba_cache_philox_rounds < 0: - raise ValueError("--mamba-cache-philox-rounds must be non-negative.") + from sglang.srt.arg_groups.mamba_hook import handle_mamba_backend - 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) + handle_mamba_backend(self) def _handle_int8_mamba_checkpoint(self): - # 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(self) - 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." - ) + from sglang.srt.arg_groups.mamba_hook import handle_int8_mamba_checkpoint + + handle_int8_mamba_checkpoint(self) def _handle_linear_attn_backend(self): - cfg = resolving_view(self) - import torch + from sglang.srt.arg_groups.attention_hook import handle_linear_attn_backend - # 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). - ): - self._declare( - "_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)." - ) - self._declare( - "_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(self)): - 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)." - ) - self._declare( - "_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, - ) + handle_linear_attn_backend(self) def _handle_legacy_cp_arguments(self): - cfg = resolving_view(self) - legacy_mode_to_strategy = { - "in-seq-split": "zigzag", - "round-robin-split": "interleave", - } - strategy_to_legacy_mode = { - "zigzag": "in-seq-split", - "interleave": "round-robin-split", - } + from sglang.srt.arg_groups.parallel_hook import handle_legacy_cp_arguments - if ( - cfg.enable_prefill_context_parallel - or cfg.enable_dsa_prefill_context_parallel - ): - self._declare( - "_handle_legacy_cp_arguments", - enable_prefill_cp=True, - ) - - if cfg.enable_prefill_context_parallel and cfg.cp_strategy is None: - self._declare( - "_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: - self._declare( - "_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( - self._resolved(), "attention_backend", None - ) in ("dsa", "dsv4") - if use_dsa_legacy_aliases: - self._declare( - "_handle_legacy_cp_arguments", - enable_dsa_prefill_context_parallel=True, - ) - self._declare( - "_handle_legacy_cp_arguments", - enable_prefill_context_parallel=False, - ) - else: - self._declare( - "_handle_legacy_cp_arguments", - enable_prefill_context_parallel=True, - ) - self._declare( - "_handle_legacy_cp_arguments", - dsa_prefill_cp_mode=mode, - ) - self._declare( - "_handle_legacy_cp_arguments", - prefill_cp_mode=mode, - ) + handle_legacy_cp_arguments(self) def _handle_context_parallelism(self): - cfg = resolving_view(self) - 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 + from sglang.srt.arg_groups.parallel_hook import handle_context_parallelism - model_config = self.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 = self._resolved() - 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, - ) + handle_context_parallelism(self) def _handle_dwdp(self): - cfg = resolving_view(self) - if cfg.dwdp_size <= 1: - return + from sglang.srt.arg_groups.parallel_hook import handle_dwdp - 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." - ) - - self._declare( - "_handle_dwdp", - dp_size=cfg.dwdp_size, - ) - self._declare( - "_handle_dwdp", - enable_dp_attention=True, - ) - self._declare("_handle_dwdp", enable_dp_attention_local_control_broadcast=True) - self._declare( - "_handle_dwdp", - enable_dp_lm_head=True, - ) - self._declare( - "_handle_dwdp", - moe_dense_tp_size=1, - ) - self._declare( - "_handle_dwdp", - ep_size=cfg.dwdp_size, - ) - self._declare( - "_handle_dwdp", - moe_dp_size=1, - ) - self._declare( - "_handle_dwdp", - moe_a2a_backend="none", - ) - - envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.set(True) - - self._declare( - "_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" - ) + handle_dwdp(self) def _handle_data_parallelism(self): - # The dp_size==1 resets moved to the resolution pipeline - # (arg_groups/overrides.py: _data_parallelism_defaults). - cfg = resolving_view(self) - from sglang.srt.arg_groups.overrides import ( - _data_parallelism_defaults, - run_post_process_pass, - ) + from sglang.srt.arg_groups.parallel_hook import handle_data_parallelism - run_post_process_pass(self, _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 self._resolved().enable_dp_attention: - self._declare( - "_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 - self._declare( - "_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 self._cuda_graph_config_locked - ): - clamped = {"max_bs": cfg.chunked_prefill_size} - if (Phase.PREFILL, "bs") not in self._cuda_graph_config_locked: - clamped["bs"] = self._generate_prefill_cuda_graph_batch_sizes( - clamped["max_bs"] - ) - self._declare( - "_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(self, _tp_lm_head_all_to_all_default) - run_post_process_pass(self, _dp_lm_head_validation) + handle_data_parallelism(self) def _handle_moe_kernel_config(self): - # 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(self) - from sglang.srt.arg_groups.overrides import ( - _moe_runner_backend_quant_constraints, - _moe_runner_fusion_disable, - run_post_process_pass, - ) + from sglang.srt.arg_groups.moe_hook import handle_moe_kernel_config - run_post_process_pass(self, _moe_runner_backend_quant_constraints) - - view = resolved_view(self) - 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 self.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(self, _moe_runner_fusion_disable) - - if resolved_view(self).moe_runner_backend == "cutlass" and resolved_view( - self - ).quantization in [ - "fp8", - "mxfp8", - ]: - assert ( - resolved_view(self).ep_size == 1 - ), "FP8/MXFP8 Cutlass MoE is only supported with ep_size == 1" + handle_moe_kernel_config(self) def cutedsl_moe_max_num_tokens(self) -> int: """Largest number of tokens a single forward routes through a CuteDSL @@ -7520,352 +4655,31 @@ class ServerArgs: return tokens def _validate_cutedsl_a2a_token_budget(self): - """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(self) + from sglang.srt.arg_groups.moe_hook import validate_cutedsl_a2a_token_budget - view = resolved_view(self) - 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 = self.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}." - ) + validate_cutedsl_a2a_token_budget(self) def _validate_deepep_v2_dispatch_token_budget(self) -> None: - """Check the configured prefill and decode-graph buffer bounds.""" - view = resolved_view(self) - 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 = self.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 = ( - self.max_speculative_num_draft_tokens or 1 - if view.speculative_algorithm - else 1 + from sglang.srt.arg_groups.moe_hook import ( + validate_deepep_v2_dispatch_token_budget, ) - 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." - ) + + validate_deepep_v2_dispatch_token_budget(self) def _validate_deepep_v2_model_architecture(self) -> None: - """Allow DeepEP v2 only where its model workflow is validated.""" - if ( - parse_connector_type(resolved_view(self).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." - ) + from sglang.srt.arg_groups.moe_hook import validate_deepep_v2_model_architecture - architectures = ( - getattr(self.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." - ) + validate_deepep_v2_model_architecture(self) def _validate_deepep_v2_speculative_draft(self) -> None: - """Reject an explicit or inherited DeepEP v2 draft backend.""" - view = resolved_view(self) - draft_backend = view.speculative_moe_a2a_backend - if draft_backend is None and view.speculative_algorithm: - from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + from sglang.srt.arg_groups.moe_hook import validate_deepep_v2_speculative_draft - 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." - ) + validate_deepep_v2_speculative_draft(self) def _handle_a2a_moe(self): - # 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(self) - from sglang.srt.arg_groups.overrides import ( - _a2a_backend_overrides, - _a2a_ep_size, - _a2a_fusion_adjustments, - run_post_process_pass, - ) + from sglang.srt.arg_groups.moe_hook import handle_a2a_moe - run_post_process_pass(self, _a2a_backend_overrides) - run_post_process_pass(self, _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(self, _a2a_fusion_adjustments) - - a2a_backend = resolved_view(self).moe_a2a_backend - if cfg.enable_waterfill: - self._declare("_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": - self._declare( - "_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`") - self._declare( - "_handle_a2a_moe", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_handle_a2a_moe", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - - if a2a_backend == "deepep_v2": - self._validate_deepep_v2_model_architecture() - if resolved_view(self).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(self).moe_runner_backend - if resolved_runner == "auto": - self._declare("_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. - self._declare( - "_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(self).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 - self._declare( - "_handle_a2a_moe", - moe_a2a_backend="none", - ) - - if cfg.moe_a2a_backend == "flashinfer": - assert ( - resolved_view(self).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(self).quantization == "modelopt_fp4" - or self.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(self).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": - self._declare( - "_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 ( - self._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": - self._declare( - "_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(self).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(self).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": - self._declare( - "_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 ( - self._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)" - ) + handle_a2a_moe(self) def _required_mori_dispatch_tokens_per_rank(self) -> int: """Max tokens a single rank dispatches through MoRI in one forward.""" @@ -7881,269 +4695,29 @@ class ServerArgs: return required def _handle_eplb_and_dispatch(self): - cfg = resolving_view(self) - if cfg.enable_eplb and (cfg.expert_distribution_recorder_mode is None): - self._declare( - "_handle_eplb_and_dispatch", - expert_distribution_recorder_mode="stat", - ) - logger.warning( - "EPLB is enabled. The expert_distribution_recorder_mode is automatically set." - ) + from sglang.srt.arg_groups.parallel_hook import handle_eplb_and_dispatch - # 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 = self._resolved().moe_a2a_backend == "none" - - if (cfg.enable_eplb or (cfg.init_expert_location != "trivial")) and ( - cfg.ep_dispatch_algorithm is None - ): - self._declare( - "_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 self._resolved().ep_size > 1 + handle_eplb_and_dispatch(self) def _handle_elastic_ep(self): - cfg = resolving_view(self) - 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." - ) - self._declare( - "_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": - self._declare( - "_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)'." + from sglang.srt.arg_groups.parallel_hook import handle_elastic_ep - assert cfg.pp_size == 1, "PP size should be set to 1 under elastic EP" - - if cfg.elastic_ep_backend == "mooncake": - self._declare( - "_handle_elastic_ep", - mooncake_ib_device=self._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 = self._resolved() - 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." - ) - self._declare( - "_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: - self._declare( - "_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})." - ) + handle_elastic_ep(self) def _validate_experimental_sgl_marlin(self): - view = self._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, + from sglang.srt.arg_groups.validation_hook import ( + validate_experimental_sgl_marlin, ) - validate_experimental_sgl_marlin_server_args(self, view) + validate_experimental_sgl_marlin(self) # ===== END TO BE REFACTORED ==== def _handle_expert_distribution_metrics(self): - cfg = resolving_view(self) - 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." - ) + from sglang.srt.arg_groups.parallel_hook import ( + handle_expert_distribution_metrics, + ) - if self.should_report_expert_balancedness() and ( - cfg.expert_distribution_recorder_mode is None - ): - self._declare( - "_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: - self._declare( - "_handle_expert_distribution_metrics", - expert_distribution_recorder_buffer_size=x, - ) - elif cfg.expert_distribution_recorder_mode is not None: - self._declare( - "_handle_expert_distribution_metrics", - expert_distribution_recorder_buffer_size=1000, - ) + handle_expert_distribution_metrics(self) def _handle_pipeline_parallelism(self): # Moved to the resolution pipeline (arg_groups/overrides.py: @@ -8156,331 +4730,55 @@ class ServerArgs: run_post_process_pass(self, _pipeline_parallel_overlap_disable) def _validate_prefill_only_disable_kv_cache_args(self): - """Validate --prefill-only-disable-kv-cache flag/precondition constraints. + from sglang.srt.arg_groups.kv_cache_hook import ( + validate_prefill_only_disable_kv_cache_args, + ) - Backend resolution is checked separately by - _handle_prefill_only_disable_kv_cache after backends settle. - """ - cfg = resolving_view(self) - 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 self._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." - ) + validate_prefill_only_disable_kv_cache_args(self) def _handle_prefill_only_disable_kv_cache(self): - """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(self) - - if not cfg.prefill_only_disable_kv_cache: - return - - assert resolved_view(self).attention_backend is not None, ( - "_handle_prefill_only_disable_kv_cache must run after " - "_handle_attention_backend_compatibility() so the prefill backend is resolved." + from sglang.srt.arg_groups.kv_cache_hook import ( + handle_prefill_only_disable_kv_cache, ) - prefill_backend, _ = self._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." - ) + handle_prefill_only_disable_kv_cache(self) def _handle_hicache_ratio_default(self): - """Default the host/device ratio per host memory mode. + from sglang.srt.arg_groups.hicache_hook import handle_hicache_ratio_default - 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(self) - if cfg.hicache_ratio is None and cfg.disaggregation_mode != "decode": - self._declare( - "_handle_hicache_ratio_default", - hicache_ratio=( - 1.2 if cfg.hicache_host_memory_mode == "buffer_only" else 2.0 - ), - ) + handle_hicache_ratio_default(self) def _handle_hicache(self): - """Normalize hicache-related knobs into a valid runtime configuration. + from sglang.srt.arg_groups.hicache_hook import handle_hicache - Resolution order: - 1) Layout <-> I/O compatibility for direct conflicts. - 2) Storage <-> layout compatibility (may rewrite layout). - """ - cfg = resolving_view(self) - # 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 - - self._validate_hicache_host_memory_mode() - - # Step 1: Initial layout-io compatibility normalization. - self._resolve_layout_io_compatibility() - - # Step 2: Storage-layout normalization without changing io backend. - self._resolve_storage_layout_compatibility() - - # Step 3: DCP compatibility for the L2 (device<->host) path. - self._resolve_hicache_dcp_compatibility() + handle_hicache(self) def _validate_hicache_host_memory_mode(self): - cfg = resolving_view(self) - 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}" - ) + from sglang.srt.arg_groups.hicache_hook import validate_hicache_host_memory_mode - # 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." - ) + validate_hicache_host_memory_mode(self) def _resolve_hicache_dcp_compatibility(self): - cfg = resolving_view(self) - 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 self.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, - ) + from sglang.srt.arg_groups.hicache_hook import resolve_hicache_dcp_compatibility + + resolve_hicache_dcp_compatibility(self) def _resolve_layout_io_compatibility(self): - cfg = resolving_view(self) - if ( - cfg.hicache_mem_layout == "page_first_direct" - and cfg.hicache_io_backend == "kernel" - ): - self._declare( - "_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" - ) + from sglang.srt.arg_groups.hicache_hook import resolve_layout_io_compatibility - if ( - cfg.hicache_mem_layout == "page_first" - and cfg.hicache_io_backend == "direct" - ): - self._declare( - "_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" - ) + resolve_layout_io_compatibility(self) def _resolve_storage_layout_compatibility(self): - cfg = resolving_view(self) - 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 - - self._declare( - "_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" + from sglang.srt.arg_groups.hicache_hook import ( + resolve_storage_layout_compatibility, ) + resolve_storage_layout_compatibility(self) + def _resolve_hf_gguf_model_path(self): - """Turn a Hub reference to a .gguf into a local file path.""" - cfg = resolving_view(self) - from sglang.srt.utils.hf_transformers_utils import resolve_hf_gguf_reference + from sglang.srt.arg_groups.model_path_hook import resolve_hf_gguf_model_path - 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: - self._declare( - "_resolve_hf_gguf_model_path", - tokenizer_path=resolved, - ) - self._declare( - "_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, - ) - self._declare( - "_resolve_hf_gguf_model_path", - speculative_draft_model_path=resolved_draft, - ) + resolve_hf_gguf_model_path(self) def _handle_expert_pack(self): from sglang.srt.arg_groups.expert_pack_hook import handle_expert_pack @@ -8491,120 +4789,9 @@ class ServerArgs: # 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(self) - from sglang.srt.arg_groups.overrides import ( - _gguf_quantization, - run_post_process_pass, - ) + from sglang.srt.arg_groups.model_path_hook import handle_load_format - run_post_process_pass(self, _gguf_quantization) - if (cfg.load_format == "auto" or cfg.load_format == "gguf") and check_gguf_file( - cfg.model_path - ): - self._declare( - "_handle_load_format", - load_format="gguf", - ) - - if cfg.load_format == "auto" and self._is_mistral_native_format(): - self._declare( - "_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): - self._declare( - "_handle_load_format", - load_format="runai_streamer", - ) - elif is_remote_url(cfg.model_path): - self._declare( - "_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 - ): - self._declare( - "_handle_load_format", - speculative_draft_load_format="runai_streamer", - ) - - if cfg.custom_weight_loader is None: - self._declare("_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." - ) - self._declare( - "_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." - ) - self._declare( - "_handle_load_format", - load_format="auto", - ) - elif ( - cfg.remote_instance_weight_loader_backend == "transfer_engine" - and not self.validate_transfer_engine() - ): - logger.warning( - "Fallback load_format to 'auto' due to 'transfer_engine' backend is not supported." - ) - self._declare( - "_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: - self._declare( - "_handle_load_format", - remote_instance_weight_loader_start_seed_via_transfer_engine=self.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." - ) + handle_load_format(self) def _is_mistral_native_format(self) -> bool: """True iff the checkpoint requires load_format=mistral. @@ -8669,1077 +4856,94 @@ class ServerArgs: LANGUAGE_MODEL_ONLY_ARCHITECTURES = ("MuseGlimmerForConditionalGeneration",) def _handle_language_model_only(self): - cfg = resolving_view(self) - 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 = self.get_model_config().hf_config.architectures - if not any(a in self.LANGUAGE_MODEL_ONLY_ARCHITECTURES for a in architectures): - raise ValueError( - f"--language-model-only does not support {architectures}. " - f"Supported: {list(self.LANGUAGE_MODEL_ONLY_ARCHITECTURES)}." - ) + from sglang.srt.arg_groups.model_hook import handle_language_model_only + + handle_language_model_only(self) def _handle_encoder_disaggregation(self): - cfg = resolving_view(self) - self._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" - ) + from sglang.srt.arg_groups.pd_disaggregation_hook import ( + handle_encoder_disaggregation, + ) - 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": - self._declare( - "_handle_encoder_disaggregation", - disaggregation_ib_device=self._validate_ib_devices( - cfg.disaggregation_ib_device - ), - ) - - # Validate model type for encoder disaggregation - hf_config = self.get_model_config().hf_config - model_arch = hf_config.architectures[0] - if cfg.encoder_transfer_backend == "auto": - self._declare( - "_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." - ) + handle_encoder_disaggregation(self) def _validate_ib_devices(self, device_str: Optional[str]) -> Optional[str]: - """ - Validate IB devices before passing to mooncake. + from sglang.srt.arg_groups.validation_hook import validate_ib_devices - 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=(",", ":")) + return validate_ib_devices(self, device_str) def _handle_tokenizer_batching(self): - cfg = resolving_view(self) - 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." - ) + from sglang.srt.arg_groups.serving_hook import handle_tokenizer_batching - 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})." - ) - self._declare("_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." - ) - self._declare( - "_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." - ) - self._declare( - "_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." - ) + handle_tokenizer_batching(self) def _handle_multimodal_feature_transport(self): - """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(self) - 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 ( - self.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(self.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" - ), - ) - - self._declare( - "_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. - self._declare( - "_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" + from sglang.srt.arg_groups.serving_hook import ( + handle_multimodal_feature_transport, ) + handle_multimodal_feature_transport(self) + def _handle_environment_variables(self): - cfg = resolving_view(self) - self._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." - ) - self._declare("_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 + from sglang.srt.arg_groups.serving_hook import handle_environment_variables - 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) + handle_environment_variables(self) def _handle_cache_compatibility(self): - cfg = resolving_view(self) - 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." - ) + from sglang.srt.arg_groups.kv_cache_hook import handle_cache_compatibility - 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 < self._resolved().swa_full_tokens_ratio <= 1.0): - raise ValueError("--swa-full-tokens-ratio should be in range (0, 1.0].") + handle_cache_compatibility(self) def _handle_deterministic_inference(self): - cfg = resolving_view(self) - if cfg.rl_on_policy_target is not None: - logger.warning( - "Enable deterministic inference because of rl_on_policy_target." - ) - self._declare( - "_handle_deterministic_inference", - enable_deterministic_inference=True, - ) + from sglang.srt.arg_groups.attention_hook import handle_deterministic_inference - # 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." - ) - self._declare( - "_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(self, _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(self, _deterministic_sampling_backend) - is_deepseek_model = False - if parse_connector_type(cfg.model_path) != ConnectorType.INSTANCE: - try: - hf_config = self.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(self, _deterministic_attention_backend) - - attention_backend = resolved_view(self).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 - self._declare( - "_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. - self._declare( - "_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. - self._declare( - "_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." - ) + handle_deterministic_inference(self) def _handle_unified_memory_pool(self): - cfg = resolving_view(self) - 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(self._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)." - ) + from sglang.srt.arg_groups.kv_cache_hook import handle_unified_memory_pool + + handle_unified_memory_pool(self) # The strided-layout Triton requirement is enforced via # --enable-page-major-kv-layout (implied by the unified pool in # _handle_page_major_kv_layout); the model-family gate is enforced at pool # construction in model_runner_kv_cache_mixin._init_pools. def _handle_page_major_kv_layout(self): - # 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(self) - if cfg.enable_unified_memory: - self._declare( - "_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 self.use_mla_backend(): - allowed_full = { - "triton", - "fa3", - "trtllm_mla", - "flashinfer", - "cutedsl_mla", - "tokenspeed_mla", - } - else: - allowed_full = {"triton"} - backends = set(self._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 self.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." - ) + from sglang.srt.arg_groups.kv_cache_hook import handle_page_major_kv_layout + + handle_page_major_kv_layout(self) def _handle_dllm_inference(self): - cfg = resolving_view(self) - 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" - ) - self._declare( - "_handle_dllm_inference", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_handle_dllm_inference", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) + from sglang.srt.arg_groups.dllm_hook import handle_dllm_inference - from sglang.srt.arg_groups.overrides import ( - _dllm_attention_backend, - _dllm_overlap_disable, - run_post_process_pass, - ) - - run_post_process_pass(self, _dllm_attention_backend) - run_post_process_pass(self, _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(self, _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" - ) - self._declare( - "_handle_dllm_inference", - enable_hierarchical_cache=False, - ) - if cfg.enable_lmcache: - logger.warning( - "LMCache is disabled because of using diffusion LLM inference" - ) - self._declare("_handle_dllm_inference", enable_lmcache=False) - if cfg.enable_flexkv: - logger.warning( - "FlexKV is disabled because of using diffusion LLM inference" - ) - self._declare("_handle_dllm_inference", enable_flexkv=False) - - if cfg.pp_size > 1: - logger.warning( - "Pipeline parallelism is disabled because of using diffusion LLM inference" - ) - self._declare( - "_handle_dllm_inference", - pp_size=1, - ) - - if cfg.enable_lora: - logger.warning( - "Currently LoRA is not supported by diffusion LLM inference." - ) - self._declare("_handle_dllm_inference", enable_lora=False) - - if cfg.disaggregation_mode != "null": - logger.warning( - "Currently disaggregation is not supported by diffusion LLM inference." - ) - self._declare( - "_handle_dllm_inference", - disaggregation_mode="null", - ) - - if cfg.enable_mixed_chunk: - logger.warning( - "Mixed chunked prefill is disabled because of using diffusion LLM inference." - ) - self._declare( - "_handle_dllm_inference", - enable_mixed_chunk=False, - ) + handle_dllm_inference(self) def _handle_asr_validation(self): - """Validate transcription/ASR-specific server args.""" - cfg = resolving_view(self) - 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})." - ) + from sglang.srt.arg_groups.serving_hook import handle_asr_validation + + handle_asr_validation(self) def _validate_prefill_decode_interval(self): - cfg = resolving_view(self) - if cfg.prefill_decode_interval < 0: - raise ValueError("--prefill-decode-interval must be non-negative.") + from sglang.srt.arg_groups.validation_hook import ( + validate_prefill_decode_interval, + ) + + validate_prefill_decode_interval(self) def _handle_other_validations(self): - cfg = resolving_view(self) - 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" - ) + from sglang.srt.arg_groups.serving_hook import handle_other_validations - # 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") - self._declare( - "_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" - ) - self._declare( - "_handle_other_validations", - optimistic_prefill_attempts=0, - ) - elif resolved_view(self).uses_mamba_radix_cache: - logger.warning( - "Optimistic prefill does not support models that use " - "mamba radix cache." - ) - self._declare( - "_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" - ) - self._declare( - "_handle_other_validations", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_handle_other_validations", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - self._declare("_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." - ) - self._declare( - "_handle_other_validations", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.DECODE, backend=Backend.DISABLED - ), - ) - self._declare( - "_handle_other_validations", - cuda_graph_config=with_phase( - cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED - ), - ) - self._declare("_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): - self._declare( - "_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): - self._declare( - "_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) + handle_other_validations(self) def _handle_crash_dump_env(self): - cfg = resolving_view(self) - 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) + from sglang.srt.arg_groups.serving_hook import handle_crash_dump_env - 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, - ) + handle_crash_dump_env(self) def _handle_debug_utils(self): - cfg = resolving_view(self) - if is_in_ci() and cfg.soft_watchdog_timeout is None: - logger.info("Set soft_watchdog_timeout since in CI") - self._declare("_handle_debug_utils", soft_watchdog_timeout=300) + from sglang.srt.arg_groups.serving_hook import handle_debug_utils + + handle_debug_utils(self) @staticmethod def add_cli_args(parser: argparse.ArgumentParser): @@ -10281,521 +5485,34 @@ class ServerArgs: # 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(self) + from sglang.srt.arg_groups.validation_hook import check_two_batch_overlap - 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)." - ) + check_two_batch_overlap(self) def check_server_args(self): - cfg = resolving_view(self) + from sglang.srt.arg_groups.validation_hook import check_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 - self.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" - self.validate_buckets_rule("--prompt-tokens-buckets", cfg.prompt_tokens_buckets) - self.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(self, _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. - self._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}" - ) - - self.check_load_publish_args() + check_server_args(self) def check_load_publish_args(self): - """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.""" - mode = (self.load_publish_endpoint or "").strip() - if not mode or mode.lower() == "off": - return # disabled; nothing to validate + from sglang.srt.arg_groups.validation_hook import check_load_publish_args - server_cfg = resolving_view(self) - - from sglang.srt.disaggregation.kv_events import ( - KVEventsConfig, - resolve_load_pub_range, - ) - - if not self.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(self.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) + check_load_publish_args(self) def check_lora_server_args(self): - cfg = resolving_view(self) + from sglang.srt.arg_groups.lora_hook import check_lora_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: - self._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: - self._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 - self._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) - self._late_resolution( - "check_lora_server_args", lora_paths=parsed_lora_paths - ) - elif isinstance(cfg.lora_paths, dict): - self._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: - self._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: - self._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." + check_lora_server_args(self) def _check_lora_speculative_compatibility(self): - """Validate LoRA + speculative decoding combinations. + from sglang.srt.arg_groups.lora_hook import check_lora_speculative_compatibility - Adapters apply to the target only; a shared draft runs unadapted. - Matches resolved algorithm names (NEXTN has collapsed to EAGLE). - """ - cfg = resolving_view(self) - 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}." - ) + check_lora_speculative_compatibility(self) def validate_buckets_rule(self, arg_name: str, buckets_rule: List[str]): - if not buckets_rule: - return + from sglang.srt.arg_groups.validation_hook import validate_buckets_rule - 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', , , ]" - 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" + validate_buckets_rule(self, arg_name, buckets_rule) def adjust_mem_fraction_for_vlm(self, model_config): cfg = resolving_view(self) @@ -10838,23 +5555,9 @@ class ServerArgs: ) def validate_transfer_engine(self): - cfg = resolving_view(self) - 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 + from sglang.srt.arg_groups.model_path_hook import validate_transfer_engine + + return validate_transfer_engine(self) @property def _parsed_modelexpress_config(self) -> dict: diff --git a/test/registered/cpu/test_server_args_backend.py b/test/registered/cpu/test_server_args_backend.py index 423b11da3..21fd73fd9 100644 --- a/test/registered/cpu/test_server_args_backend.py +++ b/test/registered/cpu/test_server_args_backend.py @@ -20,7 +20,7 @@ class TestServerArgsCPUBackend(unittest.TestCase): server_args.sampling_backend = None 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): server_args = self._make_server_args() @@ -31,7 +31,7 @@ class TestServerArgsCPUBackend(unittest.TestCase): ) 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): server_args = self._make_server_args() diff --git a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py index 5db365996..cae79b9f3 100644 --- a/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py +++ b/test/registered/unit/configs/test_multimodal_piecewise_cuda_graph.py @@ -184,7 +184,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase): with ( 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() diff --git a/test/registered/unit/layers/attention/test_kda_helion_dispatcher.py b/test/registered/unit/layers/attention/test_kda_helion_dispatcher.py index 269d26e5d..31ded97e9 100644 --- a/test/registered/unit/layers/attention/test_kda_helion_dispatcher.py +++ b/test/registered/unit/layers/attention/test_kda_helion_dispatcher.py @@ -159,8 +159,11 @@ class TestHelionKDADispatcher(unittest.TestCase): def test_replayssm_accepts_helion_and_rejects_other_backends(self): with ( - patch("sglang.srt.server_args.is_sm100_supported", return_value=False), - patch("sglang.srt.server_args.is_cuda", return_value=False), + patch( + "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( model_path="dummy", @@ -184,8 +187,11 @@ class TestHelionKDADispatcher(unittest.TestCase): mamba_ssm_dtype="bfloat16", ) with ( - patch("sglang.srt.server_args.is_sm100_supported", return_value=True), - patch("sglang.srt.server_args.is_cuda", return_value=False), + patch( + "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() diff --git a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py index be6b13f4b..4f8c7a6b1 100644 --- a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py +++ b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py @@ -37,6 +37,20 @@ _READ_BEFORE_RESOLUTION = frozenset({"is_embedding"}) # has to be looked at. _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` # builds the model configuration and *then* collects the override declarations, # 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 it above the configuration another handler already cached. """ - tree = _parsed(_SRT / "server_args.py") - handler = next( - node - for node in ast.walk(tree) - if isinstance(node, ast.FunctionDef) - and node.name == "_handle_model_specific_adjustments" - ) + handler = None + for source, wanted in ( + (_SRT / "server_args.py", "_handle_model_specific_adjustments"), + *( + (path, "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 for node in ast.walk(handler): if not isinstance(node, ast.Call): @@ -283,6 +311,21 @@ def _hook_declarations(dispatch, source_module): 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(): """(ordered steps, {step: methods it reaches}) for the resolution dispatch.""" tree = _parsed(_SRT / "server_args.py") @@ -294,6 +337,28 @@ def _pipeline(): methods = { 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"] steps = [ name @@ -313,14 +378,18 @@ def _pipeline(): return seen seen.add(name) for node in ast.walk(methods[name]): + if not isinstance(node, ast.Call): + continue if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) + isinstance(node.func, ast.Attribute) and isinstance(node.func.value, ast.Name) and node.func.value.id == "self" and node.func.attr in methods ): reaches(node.func.attr, seen) + target = slot_target.get(name) + if target is not None: + reaches(target, seen) return seen step_lines = {} @@ -512,6 +581,7 @@ class TestModelConfigReadsResolvedInput(CustomTestCase): _READ_BEFORE_RESOLUTION | _STALE_IN_THE_MODEL_CONFIG | _STALE_FROM_THE_REGISTRIES + | _STALE_BEHIND_THE_EXPERT_PACK_BUILD ) late = sorted( field diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index 762cab64d..82facf593 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -520,8 +520,11 @@ class TestProgramsResolveBeforeReadingResolution(CustomTestCase): declarers = {"_declare", "declare_resolution", "declare_late_resolution"} fields = set() field_names = {field.name for field in _dataclasses.fields(_ServerArgs)} - for name in ("server_args.py", "arg_groups/overrides.py"): - tree = ast.parse((srt / name).read_text(encoding="utf-8-sig")) + # The record plus every module under `arg_groups/`: a handler declares + # 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): # Registry data: provider dict keys are field names as # *data*, invisible to the keyword scan below. Filtered diff --git a/test/registered/unit/server_args/test_resolution_reads_no_bag.py b/test/registered/unit/server_args/test_resolution_reads_no_bag.py index 63bed4ecf..390e6fc94 100644 --- a/test/registered/unit/server_args/test_resolution_reads_no_bag.py +++ b/test/registered/unit/server_args/test_resolution_reads_no_bag.py @@ -72,6 +72,15 @@ _ATTRIBUTE_SPELLED = _BAG_ACCESSORS - {"get_device"} _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): """`sglang.srt.a.b` -> the file, if it is one of ours.""" 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): """Does `entry` in `path` reach a bag accessor, following calls in-module?""" functions = _functions_in(path) + shadowed = _locally_shadowed_accessors(path) seen = set() def walk(name): @@ -216,7 +245,7 @@ def _reaches_a_bag(path, entry): continue if not isinstance(node.func, ast.Name): 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 found = walk(node.func.id) if found is not None: @@ -241,9 +270,7 @@ class TestResolutionReadsNoBag(CustomTestCase): def test_the_walk_finds_something_to_walk(self): """A collapsed import map would make the pin vacuous.""" - imported = _imported_symbols( - [_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"] - ) + imported = _imported_symbols(_pipeline_sources()) self.assertGreater( len(imported), 20, @@ -282,9 +309,7 @@ class TestResolutionReadsNoBag(CustomTestCase): ) def test_nothing_the_pipeline_calls_reads_a_bag(self): - imported = _imported_symbols( - [_SRT / "server_args.py", _SRT / "arg_groups" / "overrides.py"] - ) + imported = _imported_symbols(_pipeline_sources()) reachable = { (path, symbol) for path, symbols in imported.items() for symbol in symbols } | _registered_entries() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 95aad9e59..486290f3c 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -9,7 +9,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch 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.speculative_hook import handle_speculative_decoding 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") # 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() @@ -223,7 +225,7 @@ class TestMmEncoderDataParallelLogging(CustomTestCase): 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() 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 ) - 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() 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): 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() self.assertEqual( @@ -285,7 +287,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): def _set_model_type(server_args, *, 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): server_args = ServerArgs( model_path="dummy", @@ -295,7 +297,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): ) 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() self.assertEqual( @@ -307,12 +309,12 @@ class TestMultimodalFeatureTransport(CustomTestCase): self.assertIn("base GPU 2", 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): server_args = ServerArgs(model_path="dummy", keep_mm_feature_on_device=True) 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() self.assertEqual( @@ -335,12 +337,12 @@ class TestMultimodalFeatureTransport(CustomTestCase): with self.assertRaisesRegex(ValueError, "conflicts.*cuda_vmm"): 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): server_args = ServerArgs(model_path="dummy", mm_feature_transport="cpu") 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() self.assertEqual( @@ -361,7 +363,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): ) 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): server_args = ServerArgs(model_path="dummy") 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()) - @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): server_args = ServerArgs(model_path="dummy") 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()) - @patch("sglang.srt.server_args.os.path.exists", return_value=True) - @patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True) - @patch("sglang.srt.server_args.is_cuda", return_value=True) + @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True) + @patch( + "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( "sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport", return_value=True, @@ -410,7 +414,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): with patch.dict(os.environ, {}, clear=False): 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() self.assertEqual( @@ -422,9 +426,11 @@ class TestMultimodalFeatureTransport(CustomTestCase): self.assertIn("auto-resolved to cuda_vmm", output) self.assertIn("CUDA FABRIC", output) - @patch("sglang.srt.server_args.os.path.exists", return_value=True) - @patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=True) - @patch("sglang.srt.server_args.is_cuda", return_value=True) + @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=True) + @patch( + "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( "sglang.srt.model_loader.utils.supports_cuda_vmm_feature_transport", return_value=False, @@ -439,15 +445,17 @@ class TestMultimodalFeatureTransport(CustomTestCase): server_args = ServerArgs(model_path="dummy", nnodes=2) 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() self.assertEqual(resolution_result(server_args, "mm_feature_transport"), "cpu") 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.server_args.is_mnnvl_fabric_device", return_value=True) - @patch("sglang.srt.server_args.is_cuda", return_value=True) + @patch("sglang.srt.arg_groups.serving_hook.os.path.exists", return_value=False) + @patch( + "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( self, _mock_is_cuda, _mock_is_mnnvl, _mock_path_exists ): @@ -456,7 +464,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): with patch.dict(os.environ, {}, clear=False): 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() self.assertEqual( @@ -465,8 +473,10 @@ class TestMultimodalFeatureTransport(CustomTestCase): self.assertIn("no IMEX channel", "\n".join(logs.output)) - @patch("sglang.srt.server_args.is_mnnvl_fabric_device", return_value=False) - @patch("sglang.srt.server_args.is_cuda", return_value=True) + @patch( + "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( self, _mock_is_cuda, _mock_is_mnnvl ): @@ -482,7 +492,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): ) 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): server_args = ServerArgs(model_path="dummy", language_only=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()) - @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): server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_ipc") with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"): 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): server_args = ServerArgs( model_path="dummy", mm_feature_transport="cuda_ipc", nnodes=2 @@ -512,7 +522,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): with self.assertRaisesRegex(ValueError, "single node"): 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): server_args = ServerArgs( model_path="dummy", @@ -525,7 +535,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): patch.dict(os.environ, {"SGLANG_USE_CUDA_IPC_TRANSPORT": "1"}), 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() self.assertEqual( @@ -539,14 +549,14 @@ class TestMultimodalFeatureTransport(CustomTestCase): self.assertIn("2 tokenizer worker", 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): server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm") with self.assertRaisesRegex(ValueError, "requires NVIDIA CUDA"): 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): server_args = ServerArgs(model_path="dummy", mm_feature_transport="cuda_vmm") @@ -556,7 +566,7 @@ class TestMultimodalFeatureTransport(CustomTestCase): ): 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): server_args = ServerArgs( 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"): 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): server_args = ServerArgs( model_path="dummy", @@ -588,8 +598,8 @@ class TestMambaCacheStochasticRounding(unittest.TestCase): with self.assertRaisesRegex(ValueError, "NVIDIA CUDA"): server_args._handle_mamba_backend() - @patch("sglang.srt.server_args.is_cuda", return_value=True) - @patch("sglang.srt.server_args.is_sm100_supported", return_value=False) + @patch("sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=True) + @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): server_args = ServerArgs( model_path="dummy", @@ -1925,11 +1935,11 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase): prefill=PhaseConfig(backend=Backend.TC_PIECEWISE) ) with ( - patch("sglang.srt.server_args.is_hip", return_value=False), - patch("sglang.srt.server_args.is_npu", return_value=False), - patch("sglang.srt.server_args.is_cpu", return_value=False), - patch("sglang.srt.server_args.is_mps", return_value=False), - patch("sglang.srt.server_args.is_xpu", return_value=False), + patch("sglang.srt.arg_groups.cuda_graph_hook.is_hip", return_value=False), + patch("sglang.srt.arg_groups.cuda_graph_hook.is_npu", return_value=False), + patch("sglang.srt.arg_groups.cuda_graph_hook.is_cpu", return_value=False), + patch("sglang.srt.arg_groups.cuda_graph_hook.is_mps", return_value=False), + patch("sglang.srt.arg_groups.cuda_graph_hook.is_xpu", return_value=False), ): 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): 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() self.assertTrue(resolution_result(sa, "smg_grpc_mode")) self.assertTrue(any("--grpc-mode is deprecated" in line for line in cm.output)) diff --git a/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py b/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py index cee36b559..32ab6baf0 100644 --- a/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py +++ b/test/registered/unit/spec/test_dflash_extra_buffer_lazy.py @@ -46,7 +46,7 @@ class TestValidateMambaExtraBufferLazyDflash(CustomTestCase): ), mock.patch( # Keep the test runnable on CPU-only hosts: the platform assert is # not what is under test here. - "sglang.srt.server_args.is_cuda", + "sglang.srt.arg_groups.mamba_hook.is_cuda", return_value=True, ): ServerArgs._validate_mamba_extra_buffer( @@ -69,6 +69,27 @@ class TestValidateMambaExtraBufferLazyDflash(CustomTestCase): _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): """prepare_for_verify calls prepare_mamba_track_for_verify after the batch diff --git a/test/registered/unit/test_chain_read_ratchet.py b/test/registered/unit/test_chain_read_ratchet.py index f7d06f8fa..113473c22 100644 --- a/test/registered/unit/test_chain_read_ratchet.py +++ b/test/registered/unit/test_chain_read_ratchet.py @@ -195,15 +195,18 @@ def _declared_by_late_resolution(): 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. """ - 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() - for node in ast.walk(tree): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) - and node.func.attr == "_late_resolution" - ): - fields |= {keyword.arg for keyword in node.keywords if keyword.arg} + for source in sources: + for node in ast.walk(ast.parse(source.read_text(encoding="utf-8-sig"))): + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "_late_resolution" + ): + fields |= {keyword.arg for keyword in node.keywords if keyword.arg} return fields diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index 891648105..ce332cd1f 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -1236,6 +1236,33 @@ class TestGoldenModelOverrides(_IsolatedPublish): "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): sa = self._construct("LlamaForCausalLM", "llama", attention_backend="triton") self.assertEqual(self._resolved(sa, "attention_backend"), "triton") diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index e8f91a5ad..a8c100f6a 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -511,12 +511,28 @@ class TestSuppliedInstanceExposure(CustomTestCase): "prefill_attention_backend", "speculative_draft_attention_backend", } - deprecated = next( - node - for node in ast.walk(sa_class) - if isinstance(node, ast.FunctionDef) - and node.name == "_handle_deprecated_args" - ) + + # The handler lives in `arg_groups/serving_hook.py`, reached either as a + # record method or as a bare-name call, so look the loop up by both. + def _deprecated_alias_handler(): + for node in ast.walk(sa_class): + if ( + isinstance(node, ast.FunctionDef) + 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 = [ {elt.value for elt in node.iter.elts if isinstance(elt, ast.Constant)} for node in ast.walk(deprecated) diff --git a/test/registered/utils/test_network_address.py b/test/registered/utils/test_network_address.py index 87194b3aa..6a9a9d773 100644 --- a/test/registered/utils/test_network_address.py +++ b/test/registered/utils/test_network_address.py @@ -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") # 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()