[DeepSeek V4] Default FP4 checkpoints to FlashInfer MXFP4 MoE (#35919)
This commit is contained in:
@@ -1164,16 +1164,27 @@ def _deepseek_v4_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
overrides["swa_full_tokens_ratio"] = 0.1
|
||||
logger.info(f"Setting swa_full_tokens_ratio to 0.1 for {model_arch}.")
|
||||
|
||||
# nvidia/DeepSeek-V4-Pro-NVFP4 uses flashinfer_trtllm_routed MoE runner backend.
|
||||
if (
|
||||
server_args.moe_runner_backend == "auto"
|
||||
and server_args.get_model_config().nvfp4_moe_meta is not None
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
|
||||
logger.info(
|
||||
"Use flashinfer_trtllm_routed as MoE runner backend for "
|
||||
f"{model_arch} hybrid FP8+NVFP4 checkpoint."
|
||||
)
|
||||
if server_args.moe_runner_backend == "auto":
|
||||
model_config = server_args.get_model_config()
|
||||
# nvidia/DeepSeek-V4-Pro-NVFP4 uses the routed TRT-LLM runner.
|
||||
if model_config.nvfp4_moe_meta is not None:
|
||||
overrides["moe_runner_backend"] = "flashinfer_trtllm_routed"
|
||||
logger.info(
|
||||
"Use flashinfer_trtllm_routed as MoE runner backend for "
|
||||
f"{model_arch} hybrid FP8+NVFP4 checkpoint."
|
||||
)
|
||||
elif (
|
||||
server_args.device == "cuda"
|
||||
and not is_hip()
|
||||
and server_args.moe_a2a_backend == "none"
|
||||
and not envs.SGLANG_DSV4_FP4_DEQUANT.get()
|
||||
and model_config.is_fp4_experts
|
||||
and (is_sm90_supported() or is_sm100_supported() or is_sm120_supported())
|
||||
):
|
||||
overrides["moe_runner_backend"] = "flashinfer_mxfp4"
|
||||
logger.info(
|
||||
"Use flashinfer_mxfp4 as MoE runner backend for " f"{model_arch}."
|
||||
)
|
||||
return overrides
|
||||
|
||||
|
||||
@@ -1921,20 +1932,6 @@ def _deepseek_v4_kv_cache_dtype(view: Any) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
@register_post_process
|
||||
def _deepseek_v4_sm120_moe(view: Any) -> dict:
|
||||
"""Default DeepSeek V4 MXFP4 experts to FlashInfer CUTLASS on SM120."""
|
||||
hf_config = view.get_model_config().hf_config
|
||||
if hf_config.architectures[0] != "DeepseekV4ForCausalLM":
|
||||
return {}
|
||||
if is_sm120_supported() and view.moe_runner_backend == "auto":
|
||||
logger.info(
|
||||
"Use flashinfer_mxfp4 as MoE runner backend on SM120 for DeepseekV4"
|
||||
)
|
||||
return {"moe_runner_backend": "flashinfer_mxfp4"}
|
||||
return {}
|
||||
|
||||
|
||||
@_register_for("MuseGlimmerForConditionalGeneration", "MuseGlimmerForCausalLM")
|
||||
def _muse_glimmer_fp4_gemm_runner_overrides(server_args: Any, hf_config: Any) -> dict:
|
||||
if is_sm120_supported() and server_args.fp4_gemm_runner_backend == "auto":
|
||||
|
||||
@@ -5540,15 +5540,6 @@ class ServerArgs:
|
||||
validate_deepseek_v4_cp(self)
|
||||
validate_deepseek_v4_mega_moe_token_budget(self)
|
||||
|
||||
# The SM120 marlin fallback moved to the resolution pipeline
|
||||
# (arg_groups/overrides.py: _deepseek_v4_sm120_moe), invoked here
|
||||
# at its legacy slot.
|
||||
from sglang.srt.arg_groups.overrides import (
|
||||
_deepseek_v4_sm120_moe,
|
||||
run_post_process_pass,
|
||||
)
|
||||
|
||||
run_post_process_pass(self, _deepseek_v4_sm120_moe)
|
||||
if is_sm120_supported():
|
||||
# SM120 lacks tcgen05/TMEM: disable features that depend on
|
||||
# DeepGEMM or require >99KB SMEM (topk_v2).
|
||||
|
||||
Reference in New Issue
Block a user