[diffusion] FLUX.1 bit-exact residual-gate fast path + tanh-GELU epilogue behind quality=high (H200 e2e -1.1% lossless / -4.3% high) (#33819)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
f8f2870a84
commit
eff6a11350
@@ -18,6 +18,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union
|
|||||||
|
|
||||||
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.modeling_outputs import Transformer2DModelOutput
|
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||||
from diffusers.models.normalization import (
|
from diffusers.models.normalization import (
|
||||||
@@ -27,6 +28,15 @@ from diffusers.models.normalization import (
|
|||||||
)
|
)
|
||||||
from torch.nn import LayerNorm as LayerNorm
|
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.configs.models.dits.flux import FluxConfig
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
divide,
|
divide,
|
||||||
@@ -76,6 +86,40 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
|
|
||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
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:
|
try:
|
||||||
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
|
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",
|
prefix=f"{prefix}.proj" if prefix else "proj",
|
||||||
)
|
)
|
||||||
self.gelu = nn.GELU(approximate="tanh")
|
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:
|
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)
|
hidden_states, _ = self.proj(hidden_states)
|
||||||
return self.gelu(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):
|
class FluxParallelFeedForward(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -625,6 +703,9 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
prefix=f"{prefix}.proj_mlp" if prefix else "proj_mlp",
|
prefix=f"{prefix}.proj_mlp" if prefix else "proj_mlp",
|
||||||
)
|
)
|
||||||
self.act_mlp = nn.GELU(approximate="tanh")
|
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 = (
|
proj_out_cls = (
|
||||||
RowParallelLinear if shard_single_block else ColumnParallelLinear
|
RowParallelLinear if shard_single_block else ColumnParallelLinear
|
||||||
)
|
)
|
||||||
@@ -724,8 +805,15 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
hidden_states = gate * hidden_states
|
hidden_states = gate * hidden_states
|
||||||
hidden_states = residual + hidden_states
|
hidden_states = residual + hidden_states
|
||||||
else:
|
else:
|
||||||
proj_hidden_states, _ = self.proj_mlp(norm_hidden_states)
|
if self._sgl_fused_gelu_enabled and can_fuse_linear_gelu(
|
||||||
mlp_hidden_states = self.act_mlp(proj_hidden_states)
|
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(
|
attn_output = self.attn(
|
||||||
x=norm_hidden_states,
|
x=norm_hidden_states,
|
||||||
@@ -737,8 +825,7 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||||
gate = gate.unsqueeze(1)
|
gate = gate.unsqueeze(1)
|
||||||
proj_out, _ = self.proj_out(hidden_states)
|
proj_out, _ = self.proj_out(hidden_states)
|
||||||
hidden_states = gate * proj_out
|
hidden_states = _flux_residual_gate_add(residual, proj_out, gate)
|
||||||
hidden_states = residual + hidden_states
|
|
||||||
|
|
||||||
if hidden_states.dtype == torch.float16:
|
if hidden_states.dtype == torch.float16:
|
||||||
hidden_states = hidden_states.clip(-65504, 65504)
|
hidden_states = hidden_states.clip(-65504, 65504)
|
||||||
@@ -831,6 +918,10 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
dim_out=dim,
|
dim_out=dim,
|
||||||
activation_fn="gelu-approximate",
|
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -869,8 +960,9 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
attn_output, context_attn_output, ip_attn_output = attention_outputs
|
attn_output, context_attn_output, ip_attn_output = attention_outputs
|
||||||
|
|
||||||
# Process attention outputs for the `hidden_states`.
|
# Process attention outputs for the `hidden_states`.
|
||||||
attn_output = gate_msa.unsqueeze(1) * attn_output
|
hidden_states = _flux_residual_gate_add(
|
||||||
hidden_states = hidden_states + attn_output
|
hidden_states, attn_output, gate_msa.unsqueeze(1)
|
||||||
|
)
|
||||||
norm_hidden_states = self.norm2(hidden_states)
|
norm_hidden_states = self.norm2(hidden_states)
|
||||||
if self.use_nunchaku_structure:
|
if self.use_nunchaku_structure:
|
||||||
norm_hidden_states = (
|
norm_hidden_states = (
|
||||||
@@ -882,15 +974,16 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
ff_output = self.ff(norm_hidden_states)
|
ff_output = self.ff(norm_hidden_states)
|
||||||
ff_output = gate_mlp.unsqueeze(1) * ff_output
|
hidden_states = _flux_residual_gate_add(
|
||||||
|
hidden_states, ff_output, gate_mlp.unsqueeze(1)
|
||||||
hidden_states = hidden_states + ff_output
|
)
|
||||||
|
|
||||||
if len(attention_outputs) == 3:
|
if len(attention_outputs) == 3:
|
||||||
hidden_states = hidden_states + ip_attn_output
|
hidden_states = hidden_states + ip_attn_output
|
||||||
# Process attention outputs for the `encoder_hidden_states`.
|
# Process attention outputs for the `encoder_hidden_states`.
|
||||||
context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output
|
encoder_hidden_states = _flux_residual_gate_add(
|
||||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
encoder_hidden_states, context_attn_output, c_gate_msa.unsqueeze(1)
|
||||||
|
)
|
||||||
|
|
||||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||||
if self.use_nunchaku_structure:
|
if self.use_nunchaku_structure:
|
||||||
@@ -904,8 +997,8 @@ class FluxTransformerBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||||
encoder_hidden_states = (
|
encoder_hidden_states = _flux_residual_gate_add(
|
||||||
encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output
|
encoder_hidden_states, context_ff_output, c_gate_mlp.unsqueeze(1)
|
||||||
)
|
)
|
||||||
if encoder_hidden_states.dtype == torch.float16:
|
if encoder_hidden_states.dtype == torch.float16:
|
||||||
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
||||||
|
|||||||
@@ -37,6 +37,23 @@ def test_fused_matches_reference(dtype):
|
|||||||
torch.testing.assert_close(site(x), ref, atol=atol, rtol=2e-2)
|
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():
|
def test_mount_guards_and_lossless_path():
|
||||||
torch.manual_seed(0)
|
torch.manual_seed(0)
|
||||||
good, bad = _Site(), _Site(torch.float32)
|
good, bad = _Site(), _Site(torch.float32)
|
||||||
|
|||||||
@@ -18,6 +18,10 @@ CASES = [
|
|||||||
((1, 512, 4096), (1, 512, 4096)),
|
((1, 512, 4096), (1, 512, 4096)),
|
||||||
((1, 17, 65), (1, 1, 65)),
|
((1, 17, 65), (1, 1, 65)),
|
||||||
((1, 17, 65), (1, 17, 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)),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user