[Diffusion] Avoid slow cuBLASLt GELU epilogue on SM120 (#34350)

This commit is contained in:
Xiaoyu Zhang
2026-08-12 12:12:00 +08:00
committed by GitHub
parent a9a355774a
commit 22e4b3a81f
@@ -31,6 +31,7 @@ from typing import Any
import torch
import torch.nn as nn
from sglang.kernels.jit.utils import get_jit_cuda_arch
from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion
from sglang.srt.utils.custom_op import register_custom_op
@@ -131,6 +132,13 @@ def can_fuse_linear_gelu(linear: Any, x: torch.Tensor) -> bool:
"""Whether ``gelu(linear(x))`` can use the fused cublasLt epilogue now."""
if not (x.is_cuda and x.dtype in (torch.bfloat16, torch.float16)):
return False
arch = get_jit_cuda_arch()
if arch.major * 10 + arch.minor >= 120:
# The cublasLt GELU epilogue selected by current SM120 PyTorch/CUDA
# builds is slower than GEMM + the native GELU kernel for the FLUX
# production shape (1, 512, 3072 -> 12288). Keep quality-gated sites
# on their existing eager path on RTX 5090; SM90 dispatch is unchanged.
return False
if getattr(linear, "weight", None) is None or x.dtype != linear.weight.dtype:
return False
return can_fuse_linear_gelu_static(linear)