[diffusion] FLUX.2 bit-exact residual-gate fast path (H200 klein-4B 50-step denoise -1.2%) (#33823)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
dd98c9572a
commit
591cfb0881
@@ -20,6 +20,10 @@ from diffusers.models.attention import AttentionModuleMixin
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
|
||||
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,
|
||||
@@ -65,6 +69,40 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
_FLUX2_RESIDUAL_GATE_CUDA_DISABLED = False
|
||||
|
||||
|
||||
def _flux2_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 _FLUX2_RESIDUAL_GATE_CUDA_DISABLED
|
||||
|
||||
if (
|
||||
not _FLUX2_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.2 residual-gate CUDA fast path: {exc}")
|
||||
_FLUX2_RESIDUAL_GATE_CUDA_DISABLED = True
|
||||
|
||||
return residual + gate * update
|
||||
|
||||
|
||||
def _get_qkv_projections(
|
||||
attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
|
||||
@@ -656,7 +694,7 @@ class Flux2SingleTransformerBlock(nn.Module):
|
||||
**joint_attention_kwargs,
|
||||
)
|
||||
|
||||
hidden_states = hidden_states + mod_gate * attn_output
|
||||
hidden_states = _flux2_residual_gate_add(hidden_states, attn_output, mod_gate)
|
||||
if hidden_states.dtype == torch.float16:
|
||||
hidden_states = hidden_states.clip(-65504, 65504)
|
||||
|
||||
@@ -774,18 +812,18 @@ class Flux2TransformerBlock(nn.Module):
|
||||
attn_output, context_attn_output = attention_outputs
|
||||
|
||||
# Process attention outputs for the image stream (`hidden_states`).
|
||||
attn_output = gate_msa * attn_output
|
||||
hidden_states = hidden_states + attn_output
|
||||
hidden_states = _flux2_residual_gate_add(hidden_states, attn_output, gate_msa)
|
||||
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp
|
||||
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + gate_mlp * ff_output
|
||||
hidden_states = _flux2_residual_gate_add(hidden_states, ff_output, gate_mlp)
|
||||
|
||||
# Process attention outputs for the text stream (`encoder_hidden_states`).
|
||||
context_attn_output = c_gate_msa * context_attn_output
|
||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
||||
encoder_hidden_states = _flux2_residual_gate_add(
|
||||
encoder_hidden_states, context_attn_output, c_gate_msa
|
||||
)
|
||||
|
||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||
norm_encoder_hidden_states = (
|
||||
@@ -793,7 +831,9 @@ class Flux2TransformerBlock(nn.Module):
|
||||
)
|
||||
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states + c_gate_mlp * context_ff_output
|
||||
encoder_hidden_states = _flux2_residual_gate_add(
|
||||
encoder_hidden_states, context_ff_output, c_gate_mlp
|
||||
)
|
||||
if encoder_hidden_states.dtype == torch.float16:
|
||||
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user