[diffusion] Fuse DiT FFN tanh-GELU into up-proj GEMM (cublasLt epilogue) behind quality=high (Qwen-Image 1024^2 denoise 12.36 -> 12.05 s on H200) (#33536)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Xiaoyu Zhang
2026-08-04 23:49:43 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 0d0c7d853f
commit 95d0e57e83
6 changed files with 317 additions and 0 deletions
@@ -42,6 +42,18 @@ register_kernel(
description="Fused residual gate-add (sglang.kernels.jit).",
)
)
register_kernel(
KernelSpec(
op="diffusion.fused_linear_gelu_tanh",
backend=KernelBackend.TORCH,
target="sglang.kernels.ops.diffusion.fused_linear_gelu:fused_linear_gelu_tanh",
capabilities=_CUDA,
format_signature=FormatSignature(
description="linear + tanh-GELU via the cublasLt GELU epilogue"
),
description="Fused up-proj GEMM + tanh-GELU (torch._addmm_activation).",
)
)
register_kernel(
KernelSpec(
op="diffusion.fused_inplace_qknorm_rope",
@@ -0,0 +1,182 @@
"""Fused linear + tanh-GELU via the cublasLt GELU epilogue, gated by quality.
Many diffusion DiT FeedForwards compute ``gelu(linear(x), approximate="tanh")``
as a standalone up-projection GEMM followed by a separate, bandwidth-bound GELU
kernel over the ``[tokens, 4*dim]`` MLP intermediate. ``torch._addmm_activation``
folds the bias-add and GELU into the GEMM epilogue (cublasLt), removing the
extra kernel launch and the intermediate HBM round-trip. cublasLt's GELU is the
tanh-approximate GELU (max abs diff ~5e-6 vs ``F.gelu(approximate="tanh")`` in
fp32), so for half-precision inference the fused path differs from the
reference only at bf16/fp16 rounding-order level -- close, but not bit-exact.
Because it is not bit-exact, the fused path is **mounted only for
``quality="high"`` requests** (see ``SamplingParams.quality``): model code marks
its GELU up-projection sites with :func:`mark_fused_gelu_site` (default: off,
reference path, bit-exact), and the denoising stage calls
:func:`mount_fused_linear_gelu` / :func:`unmount_fused_linear_gelu` at batch
boundaries. Mounting is all-or-nothing per transformer: if any marked site
fails the static guards (quantized weights, missing bias, non-half dtype, ...)
no site on that transformer is fused.
The fused GEMM is exposed as a registered custom op (``register_custom_op``)
exactly like the other diffusion kernels (e.g. qknorm_rope), so it stays a
single opaque op under ``torch.compile`` -- no graph break.
"""
from __future__ import annotations
import logging
from typing import Any, Iterator
import torch
import torch.nn as nn
from sglang.srt.utils.custom_op import register_custom_op
logger = logging.getLogger(__name__)
# ``torch._addmm_activation`` is the (private but stable) entry point to the
# cublasLt GEMM+bias+activation epilogue. Guard for builds where it is absent
# so the reference path is always available.
_HAS_ADDMM_ACTIVATION = hasattr(torch, "_addmm_activation")
# Attributes of the site protocol (set by ``mark_fused_gelu_site``).
_SITE_LINEAR_ATTR = "_sgl_fused_gelu_linear_attr"
_SITE_ENABLED_ATTR = "_sgl_fused_gelu_enabled"
def _fused_linear_gelu_tanh_fake(
x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor
) -> torch.Tensor:
return x.new_empty((*x.shape[:-1], weight.shape[0]))
@register_custom_op(
op_name="diffusion_fused_linear_gelu_tanh",
mutates_args=[],
fake_impl=_fused_linear_gelu_tanh_fake,
)
def fused_linear_gelu_tanh(
x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor
) -> torch.Tensor:
"""``gelu_tanh(x @ weight.T + bias)`` fused in the cublasLt GELU epilogue.
``weight`` is ``[out, in]`` (nn.Linear / sglang linear layout). Registered
as a custom op so it is opaque under torch.compile.
"""
x2d = x.reshape(-1, x.shape[-1])
out = torch._addmm_activation(bias, x2d, weight.t(), use_gelu=True)
return out.view(*x.shape[:-1], weight.shape[0])
def _is_unquantized(linear: Any) -> bool:
"""True iff ``linear`` carries plain, unquantized weights."""
# Quantized checkpoints can leave selected layers unquantized via their
# exclude list; keep every layer of a quantized model on the reference
# path (all-or-nothing would reject the model anyway).
if getattr(linear, "quant_config", None) is not None:
return False
quant_method = getattr(linear, "quant_method", None)
if quant_method is None:
# Plain nn.Linear has no quant_method.
return True
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod
return isinstance(quant_method, UnquantizedLinearMethod)
def _static_reject_reason(linear: Any) -> str | None:
"""Why ``linear`` may never use the epilogue, or None if it may.
Input-independent guards: requires the fused API and an unquantized,
bias'd, non-bias-deferring linear with a half-precision 2D weight. A
column-parallel layer that gathers across multiple ranks is excluded (the
per-shard fused op would skip the cross-rank gather);
``gather_output=False`` sharded layers are fine because the local shard is
exactly what the reference forward multiplies. The weight's *device* is
deliberately not checked here -- under CPU offload the weights live on CPU
between requests -- the runtime guard checks the input device per call.
"""
if not _HAS_ADDMM_ACTIVATION:
return "torch._addmm_activation unavailable"
if not _is_unquantized(linear):
return "quantized linear"
if getattr(linear, "skip_bias_add", False):
return "skip_bias_add (bias returned separately)"
if getattr(linear, "gather_output", False) and getattr(linear, "tp_size", 1) > 1:
return "multi-rank gather_output"
weight = getattr(linear, "weight", None)
bias = getattr(linear, "bias", None)
if weight is None or bias is None or weight.dim() != 2:
return "missing bias or non-2D weight"
if weight.dtype not in (torch.bfloat16, torch.float16):
return f"non-half weight dtype {weight.dtype}"
if bias.dtype != weight.dtype:
return f"bias dtype {bias.dtype} != weight dtype {weight.dtype}"
return None
def can_fuse_linear_gelu_static(linear: Any) -> bool:
"""Input-independent guards: whether ``linear`` may ever use the epilogue."""
return _static_reject_reason(linear) is None
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
if getattr(linear, "weight", None) is None or x.dtype != linear.weight.dtype:
return False
return can_fuse_linear_gelu_static(linear)
def mark_fused_gelu_site(module: nn.Module, linear_attr: str) -> None:
"""Declare ``module`` as a tanh-GELU up-projection fusion site.
``getattr(module, linear_attr)`` must be the up-projection linear whose
output feeds ``F.gelu(..., approximate="tanh")``. The site starts unmounted
(``_sgl_fused_gelu_enabled = False``): the module's forward must keep the
reference path bit-exact until :func:`mount_fused_linear_gelu` enables it.
"""
setattr(module, _SITE_LINEAR_ATTR, linear_attr)
setattr(module, _SITE_ENABLED_ATTR, False)
def iter_fused_gelu_sites(root: nn.Module) -> Iterator[nn.Module]:
"""Yield every marked fusion site under ``root`` (including ``root``)."""
for module in root.modules():
if getattr(module, _SITE_LINEAR_ATTR, None) is not None:
yield module
def mount_fused_linear_gelu(root: nn.Module) -> bool:
"""Enable the fused epilogue on every marked site under ``root``.
All-or-nothing: if any marked site fails the static guards, every site is
left (or reset) on the reference path and False is returned. Returns False
as well when ``root`` has no marked sites.
"""
sites = list(iter_fused_gelu_sites(root))
if not sites:
return False
for site in sites:
linear = getattr(site, getattr(site, _SITE_LINEAR_ATTR), None)
reason = "missing linear" if linear is None else _static_reject_reason(linear)
if reason is not None:
unmount_fused_linear_gelu(root)
logger.info(
"fused linear+GELU: %s site failed static guards (%s); "
"keeping the whole model on the reference path",
type(site).__name__,
reason,
)
return False
for site in sites:
setattr(site, _SITE_ENABLED_ATTR, True)
return True
def unmount_fused_linear_gelu(root: nn.Module) -> None:
"""Reset every marked site under ``root`` to the bit-exact reference path."""
for site in iter_fused_gelu_sites(root):
setattr(site, _SITE_ENABLED_ATTR, False)
@@ -18,6 +18,11 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from sglang.kernels.ops.diffusion.fused_linear_gelu import (
can_fuse_linear_gelu,
fused_linear_gelu_tanh,
mark_fused_gelu_site,
)
from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_sp_parallel_rank,
@@ -330,8 +335,17 @@ class GlmImageGELU(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.proj" if prefix else "proj",
)
# quality="high" fusion site: up-proj GEMM + tanh-GELU in the cublasLt
# epilogue. Off by default; mounted per batch by the denoising stage.
mark_fused_gelu_site(self, "proj")
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if self._sgl_fused_gelu_enabled and can_fuse_linear_gelu(
self.proj, hidden_states
):
return fused_linear_gelu_tanh(
hidden_states, self.proj.weight, self.proj.bias
)
hidden_states, _ = self.proj(hidden_states)
return F.gelu(hidden_states, approximate="tanh")
@@ -14,6 +14,11 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.kernels.ops.diffusion.fused_linear_gelu import (
can_fuse_linear_gelu,
fused_linear_gelu_tanh,
mark_fused_gelu_site,
)
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
from sglang.multimodal_gen.runtime.distributed import (
get_local_torch_device,
@@ -851,8 +856,17 @@ class QwenImageGELU(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.proj",
)
# quality="high" fusion site: up-proj GEMM + tanh-GELU in the cublasLt
# epilogue. Off by default; mounted per batch by the denoising stage.
mark_fused_gelu_site(self, "proj")
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if self._sgl_fused_gelu_enabled and can_fuse_linear_gelu(
self.proj, hidden_states
):
return fused_linear_gelu_tanh(
hidden_states, self.proj.weight, self.proj.bias
)
hidden_states, _ = self.proj(hidden_states)
return F.gelu(hidden_states, approximate="tanh")
@@ -20,6 +20,10 @@ from typing import Any
import torch
import torch.nn as nn
from sglang.kernels.ops.diffusion.fused_linear_gelu import (
mount_fused_linear_gelu,
unmount_fused_linear_gelu,
)
from sglang.multimodal_gen import envs
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
from sglang.multimodal_gen.configs.pipeline_configs.flux import (
@@ -222,6 +226,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# cache-dit state (for delayed mounting and idempotent control)
self._cache_dit_enabled = False
self._cached_num_steps = None
# fused linear+GELU state: whether the cublasLt-epilogue fusion is
# currently mounted on the transformers (quality="high" batches only).
self._fused_gelu_mounted = False
self._torch_compile_registry = CompiledModuleRegistry()
# Breakable CUDA graph runners, one per transformer module (lazy).
self._bcg_runners: dict[int, Any] = {}
@@ -443,10 +450,39 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
self, num_inference_steps: int | tuple[int, int], batch: Req
) -> None:
"""Apply request-dependent transformer acceleration in trace-safe order."""
self._maybe_toggle_fused_gelu(batch)
self._maybe_enable_cache_dit(num_inference_steps, batch)
for transformer in filter(None, [self.transformer, self.transformer_2]):
self._maybe_torch_compile(transformer)
def _maybe_toggle_fused_gelu(self, batch: Req) -> None:
"""Mount/unmount the cublasLt linear+GELU fusion for this batch.
The fused epilogue is numerically equivalent only at half-precision
rounding level (not bit-exact), so it is mounted for
``quality="high"`` requests and unmounted otherwise -- the
``"lossless"`` default runs the unmodified reference path bit-for-bit.
``quality`` participates in the dynamic-batch signature, so a worker
batch is uniform in ``quality`` and this process-wide transition is
safe at the batch boundary. Mounting is all-or-nothing per
transformer (any ineligible marked site keeps the whole transformer
on the reference path); models without marked sites are no-ops.
"""
want = getattr(batch.sampling_params, "quality", "lossless") == "high"
if want == self._fused_gelu_mounted:
return
mounted = False
for transformer in filter(None, [self.transformer, self.transformer_2]):
if want:
mounted |= mount_fused_linear_gelu(transformer)
else:
unmount_fused_linear_gelu(transformer)
self._fused_gelu_mounted = want
if want and mounted:
logger.info(
"Mounted fused linear+GELU (cublasLt epilogue) for quality=high"
)
def _cache_dit_dual_model_name(self) -> str:
return "wan2.2"