diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index a5423edd5..5aa3158a1 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -18,6 +18,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union import torch import torch.nn as nn +import torch.nn.functional as F from diffusers.models.attention import AttentionModuleMixin from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.normalization import ( @@ -27,6 +28,15 @@ from diffusers.models.normalization import ( ) from torch.nn import LayerNorm as LayerNorm +from sglang.kernels.ops.diffusion.fused_linear_gelu import ( + can_fuse_linear_gelu, + fused_linear_gelu_tanh, + mark_fused_gelu_site, +) +from sglang.kernels.ops.diffusion.residual_gate_add import ( + can_use_residual_gate_add_cuda, + residual_gate_add_cuda, +) from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( divide, @@ -76,6 +86,40 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name +_FLUX_RESIDUAL_GATE_CUDA_DISABLED = False + + +def _flux_residual_gate_add( + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, +) -> torch.Tensor: + """Single-kernel ``residual + gate * update``, bit-exact vs the eager pair. + + Restricted to half dtypes: there the kernel reproduces the eager pair's + two-step rounding exactly (verified by ``torch.equal``), while for fp32 it + would contract to an fma (one rounding) and stop being bit-exact. The + kernel's row-broadcast gate only covers ``[1, ..., 1, D]``; batched + ``[B>1, 1, D]`` gates fail ``can_use_residual_gate_add_cuda`` and take the + eager fallback below. + """ + global _FLUX_RESIDUAL_GATE_CUDA_DISABLED + + if ( + not _FLUX_RESIDUAL_GATE_CUDA_DISABLED + and residual.dtype in (torch.float16, torch.bfloat16) + and can_use_residual_gate_add_cuda(residual, update, gate) + ): + try: + return residual_gate_add_cuda(residual, update, gate) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling FLUX residual-gate CUDA fast path: {exc}") + _FLUX_RESIDUAL_GATE_CUDA_DISABLED = True + + return residual + gate * update + try: from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import] @@ -244,12 +288,46 @@ class FluxGELU(nn.Module): prefix=f"{prefix}.proj" if prefix else "proj", ) self.gelu = nn.GELU(approximate="tanh") + # 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 self.gelu(hidden_states) +class FluxFusedGELUProj(nn.Module): + """tanh-GELU up-projection site for the shared (diffusers-style) FeedForward. + + Drop-in replacement for ``diffusers.models.activations.GELU`` with + ``approximate="tanh"`` that keeps the ``net.0.proj`` parameter path. The + default path is the bit-exact reference (plain Linear + tanh-GELU); the + cublasLt GELU epilogue is mounted per batch by the denoising stage for + quality="high" requests only. + """ + + def __init__(self, proj: nn.Linear): + super().__init__() + self.proj = proj + 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 + ) + return F.gelu(self.proj(hidden_states), approximate="tanh") + + class FluxParallelFeedForward(nn.Module): def __init__( self, @@ -625,6 +703,9 @@ class FluxSingleTransformerBlock(nn.Module): prefix=f"{prefix}.proj_mlp" if prefix else "proj_mlp", ) self.act_mlp = nn.GELU(approximate="tanh") + # quality="high" fusion site: proj_mlp GEMM + tanh-GELU in the + # cublasLt epilogue (mounted per batch by the denoising stage). + mark_fused_gelu_site(self, "proj_mlp") proj_out_cls = ( RowParallelLinear if shard_single_block else ColumnParallelLinear ) @@ -724,8 +805,15 @@ class FluxSingleTransformerBlock(nn.Module): hidden_states = gate * hidden_states hidden_states = residual + hidden_states else: - proj_hidden_states, _ = self.proj_mlp(norm_hidden_states) - mlp_hidden_states = self.act_mlp(proj_hidden_states) + if self._sgl_fused_gelu_enabled and can_fuse_linear_gelu( + self.proj_mlp, norm_hidden_states + ): + mlp_hidden_states = fused_linear_gelu_tanh( + norm_hidden_states, self.proj_mlp.weight, self.proj_mlp.bias + ) + else: + proj_hidden_states, _ = self.proj_mlp(norm_hidden_states) + mlp_hidden_states = self.act_mlp(proj_hidden_states) attn_output = self.attn( x=norm_hidden_states, @@ -737,8 +825,7 @@ class FluxSingleTransformerBlock(nn.Module): hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) gate = gate.unsqueeze(1) proj_out, _ = self.proj_out(hidden_states) - hidden_states = gate * proj_out - hidden_states = residual + hidden_states + hidden_states = _flux_residual_gate_add(residual, proj_out, gate) if hidden_states.dtype == torch.float16: hidden_states = hidden_states.clip(-65504, 65504) @@ -831,6 +918,10 @@ class FluxTransformerBlock(nn.Module): dim_out=dim, activation_fn="gelu-approximate", ) + # Re-home each FF's tanh-GELU up-projection onto a marked + # quality="high" fusion site (bit-exact reference by default). + self.ff.net[0] = FluxFusedGELUProj(self.ff.net[0].proj) + self.ff_context.net[0] = FluxFusedGELUProj(self.ff_context.net[0].proj) def forward( self, @@ -869,8 +960,9 @@ class FluxTransformerBlock(nn.Module): attn_output, context_attn_output, ip_attn_output = attention_outputs # Process attention outputs for the `hidden_states`. - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output + hidden_states = _flux_residual_gate_add( + hidden_states, attn_output, gate_msa.unsqueeze(1) + ) norm_hidden_states = self.norm2(hidden_states) if self.use_nunchaku_structure: norm_hidden_states = ( @@ -882,15 +974,16 @@ class FluxTransformerBlock(nn.Module): ) ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - - hidden_states = hidden_states + ff_output + hidden_states = _flux_residual_gate_add( + hidden_states, ff_output, gate_mlp.unsqueeze(1) + ) if len(attention_outputs) == 3: hidden_states = hidden_states + ip_attn_output # Process attention outputs for the `encoder_hidden_states`. - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output + encoder_hidden_states = _flux_residual_gate_add( + encoder_hidden_states, context_attn_output, c_gate_msa.unsqueeze(1) + ) norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) if self.use_nunchaku_structure: @@ -904,8 +997,8 @@ class FluxTransformerBlock(nn.Module): ) context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = ( - encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output + encoder_hidden_states = _flux_residual_gate_add( + encoder_hidden_states, context_ff_output, c_gate_mlp.unsqueeze(1) ) if encoder_hidden_states.dtype == torch.float16: encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) diff --git a/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py b/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py index df9afd0d8..0976a53d9 100644 --- a/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py +++ b/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py @@ -37,6 +37,23 @@ def test_fused_matches_reference(dtype): torch.testing.assert_close(site(x), ref, atol=atol, rtol=2e-2) +def test_flux_gelu_proj_site(): + """FLUX.1 shared-FF site: gate off is bit-exact, gate on is close.""" + from sglang.multimodal_gen.runtime.models.dits.flux import FluxFusedGELUProj + + torch.manual_seed(0) + proj = nn.Linear(3072, 12288, device="cuda", dtype=torch.bfloat16) + site = FluxFusedGELUProj(proj) + x = torch.randn(1, 512, 3072, device="cuda", dtype=torch.bfloat16) + ref = F.gelu(proj(x), approximate="tanh") + + assert torch.equal(site(x), ref) # unmounted default: bit-exact reference + assert gelu.mount_fused_linear_gelu(site) + torch.testing.assert_close(site(x), ref, atol=2e-2, rtol=2e-2) + gelu.unmount_fused_linear_gelu(site) + assert torch.equal(site(x), ref) + + def test_mount_guards_and_lossless_path(): torch.manual_seed(0) good, bad = _Site(), _Site(torch.float32) diff --git a/test/registered/kernels/ops/diffusion/test_residual_gate_add.py b/test/registered/kernels/ops/diffusion/test_residual_gate_add.py index 935dc4884..81480a31e 100644 --- a/test/registered/kernels/ops/diffusion/test_residual_gate_add.py +++ b/test/registered/kernels/ops/diffusion/test_residual_gate_add.py @@ -18,6 +18,10 @@ CASES = [ ((1, 512, 4096), (1, 512, 4096)), ((1, 17, 65), (1, 1, 65)), ((1, 17, 65), (1, 17, 65)), + # FLUX.1 1024^2 shapes: dual-stream image/text and single-stream joint. + ((1, 4096, 3072), (1, 1, 3072)), + ((1, 512, 3072), (1, 1, 3072)), + ((1, 4608, 3072), (1, 1, 3072)), ]