fix: default FP4 GEMM backend to flashinfer_cudnn on SM120 (Blackwell) (#20047)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com> Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
This commit is contained in:
co-authored by
Claude Opus 4.6
Brayden Zhong
parent
61d530e8ac
commit
d39ed074cf
@@ -5,6 +5,7 @@ from enum import Enum
|
|||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils.common import is_sm120_supported
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -75,6 +76,19 @@ def initialize_fp4_gemm_config(server_args: ServerArgs) -> None:
|
|||||||
"Using server argument value."
|
"Using server argument value."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if backend == "auto":
|
||||||
|
if is_sm120_supported():
|
||||||
|
# flashinfer_cutlass produces NaN in dense MLP layers with
|
||||||
|
# heterogeneous batches on SM120 (Blackwell). cudnn is stable.
|
||||||
|
# See: https://github.com/sgl-project/sglang/issues/20043
|
||||||
|
backend = "flashinfer_cudnn"
|
||||||
|
logger.info(
|
||||||
|
"SM120 (Blackwell) detected: auto-selecting "
|
||||||
|
"fp4-gemm-backend=flashinfer_cudnn"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
backend = "flashinfer_cutlass"
|
||||||
|
|
||||||
FP4_GEMM_RUNNER_BACKEND = Fp4GemmRunnerBackend(backend)
|
FP4_GEMM_RUNNER_BACKEND = Fp4GemmRunnerBackend(backend)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -466,7 +466,7 @@ class ServerArgs:
|
|||||||
grammar_backend: Optional[str] = None
|
grammar_backend: Optional[str] = None
|
||||||
mm_attention_backend: Optional[str] = None
|
mm_attention_backend: Optional[str] = None
|
||||||
fp8_gemm_runner_backend: str = "auto"
|
fp8_gemm_runner_backend: str = "auto"
|
||||||
fp4_gemm_runner_backend: str = "flashinfer_cutlass"
|
fp4_gemm_runner_backend: str = "auto"
|
||||||
nsa_prefill_backend: Optional[str] = (
|
nsa_prefill_backend: Optional[str] = (
|
||||||
None # None = auto-detect based on hardware/kv_cache_dtype
|
None # None = auto-detect based on hardware/kv_cache_dtype
|
||||||
)
|
)
|
||||||
@@ -4308,8 +4308,8 @@ class ServerArgs:
|
|||||||
default=ServerArgs.fp4_gemm_runner_backend,
|
default=ServerArgs.fp4_gemm_runner_backend,
|
||||||
dest="fp4_gemm_runner_backend",
|
dest="fp4_gemm_runner_backend",
|
||||||
help="Choose the runner backend for NVFP4 GEMM operations. "
|
help="Choose the runner backend for NVFP4 GEMM operations. "
|
||||||
"Options: 'flashinfer_cutlass' (default), "
|
"Options: 'auto' (default; selects flashinfer_cudnn on SM120, flashinfer_cutlass otherwise), "
|
||||||
"'auto' (auto-selects between flashinfer_cudnn/flashinfer_cutlass based on CUDA/cuDNN version), "
|
"'flashinfer_cutlass' (CUTLASS backend), "
|
||||||
"'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), "
|
"'flashinfer_cudnn' (FlashInfer cuDNN backend, optimal on CUDA 13+ with cuDNN 9.15+), "
|
||||||
"'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling). "
|
"'flashinfer_trtllm' (FlashInfer TensorRT-LLM backend, requires different weight preparation with shuffling). "
|
||||||
"NOTE: This replaces the deprecated environment variable "
|
"NOTE: This replaces the deprecated environment variable "
|
||||||
|
|||||||
Reference in New Issue
Block a user