diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 2e675ba4c..988518906 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -38,6 +38,8 @@ from sglang.srt.runtime_context import resolve_flag_leaf from sglang.srt.utils.common import ( cpu_has_amx_support, get_device_sm, + get_nvidia_driver_version, + get_quantization_config, is_blackwell_supported, is_cpu, is_cuda, @@ -48,6 +50,7 @@ from sglang.srt.utils.common import ( is_sm90_supported, is_sm100_supported, is_sm120_supported, + is_triton_kernels_available, is_xpu, xpu_has_xmx_support, ) @@ -295,12 +298,70 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: f"but got '{server_args.dtype}'. Please use --dtype bfloat16 or remove --dtype to use auto." ) quantization_config = getattr(hf_config, "quantization_config", None) - if ( + is_mxfp4_quant_format = ( quantization_config is not None and quantization_config.get("quant_method") == "mxfp4" - ): + ) + if is_mxfp4_quant_format: # use bf16 for mxfp4 triton kernels overrides["dtype"] = "bfloat16" + if server_args.moe_runner_backend == "auto": + from sglang.srt.environ import envs + + if is_sm100_supported() and is_mxfp4_quant_format: + overrides["moe_runner_backend"] = "flashinfer_mxfp4" + logger.warning( + "Detected SM100 and MXFP4 quantization format for GPT-OSS model, enabling FlashInfer MXFP4 MOE kernel." + ) + elif is_sm120_supported() and is_mxfp4_quant_format: + # trtllm-gen only supports SM100 + overrides["moe_runner_backend"] = "marlin" + logger.warning( + "Detected SM120 and MXFP4 quantization format for GPT-OSS model, enabling Marlin MOE kernel." + ) + elif (is_hip() and envs.SGLANG_USE_AITER.get()) and is_mxfp4_quant_format: + overrides["moe_runner_backend"] = "auto" + logger.warning( + "Detected ROCm and MXFP4 quantization format for GPT-OSS model, enabling aiter MXFP4 MOE kernel." + ) + ## The AITER MXFP4 fused-MoE path for GPT-OSS expects the + ## SEPARATED gate/up tile layout (matches the + ## `gptoss_fp4_tuned_fmoe.csv` flydsl entries and the + ## Mxfp4MoEMethod weight shuffle). Other AITER MXFP4 + ## callers default to INTERLEAVE; opt this path out + ## unless the user explicitly overrode it. + # envs.SGLANG_USE_AITER_MOE_GU_ITLV.set(False) + elif is_hip() and envs.SGLANG_USE_AITER.get(): + # For GPT-OSS bf16 on ROCm with aiter, use triton backend + # because aiter CK kernel doesn't support all GEMM dimensions + overrides["moe_runner_backend"] = "triton" + logger.warning( + "Detected ROCm with SGLANG_USE_AITER for GPT-OSS bf16 model, using triton MOE kernel." + ) + elif is_musa() and envs.SGLANG_DEEPEP_BF16_DISPATCH.get(): + overrides["moe_runner_backend"] = "deep_gemm" + logger.warning( + "Detected MUSA with SGLANG_DEEPEP_BF16_DISPATCH for bf16 model, using deep_gemm kernel." + ) + elif ( + server_args.ep_size == 1 + and is_triton_kernels_available() + and server_args.quantization is None + and not (is_cpu() and cpu_has_amx_support()) + ): + # The triton_kernels package segfaults on Blackwell (B200) + # with NVIDIA driver >= 595. Fall back to triton backend. + if is_blackwell_supported() and get_nvidia_driver_version() >= (595,): + overrides["moe_runner_backend"] = "triton" + logger.warning( + "Detected GPT-OSS model on Blackwell with driver >= 595, " + "using triton MOE kernel to avoid triton_kernels SIGSEGV." + ) + else: + overrides["moe_runner_backend"] = "triton_kernel" + logger.warning( + "Detected GPT-OSS model, enabling triton_kernels MOE kernel." + ) return overrides @@ -309,6 +370,7 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: if server_args.device == "cpu": return {} + overrides: Dict[str, Any] = {} # Auto-select attention backend for Llama4 if not specified if server_args.attention_backend is None: if is_sm100_supported(): @@ -324,8 +386,14 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: logger.warning( f"Use {backend} as attention backend on {platform} for Llama4 model" ) - return {"attention_backend": backend} - return {} + overrides["attention_backend"] = backend + if is_sm100_supported() and server_args.moe_runner_backend == "auto": + if server_args.quantization in {"fp8", "modelopt_fp8"}: + overrides["moe_runner_backend"] = "flashinfer_trtllm" + logger.info( + "Use flashinfer_trtllm as MoE runner backend on SM100 for Llama4" + ) + return overrides @_register_for( @@ -334,18 +402,27 @@ def _llama4_overrides(server_args: Any, hf_config: Any) -> dict: "Gemma4UnifiedForConditionalGeneration", ) def _gemma4_overrides(server_args: Any, hf_config: Any) -> dict: + overrides: Dict[str, Any] = {} default_attention_backend = "trtllm_mha" if is_sm100_supported() else "triton" if server_args.is_attention_backend_not_set(): logger.info( f"Use {default_attention_backend} as default attention backend for Gemma4" ) - return {"attention_backend": default_attention_backend} + overrides["attention_backend"] = default_attention_backend # If only one split backend is set, keep the other side on a # Gemma4-compatible fallback instead of letting generic backend selection # choose an unsupported backend later. - if server_args.attention_backend is None: - return {"attention_backend": default_attention_backend} - return {} + elif server_args.attention_backend is None: + overrides["attention_backend"] = default_attention_backend + if is_sm100_supported() and server_args.moe_runner_backend == "auto": + if server_args.get_model_config().quantization == "modelopt_fp4": + overrides["quantization"] = "modelopt_fp4" + overrides["moe_runner_backend"] = "flashinfer_trtllm" + logger.info( + "Use flashinfer_trtllm as MoE runner backend on " + "SM100 for Gemma-4 (modelopt_fp4)" + ) + return overrides @_register_for("MiniCPMV4_6ForConditionalGeneration") @@ -428,12 +505,71 @@ def _qwen3vl_overrides(server_args: Any, hf_config: Any) -> dict: return {} +@_register_for( + "Qwen3MoeForCausalLM", + "Qwen3VLMoeForConditionalGeneration", + "Qwen3NextForCausalLM", + "Qwen3_5MoeForConditionalGeneration", + "InternS2PreviewForConditionalGeneration", + "Qwen3_5ForConditionalGeneration", +) +def _qwen3_moe_family_overrides(server_args: Any, hf_config: Any) -> dict: + overrides: Dict[str, Any] = {} + if is_sm100_supported(): + quant_method = get_quantization_config(hf_config) + quantization = server_args.quantization + if ( + quantization is None + and not server_args._quantization_explicitly_unset + and quant_method is not None + ): + overrides["quantization"] = quant_method + quantization = quant_method + if ( + (quantization in ("fp8", "modelopt_fp4") or quantization is None) + and server_args.moe_a2a_backend == "none" + and server_args.moe_runner_backend == "auto" + ): + overrides["moe_runner_backend"] = "flashinfer_trtllm" + logger.info( + "Use flashinfer_trtllm as MoE runner backend on sm100 for " + f"{hf_config.architectures[0]}" + ) + return overrides + + @_register_for("Glm4MoeForCausalLM") def _glm4_moe_overrides(server_args: Any, hf_config: Any) -> dict: + overrides: Dict[str, Any] = {} + if is_sm100_supported(): + quantization_config = getattr(hf_config, "quantization_config", None) + quant_method = ( + quantization_config.get("quant_method") + if quantization_config is not None + else None + ) + quantization = server_args.quantization + if ( + quantization is None + and not server_args._quantization_explicitly_unset + and quant_method is not None + ): + overrides["quantization"] = quant_method + quantization = quant_method + if ( + quantization in {"modelopt_fp4", None} + and server_args.moe_a2a_backend == "none" + and server_args.moe_runner_backend == "auto" + ): + overrides["moe_runner_backend"] = "flashinfer_trtllm" + logger.info( + "Use flashinfer_trtllm as MoE runner backend on sm100 for Glm4MoeForCausalLM" + ) logger.info( "Enable TF32 matmul for Glm4MoeForCausalLM model to improve gate gemm performance." ) - return {"enable_tf32_matmul": True} + overrides["enable_tf32_matmul"] = True + return overrides @_register_for("Olmo2ForCausalLM") @@ -766,6 +902,86 @@ def _page_size_default(view: Any) -> dict: return {"page_size": 64} +@register_post_process +def _moe_runner_backend_quant_constraints(view: Any) -> dict: + """The quantization-driven moe_runner_backend resolutions at the head of + _handle_moe_kernel_config. The backend-compatibility asserts and the + disable_shared_experts_fusion writes (post-publish writers exist for that + field) stay in the handler.""" + moe_runner_backend = view.moe_runner_backend + if view.quantization == "nvfp4_online": + if not is_sm100_supported(): + raise ValueError( + "--quantization nvfp4_online is supported only on " + "NVIDIA Blackwell SM100/SM103 GPUs." + ) + if moe_runner_backend == "auto": + moe_runner_backend = "flashinfer_trtllm" + elif moe_runner_backend not in [ + "flashinfer_trtllm", + "flashinfer_trtllm_routed", + ]: + raise ValueError( + "--quantization nvfp4_online supports only " + "--moe-runner-backend flashinfer_trtllm or " + "flashinfer_trtllm_routed." + ) + if view.quantization == "mxfp8": + if moe_runner_backend == "auto": + moe_runner_backend = "flashinfer_trtllm" + elif moe_runner_backend not in [ + "cutlass", + "flashinfer_trtllm", + "flashinfer_trtllm_routed", + ]: + logger.warning( + "mxfp8 quantization supports only cutlass, flashinfer_trtllm, " + "or flashinfer_trtllm_routed backends. " + f"Overriding {moe_runner_backend!r}." + ) + moe_runner_backend = "flashinfer_trtllm" + if ( + moe_runner_backend == "auto" + and view.quantization == "modelopt_fp4" + and is_sm120_supported() + ): + moe_runner_backend = "flashinfer_cutlass" + logger.info( + "Use flashinfer_cutlass as MoE runner backend on SM120 for " + "modelopt_fp4 (trtllm-gen MoE kernels are SM100-only)" + ) + if moe_runner_backend != view.moe_runner_backend: + return {"moe_runner_backend": moe_runner_backend} + return {} + + +@register_post_process +def _cutlass_moe_env_override(view: Any) -> dict: + from sglang.srt.environ import envs + + if envs.SGLANG_CUTLASS_MOE.get(): + logger.warning( + "SGLANG_CUTLASS_MOE is deprecated, use --moe-runner-backend=cutlass and/or --speculative-moe-runner-backend=cutlass instead" + ) + assert view.quantization in [ + "fp8", + "mxfp8", + ], "cutlass MoE is only supported with fp8/mxfp8 quantization" + return {"moe_runner_backend": "cutlass"} + return {} + + +@register_post_process +def _gguf_quantization(view: Any) -> dict: + from sglang.srt.utils.hf_transformers_utils import check_gguf_file + + if (view.load_format == "auto" or view.load_format == "gguf") and check_gguf_file( + view.model_path + ): + return {"quantization": "gguf"} + return {} + + @register_post_process def _dllm_attention_backend(view: Any) -> dict: if view.dllm_algorithm is None: diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 4dcc6297a..d93d2ee12 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -290,6 +290,10 @@ class AttnFlags(_StaticFlags): class MoeFlags(_StaticFlags): """MoE-family resolved flags (leaves arrive with the V3 sweeps).""" + # Resolved MoE runner backend; the pristine user request stays on + # server_args.moe_runner_backend. + runner_backend: str = "auto" + @dataclasses.dataclass class CaptureFlags(_FlagGroupBase): @@ -320,6 +324,7 @@ class Flags(_StaticFlags): disable_hybrid_swa_memory: bool = False sampling_backend: str | None = None page_size: int | None = None + quantization: str | None = None def freeze(self) -> None: for field in dataclasses.fields(self): @@ -335,6 +340,7 @@ class Flags(_StaticFlags): # family as readers migrate. FLAG_LEAF_MAP: dict[str, str] = { "attention_backend": "attn.backend", + "moe_runner_backend": "moe.runner_backend", } diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 822207431..3c939a303 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -63,12 +63,10 @@ from sglang.srt.speculative.decoupled_spec_io import DecoupledSpecIpcConfig from sglang.srt.utils.common import ( LORA_TARGET_ALL_MODULES, SUPPORTED_LORA_TARGET_MODULES, - cpu_has_amx_support, get_device, get_device_memory_capacity, get_device_sm, get_int_env_var, - get_nvidia_driver_version, get_quantization_config, has_fp8_weights_in_checkpoint, human_readable_int, @@ -87,7 +85,6 @@ from sglang.srt.utils.common import ( is_sm90_supported, is_sm100_supported, is_sm120_supported, - is_triton_kernels_available, is_xpu, json_list_type, nullable_str, @@ -576,7 +573,11 @@ class ServerArgs: ] = "auto" quantization: A[ Optional[str], - Arg(help="The quantization method.", choices=QUANTIZATION_CHOICES), + Arg( + help="The quantization method.", + choices=QUANTIZATION_CHOICES, + model_overridable=True, + ), ] = None quantization_param_path: A[ Optional[str], @@ -1721,6 +1722,7 @@ class ServerArgs: Arg( help="Choose the runner backend for MoE.", choices=MOE_RUNNER_BACKEND_CHOICES, + model_overridable=True, ), ] = "auto" flashinfer_mxfp4_moe_precision: A[ @@ -4099,65 +4101,8 @@ class ServerArgs: # The mxfp4 dtype override moved to the override registry # (arg_groups/overrides.py: _gpt_oss_overrides). - if self.moe_runner_backend == "auto": - if is_sm100_supported() and is_mxfp4_quant_format: - self.moe_runner_backend = "flashinfer_mxfp4" - logger.warning( - "Detected SM100 and MXFP4 quantization format for GPT-OSS model, enabling FlashInfer MXFP4 MOE kernel." - ) - elif is_sm120_supported() and is_mxfp4_quant_format: - # trtllm-gen only supports SM100 - self.moe_runner_backend = "marlin" - logger.warning( - "Detected SM120 and MXFP4 quantization format for GPT-OSS model, enabling Marlin MOE kernel." - ) - elif ( - is_hip() and envs.SGLANG_USE_AITER.get() - ) and is_mxfp4_quant_format: - self.moe_runner_backend = "auto" - logger.warning( - "Detected ROCm and MXFP4 quantization format for GPT-OSS model, enabling aiter MXFP4 MOE kernel." - ) - ## The AITER MXFP4 fused-MoE path for GPT-OSS expects the - ## SEPARATED gate/up tile layout (matches the - ## `gptoss_fp4_tuned_fmoe.csv` flydsl entries and the - ## Mxfp4MoEMethod weight shuffle). Other AITER MXFP4 - ## callers default to INTERLEAVE; opt this path out - ## unless the user explicitly overrode it. - # envs.SGLANG_USE_AITER_MOE_GU_ITLV.set(False) - elif is_hip() and envs.SGLANG_USE_AITER.get(): - # For GPT-OSS bf16 on ROCm with aiter, use triton backend - # because aiter CK kernel doesn't support all GEMM dimensions - self.moe_runner_backend = "triton" - logger.warning( - "Detected ROCm with SGLANG_USE_AITER for GPT-OSS bf16 model, using triton MOE kernel." - ) - elif is_musa() and envs.SGLANG_DEEPEP_BF16_DISPATCH.get(): - self.moe_runner_backend = "deep_gemm" - logger.warning( - "Detected MUSA with SGLANG_DEEPEP_BF16_DISPATCH for bf16 model, using deep_gemm kernel." - ) - elif ( - self.ep_size == 1 - and is_triton_kernels_available() - and self.quantization is None - and not (is_cpu() and cpu_has_amx_support()) - ): - # The triton_kernels package segfaults on Blackwell (B200) - # with NVIDIA driver >= 595. Fall back to triton backend. - if is_blackwell_supported() and get_nvidia_driver_version() >= ( - 595, - ): - self.moe_runner_backend = "triton" - logger.warning( - "Detected GPT-OSS model on Blackwell with driver >= 595, " - "using triton MOE kernel to avoid triton_kernels SIGSEGV." - ) - else: - self.moe_runner_backend = "triton_kernel" - logger.warning( - "Detected GPT-OSS model, enabling triton_kernels MOE kernel." - ) + # The moe_runner_backend selection moved to the override registry + # (arg_groups/overrides.py: _gpt_oss_overrides). if self.moe_runner_backend == "triton_kernel": assert ( @@ -4224,12 +4169,8 @@ class ServerArgs: "trtllm_mha", "intel_xpu", }, f"fa3, aiter, triton, ascend, trtllm_mha or intel_xpu is required for Llama4 model but got {self.attention_backend}" - if is_sm100_supported() and self.moe_runner_backend == "auto": - if self.quantization in {"fp8", "modelopt_fp8"}: - self.moe_runner_backend = "flashinfer_trtllm" - logger.info( - "Use flashinfer_trtllm as MoE runner backend on SM100 for Llama4" - ) + # 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 ( @@ -4249,14 +4190,8 @@ class ServerArgs: f"got prefill={prefill_backend}, decode={decode_backend}" ) - if is_sm100_supported() and self.moe_runner_backend == "auto": - if self.get_model_config().quantization == "modelopt_fp4": - self.quantization = "modelopt_fp4" - self.moe_runner_backend = "flashinfer_trtllm" - logger.info( - "Use flashinfer_trtllm as MoE runner backend on " - "SM100 for Gemma-4 (modelopt_fp4)" - ) + # The quantization/moe_runner_backend resolution moved to the override + # registry (arg_groups/overrides.py: _gemma4_overrides). elif model_arch == "MossVLForConditionalGeneration": if self.is_attention_backend_not_set(): self.prefill_attention_backend = "flashinfer" @@ -4309,27 +4244,8 @@ class ServerArgs: "InternS2PreviewForConditionalGeneration", "Qwen3_5ForConditionalGeneration", ]: - if is_sm100_supported(): - quant_method = get_quantization_config(hf_config) - if ( - self.quantization is None - and not self._quantization_explicitly_unset - and quant_method is not None - ): - self.quantization = quant_method - if ( - ( - self.quantization in ("fp8", "modelopt_fp4") - or self.quantization is None - ) - and self.moe_a2a_backend == "none" - and self.moe_runner_backend == "auto" - ): - self.moe_runner_backend = "flashinfer_trtllm" - logger.info( - "Use flashinfer_trtllm as MoE runner backend on sm100 for " - f"{model_arch}" - ) + # The quantization/moe_runner_backend resolution moved to the override + # registry (arg_groups/overrides.py: _qwen3_moe_family_overrides). if model_arch in [ "Qwen3NextForCausalLM", @@ -4349,30 +4265,10 @@ class ServerArgs: self._handle_mamba_radix_cache(model_arch=model_arch) elif model_arch in ["Glm4MoeForCausalLM"]: - if is_sm100_supported(): - quantization_config = getattr(hf_config, "quantization_config", None) - quant_method = ( - quantization_config.get("quant_method") - if quantization_config is not None - else None - ) - if ( - self.quantization is None - and not self._quantization_explicitly_unset - and quant_method is not None - ): - self.quantization = quant_method - if ( - self.quantization in {"modelopt_fp4", None} - and self.moe_a2a_backend == "none" - and self.moe_runner_backend == "auto" - ): - self.moe_runner_backend = "flashinfer_trtllm" - logger.info( - "Use flashinfer_trtllm as MoE runner backend on sm100 for Glm4MoeForCausalLM" - ) - # enable_tf32_matmul moved to the override registry - # (arg_groups/overrides.py: _glm4_moe_overrides). + # 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 [ "FalconH1ForCausalLM", @@ -5187,48 +5083,16 @@ class ServerArgs: ), "Please enable dp attention when setting enable_dp_lm_head. " def _handle_moe_kernel_config(self): - if self.quantization == "nvfp4_online": - if not is_sm100_supported(): - raise ValueError( - "--quantization nvfp4_online is supported only on " - "NVIDIA Blackwell SM100/SM103 GPUs." - ) - if self.moe_runner_backend == "auto": - self.moe_runner_backend = "flashinfer_trtllm" - elif self.moe_runner_backend not in [ - "flashinfer_trtllm", - "flashinfer_trtllm_routed", - ]: - raise ValueError( - "--quantization nvfp4_online supports only " - "--moe-runner-backend flashinfer_trtllm or " - "flashinfer_trtllm_routed." - ) - if self.quantization == "mxfp8": - if self.moe_runner_backend == "auto": - self.moe_runner_backend = "flashinfer_trtllm" - elif self.moe_runner_backend not in [ - "cutlass", - "flashinfer_trtllm", - "flashinfer_trtllm_routed", - ]: - logger.warning( - "mxfp8 quantization supports only cutlass, flashinfer_trtllm, " - "or flashinfer_trtllm_routed backends. " - f"Overriding {self.moe_runner_backend!r}." - ) - self.moe_runner_backend = "flashinfer_trtllm" + # The quantization-driven runner resolutions moved to the pipeline + # (arg_groups/overrides.py: _moe_runner_backend_quant_constraints); + # the compatibility asserts and fusion writes stay below. + from sglang.srt.arg_groups.overrides import ( + _cutlass_moe_env_override, + _moe_runner_backend_quant_constraints, + run_post_process_pass, + ) - if ( - self.moe_runner_backend == "auto" - and self.quantization == "modelopt_fp4" - and is_sm120_supported() - ): - self.moe_runner_backend = "flashinfer_cutlass" - logger.info( - "Use flashinfer_cutlass as MoE runner backend on SM120 for " - "modelopt_fp4 (trtllm-gen MoE kernels are SM100-only)" - ) + run_post_process_pass(self, _moe_runner_backend_quant_constraints) if self.moe_runner_backend == "flashinfer_cutlass": assert self.quantization in [ @@ -5293,15 +5157,11 @@ class ServerArgs: "FlashInfer TRTLLM routed MoE is enabled. --disable-shared-experts-fusion is automatically set." ) - if envs.SGLANG_CUTLASS_MOE.get(): - logger.warning( - "SGLANG_CUTLASS_MOE is deprecated, use --moe-runner-backend=cutlass and/or --speculative-moe-runner-backend=cutlass instead" - ) - assert self.quantization in [ - "fp8", - "mxfp8", - ], "cutlass MoE is only supported with fp8/mxfp8 quantization" - self.moe_runner_backend = "cutlass" + # The deprecated SGLANG_CUTLASS_MOE override moved to the pipeline + # (arg_groups/overrides.py: _cutlass_moe_env_override). It sits after + # the fusion blocks above on purpose: they must observe the + # pre-override runner value, exactly as they did imperatively. + run_post_process_pass(self, _cutlass_moe_env_override) if self.moe_runner_backend == "cutlass" and self.quantization in [ "fp8", "mxfp8", @@ -5751,10 +5611,19 @@ class ServerArgs: ) def _handle_load_format(self): + # 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. + from sglang.srt.arg_groups.overrides import ( + _gguf_quantization, + run_post_process_pass, + ) + + run_post_process_pass(self, _gguf_quantization) if ( self.load_format == "auto" or self.load_format == "gguf" ) and check_gguf_file(self.model_path): - self.quantization = self.load_format = "gguf" + self.load_format = "gguf" if self.load_format == "auto" and self._is_mistral_native_format(): self.load_format = "mistral" diff --git a/test/registered/unit/test_model_overrides.py b/test/registered/unit/test_model_overrides.py index f50465889..790e934a3 100644 --- a/test/registered/unit/test_model_overrides.py +++ b/test/registered/unit/test_model_overrides.py @@ -70,6 +70,8 @@ class TestModelOverridableWhitelist(CustomTestCase): "sampling_backend", "attention_backend", "page_size", + "moe_runner_backend", + "quantization", } ), ) @@ -841,6 +843,88 @@ class TestGoldenModelOverrides(_IsolatedPublish): _qwen3vl_overrides(SimpleNamespace(page_size=64), None), {} ) + def test_moe_runner_quant_constraint_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _moe_runner_backend_quant_constraints, + ) + + def _view(**kw): + defaults = dict(quantization=None, moe_runner_backend="auto") + defaults.update(kw) + return ResolvedView(SimpleNamespace(**defaults)) + + with patch.object(overrides_module, "is_sm100_supported", return_value=True): + self.assertEqual( + _moe_runner_backend_quant_constraints( + _view(quantization="nvfp4_online") + ), + {"moe_runner_backend": "flashinfer_trtllm"}, + ) + with self.assertRaises(ValueError): # incompatible explicit backend + _moe_runner_backend_quant_constraints( + _view(quantization="nvfp4_online", moe_runner_backend="triton") + ) + self.assertEqual( + _moe_runner_backend_quant_constraints(_view(quantization="mxfp8")), + {"moe_runner_backend": "flashinfer_trtllm"}, + ) + with patch.object(overrides_module, "is_sm120_supported", return_value=True): + self.assertEqual( + _moe_runner_backend_quant_constraints( + _view(quantization="modelopt_fp4") + ), + {"moe_runner_backend": "flashinfer_cutlass"}, + ) + self.assertEqual(_moe_runner_backend_quant_constraints(_view()), {}) + + def test_cutlass_moe_env_override_pass(self): + from sglang.srt.arg_groups.overrides import ( + ResolvedView, + _cutlass_moe_env_override, + ) + + with patch("sglang.srt.environ.envs.SGLANG_CUTLASS_MOE") as e: + e.get.return_value = True + self.assertEqual( + _cutlass_moe_env_override( + ResolvedView(SimpleNamespace(quantization="fp8")) + ), + {"moe_runner_backend": "cutlass"}, + ) + with self.assertRaises(AssertionError): + _cutlass_moe_env_override( + ResolvedView(SimpleNamespace(quantization=None)) + ) + e.get.return_value = False + self.assertEqual( + _cutlass_moe_env_override(ResolvedView(SimpleNamespace())), {} + ) + + def test_gguf_quantization_pass(self): + from sglang.srt.arg_groups.overrides import ResolvedView, _gguf_quantization + + with patch( + "sglang.srt.utils.hf_transformers_utils.check_gguf_file", + return_value=True, + ): + self.assertEqual( + _gguf_quantization( + ResolvedView( + SimpleNamespace(load_format="auto", model_path="x.gguf") + ) + ), + {"quantization": "gguf"}, + ) + self.assertEqual( + _gguf_quantization( + ResolvedView( + SimpleNamespace(load_format="safetensors", model_path="x") + ) + ), + {}, + ) + def test_page_constraint_passes_at_callable_level(self): from sglang.srt.arg_groups.overrides import ( ResolvedView, @@ -942,6 +1026,10 @@ class TestGoldenModelOverrides(_IsolatedPublish): device="cuda", attention_backend=None, is_attention_backend_not_set=lambda: True, + # keep the (now-absorbed) quant/moe blocks inert so these + # assertions stay attention-only + moe_runner_backend="triton", + quantization=None, ) defaults.update(kw) return SimpleNamespace(**defaults) @@ -989,9 +1077,55 @@ class TestGoldenModelOverrides(_IsolatedPublish): self.assertEqual( _gemma4_overrides(_args(), None), {"attention_backend": "triton"} ) - # Glm4Moe: unconditional tf32 declaration (quant/moe writes stay in - # the branch until their field chains migrate) - self.assertEqual(_glm4_moe_overrides(None, None), {"enable_tf32_matmul": True}) + # Glm4Moe: unconditional tf32 declaration + (sm100) quant/moe absorption + with patch.object(overrides_module, "is_sm100_supported", return_value=False): + self.assertEqual( + _glm4_moe_overrides(None, None), {"enable_tf32_matmul": True} + ) + with patch.object(overrides_module, "is_sm100_supported", return_value=True): + self.assertEqual( + _glm4_moe_overrides( + SimpleNamespace( + quantization=None, + _quantization_explicitly_unset=False, + moe_a2a_backend="none", + moe_runner_backend="auto", + ), + SimpleNamespace( + quantization_config={"quant_method": "modelopt_fp4"} + ), + ), + { + "quantization": "modelopt_fp4", + "moe_runner_backend": "flashinfer_trtllm", + "enable_tf32_matmul": True, + }, + ) + + def test_qwen3_moe_family_quant_absorption(self): + from sglang.srt.arg_groups.overrides import _qwen3_moe_family_overrides + + with patch.object(overrides_module, "is_sm100_supported", return_value=True): + with patch.object( + overrides_module, "get_quantization_config", return_value="fp8" + ): + self.assertEqual( + _qwen3_moe_family_overrides( + SimpleNamespace( + quantization=None, + _quantization_explicitly_unset=False, + moe_a2a_backend="none", + moe_runner_backend="auto", + ), + SimpleNamespace(architectures=["Qwen3MoeForCausalLM"]), + ), + { + "quantization": "fp8", + "moe_runner_backend": "flashinfer_trtllm", + }, + ) + with patch.object(overrides_module, "is_sm100_supported", return_value=False): + self.assertEqual(_qwen3_moe_family_overrides(None, None), {}) def test_step3p_declarations_at_callable_level(self): from sglang.srt.arg_groups.overrides import _step3p_overrides