[refactor] Migrate the moe_runner_backend / quantization resolution chains (stack 13/15) (#30075)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
276fbfe880
commit
3836cba9ee
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user