[Diffusion][FLUX.2] Fuse eager AdaLN and packed SwiGLU (#34616)

This commit is contained in:
Xiaoyu Zhang
2026-08-13 19:55:53 +08:00
committed by GitHub
parent 3c1791a7df
commit 74c0322342
3 changed files with 293 additions and 20 deletions
@@ -40,6 +40,26 @@ def _silu_mul_kernel(
tl.store(out_ptr + offs, s * b, mask=mask) # store rounds the multiply tl.store(out_ptr + offs, s * b, mask=mask) # store rounds the multiply
@triton.jit
def _packed_silu_mul_kernel(
out_ptr,
x_ptr,
num_rows,
row_stride,
D: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
block = tl.program_id(1).to(tl.int64)
cols = block * BLOCK + tl.arange(0, BLOCK)
mask = (row < num_rows) & (cols < D)
row_base = row * row_stride
a = tl.load(x_ptr + row_base + cols, mask=mask, other=0.0).to(tl.float32)
b = tl.load(x_ptr + row_base + D + cols, mask=mask, other=0.0).to(tl.float32)
s = round_bf16_to_fp32(a * tl.sigmoid(a))
tl.store(out_ptr + row * D + cols, s * b, mask=mask)
def can_use_fused_silu_mul(a: torch.Tensor, b: torch.Tensor) -> bool: def can_use_fused_silu_mul(a: torch.Tensor, b: torch.Tensor) -> bool:
return ( return (
a.dtype is torch.bfloat16 a.dtype is torch.bfloat16
@@ -76,3 +96,32 @@ def fused_silu_mul_bitexact(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
BLOCK=1024, BLOCK=1024,
) )
return out return out
def fused_packed_silu_mul_bitexact(x: torch.Tensor) -> torch.Tensor:
"""Bit-exact SwiGLU over a contiguous packed ``[..., 2 * D]`` input."""
if not (
x.is_cuda
and x.dtype is torch.bfloat16
and x.dim() == 3
and x.stride(-1) == 1
and x.stride(-2) >= x.shape[-1]
and x.stride(0) == x.shape[1] * x.stride(1)
and x.shape[-1] % 2 == 0
and x.numel() > 0
):
raise RuntimeError("unsupported input for packed fused SiLU-mul")
hidden = x.shape[-1] // 2
rows = x.numel() // x.shape[-1]
row_stride = x.stride(-2)
out = torch.empty((*x.shape[:-1], hidden), dtype=x.dtype, device=x.device)
with torch.cuda.device(x.device):
_packed_silu_mul_kernel[(rows, triton.cdiv(hidden, 1024))](
out,
x,
rows,
row_stride,
D=hidden,
BLOCK=1024,
)
return out
@@ -16,11 +16,21 @@ from typing import Any, Dict, List, Optional, Tuple
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F
from diffusers.models.attention import AttentionModuleMixin from diffusers.models.attention import AttentionModuleMixin
from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.normalization import AdaLayerNormContinuous from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate
from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add
from sglang.kernels.ops.diffusion.triton.layernorm_modulate import (
can_use_fused_layernorm_modulate,
fused_layernorm_modulate_raw,
is_plain_layer_norm,
)
from sglang.kernels.ops.diffusion.triton.silu_mul_bitexact import (
fused_packed_silu_mul_bitexact,
)
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
divide, divide,
@@ -69,6 +79,118 @@ logger = init_logger(__name__) # pylint: disable=invalid-name
_get_qkv_projections = get_qkv_projections _get_qkv_projections = get_qkv_projections
_FLUX2_LN_MOD = BitExactFusionGate("FLUX.2 fused LN+modulate", per_signature=True)
_FLUX2_LN_MOD_SIGS = _FLUX2_LN_MOD.verified_sigs
assert _FLUX2_LN_MOD_SIGS is not None
_FLUX2_SWIGLU = BitExactFusionGate("FLUX.2 fused SwiGLU", per_signature=True)
_FLUX2_SWIGLU_SIGS = _FLUX2_SWIGLU.verified_sigs
assert _FLUX2_SWIGLU_SIGS is not None
def _flux2_norm_modulate(
norm: nn.Module,
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
"""Bit-exact single-kernel ``LN(x) * (1 + scale) + shift``."""
# Preserve the original expression for Dynamo/Inductor. This direct
# Triton dispatch is intentionally an eager fast path.
if torch.compiler.is_compiling():
return norm(x) * (1 + scale) + shift
scale_row = scale.squeeze(1) if scale.dim() == 3 and scale.shape[1] == 1 else scale
shift_row = shift.squeeze(1) if shift.dim() == 3 and shift.shape[1] == 1 else shift
if (
_FLUX2_LN_MOD.disabled
or not is_plain_layer_norm(norm, x.shape[-1])
or not can_use_fused_layernorm_modulate(x, scale_row, shift_row)
):
return norm(x) * (1 + scale) + shift
# The bit-exact contract is set by dtype/reduction width/affine-row
# layout, not by the number of independent rows. Excluding sequence
# length lets the representative warmup verify the real prompt path too.
sig = (
x.dtype,
x.device,
x.shape[0],
x.shape[-1],
x.stride(-1),
scale_row.stride(0) if scale_row.shape[0] > 1 else x.shape[-1],
shift_row.stride(0) if shift_row.shape[0] > 1 else x.shape[-1],
norm.eps,
)
verified = sig in _FLUX2_LN_MOD_SIGS
if not verified and torch.cuda.is_current_stream_capturing():
return norm(x) * (1 + scale) + shift
try:
# Direct dispatch avoids custom-op overhead on this eager-only path.
out = fused_layernorm_modulate_raw(x, scale_row, shift_row, norm.eps)
except Exception as exc:
_FLUX2_LN_MOD.on_exception(exc, logger=logger)
return norm(x) * (1 + scale) + shift
if verified:
return out
ref = norm(x) * (1 + scale) + shift
return _FLUX2_LN_MOD.accept_or_fallback(
out,
ref,
sig=sig,
logger=logger,
mismatch_msg=(
"FLUX.2 fused LN+modulate fast path is not bit-exact on this "
"platform; falling back to eager"
),
)
def _flux2_swiglu(x: torch.Tensor) -> torch.Tensor:
"""Bit-exact fused SwiGLU for the packed FLUX.2 FFN projection."""
half = x.shape[-1] // 2
# Let Inductor fuse the reference expression in torch.compile mode.
if torch.compiler.is_compiling():
return F.silu(x[..., :half]) * x[..., half:]
# Sequence length only changes the launch grid; D and row stride define
# how the two packed halves are addressed and therefore need verification.
sig = (x.dtype, x.device, x.shape[0], x.shape[-1], x.stride(-2), x.stride(-1))
verified = sig in _FLUX2_SWIGLU_SIGS
can_fuse = (
not _FLUX2_SWIGLU.disabled
and x.is_cuda
and x.dtype is torch.bfloat16
and x.dim() == 3
and x.stride(-1) == 1
and x.stride(-2) >= x.shape[-1]
and x.stride(0) == x.shape[1] * x.stride(1)
and x.shape[-1] % 2 == 0
and x.numel() > 0
)
# Per-signature verification may compare tensors and synchronize. Never
# verify a new layout while a CUDA graph is being captured.
if can_fuse and not verified and torch.cuda.is_current_stream_capturing():
return F.silu(x[..., :half]) * x[..., half:]
if can_fuse:
try:
out = fused_packed_silu_mul_bitexact(x)
except Exception as exc:
_FLUX2_SWIGLU.on_exception(exc, logger=logger)
else:
if verified:
return out
return _FLUX2_SWIGLU.accept_or_fallback(
out,
F.silu(x[..., :half]) * x[..., half:],
sig=sig,
logger=logger,
mismatch_msg=(
"FLUX.2 fused SwiGLU fast path is not bit-exact on this "
"platform; falling back to eager"
),
)
return F.silu(x[..., :half]) * x[..., half:]
class Flux2SwiGLU(nn.Module): class Flux2SwiGLU(nn.Module):
""" """
@@ -76,14 +198,8 @@ class Flux2SwiGLU(nn.Module):
layer fused into the first linear layer of the FF sub-block. Thus, this module has no trainable parameters. layer fused into the first linear layer of the FF sub-block. Thus, this module has no trainable parameters.
""" """
def __init__(self):
super().__init__()
self.gate_fn = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
x1, x2 = x.chunk(2, dim=-1) return _flux2_swiglu(x)
x = self.gate_fn(x1) * x2
return x
class Flux2FeedForward(nn.Module): class Flux2FeedForward(nn.Module):
@@ -623,8 +739,9 @@ class Flux2SingleTransformerBlock(nn.Module):
mod_shift, mod_scale, mod_gate = temb_mod_params mod_shift, mod_scale, mod_gate = temb_mod_params
norm_hidden_states = self.norm(hidden_states) norm_hidden_states = _flux2_norm_modulate(
norm_hidden_states = (1 + mod_scale) * norm_hidden_states + mod_shift self.norm, hidden_states, mod_scale, mod_shift
)
joint_attention_kwargs = joint_attention_kwargs or {} joint_attention_kwargs = joint_attention_kwargs or {}
attn_output = self.attn( attn_output = self.attn(
@@ -737,14 +854,17 @@ class Flux2TransformerBlock(nn.Module):
) = temb_mod_params_txt ) = temb_mod_params_txt
# Img stream # Img stream
norm_hidden_states = self.norm1(hidden_states) norm_hidden_states = _flux2_norm_modulate(
norm_hidden_states = (1 + scale_msa) * norm_hidden_states + shift_msa self.norm1, hidden_states, scale_msa, shift_msa
)
# Conditioning txt stream # Conditioning txt stream
norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states) norm_encoder_hidden_states = _flux2_norm_modulate(
norm_encoder_hidden_states = ( self.norm1_context,
1 + c_scale_msa encoder_hidden_states,
) * norm_encoder_hidden_states + c_shift_msa c_scale_msa,
c_shift_msa,
)
# Attention on concatenated img + txt stream # Attention on concatenated img + txt stream
attention_outputs = self.attn( attention_outputs = self.attn(
@@ -760,8 +880,9 @@ class Flux2TransformerBlock(nn.Module):
# Process attention outputs for the image stream (`hidden_states`). # Process attention outputs for the image stream (`hidden_states`).
hidden_states = residual_gate_add(hidden_states, attn_output, gate_msa) hidden_states = residual_gate_add(hidden_states, attn_output, gate_msa)
norm_hidden_states = self.norm2(hidden_states) norm_hidden_states = _flux2_norm_modulate(
norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp self.norm2, hidden_states, scale_mlp, shift_mlp
)
ff_output = self.ff(norm_hidden_states) ff_output = self.ff(norm_hidden_states)
hidden_states = residual_gate_add(hidden_states, ff_output, gate_mlp) hidden_states = residual_gate_add(hidden_states, ff_output, gate_mlp)
@@ -771,9 +892,11 @@ class Flux2TransformerBlock(nn.Module):
encoder_hidden_states, context_attn_output, c_gate_msa encoder_hidden_states, context_attn_output, c_gate_msa
) )
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) norm_encoder_hidden_states = _flux2_norm_modulate(
norm_encoder_hidden_states = ( self.norm2_context,
norm_encoder_hidden_states * (1 + c_scale_mlp) + c_shift_mlp encoder_hidden_states,
c_scale_mlp,
c_shift_mlp,
) )
context_ff_output = self.ff_context(norm_encoder_hidden_states) context_ff_output = self.ff_context(norm_encoder_hidden_states)
@@ -0,0 +1,101 @@
"""FLUX.2 eager fusions must be bit-exact for real packed/view layouts."""
import unittest
from unittest.mock import patch
import torch
import torch.nn.functional as F
import sglang.multimodal_gen.runtime.models.dits.flux_2 as flux2
from sglang.multimodal_gen.runtime.models.dits.flux_2 import (
_flux2_norm_modulate,
_flux2_swiglu,
)
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=12, stage="base-b-kernel-unit", runner_config="1-gpu-large")
@unittest.skipUnless(torch.cuda.is_available(), "CUDA required")
class TestFlux2EagerFusions(CustomTestCase):
def setUp(self):
flux2._FLUX2_LN_MOD.disabled = False
flux2._FLUX2_LN_MOD.verified = False
flux2._FLUX2_LN_MOD_SIGS.clear()
flux2._FLUX2_SWIGLU.disabled = False
flux2._FLUX2_SWIGLU.verified = False
flux2._FLUX2_SWIGLU_SIGS.clear()
def test_norm_modulate_is_bit_exact_across_sequence_lengths(self):
torch.manual_seed(0)
hidden = 256
norm = torch.nn.LayerNorm(
hidden, eps=1e-6, elementwise_affine=False, device="cuda"
)
# FLUX.2 modulation values are views of one packed projection.
params = torch.randn(1, 1, 6 * hidden, device="cuda").bfloat16()
shift, scale = params.chunk(6, dim=-1)[:2]
for seq in (17, 65):
x = torch.randn(1, seq, hidden, device="cuda").bfloat16()
expected = norm(x) * (1 + scale) + shift
actual = _flux2_norm_modulate(norm, x, scale, shift)
self.assertTrue(torch.equal(actual, expected))
self.assertFalse(flux2._FLUX2_LN_MOD.disabled)
self.assertEqual(len(flux2._FLUX2_LN_MOD_SIGS), 1)
def test_packed_swiglu_is_bit_exact_for_contiguous_and_strided_views(self):
torch.manual_seed(1)
hidden = 384
inputs = [
torch.randn(1, 19, 2 * hidden, device="cuda").bfloat16(),
torch.randn(1, 19, 3 * hidden, device="cuda").bfloat16()[..., : 2 * hidden],
]
for x in inputs:
expected = F.silu(x[..., :hidden]) * x[..., hidden:]
actual = _flux2_swiglu(x)
self.assertTrue(torch.equal(actual, expected))
self.assertFalse(flux2._FLUX2_SWIGLU.disabled)
self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 2)
def test_fp16_preserves_reference_path(self):
x = torch.randn(1, 17, 512, device="cuda", dtype=torch.float16)
expected = F.silu(x[..., :256]) * x[..., 256:]
actual = _flux2_swiglu(x)
self.assertTrue(torch.equal(actual, expected))
self.assertFalse(flux2._FLUX2_SWIGLU.disabled)
def test_packed_swiglu_rejects_non_dense_outer_stride(self):
base = torch.randn(2, 23, 512, device="cuda", dtype=torch.bfloat16)
x = base[:, :19]
self.assertNotEqual(x.stride(0), x.shape[1] * x.stride(1))
expected = F.silu(x[..., :256]) * x[..., 256:]
actual = _flux2_swiglu(x)
self.assertTrue(torch.equal(actual, expected))
self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 0)
def test_new_swiglu_signature_is_not_verified_during_graph_capture(self):
first = torch.randn(1, 17, 512, device="cuda", dtype=torch.bfloat16)
self.assertTrue(
torch.equal(
_flux2_swiglu(first),
F.silu(first[..., :256]) * first[..., 256:],
)
)
self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1)
second = torch.randn(1, 19, 768, device="cuda", dtype=torch.bfloat16)
with patch("torch.cuda.is_current_stream_capturing", return_value=True):
actual = _flux2_swiglu(second)
expected = F.silu(second[..., :384]) * second[..., 384:]
self.assertTrue(torch.equal(actual, expected))
self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1)
if __name__ == "__main__":
unittest.main()