From 38c007dfe587f1946c84cbffbf728cd87b814b94 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Sun, 9 Aug 2026 09:52:37 +0800 Subject: [PATCH] [diffusion] FLUX.1: route the adaLN LN+modulate sites through the bit-exact fused LayerNorm+modulate kernel (H200 1024^2 lossless denoise -1.2%, e2e wall -2.9%) (#34126) --- .../runtime/models/dits/flux.py | 84 ++++++++++++++++++- .../ops/diffusion/test_flux_ln_modulate.py | 75 +++++++++++++++++ 2 files changed, 156 insertions(+), 3 deletions(-) create mode 100644 test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index 3b6cdcafd..7786b2827 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -42,6 +42,11 @@ from sglang.kernels.ops.diffusion.fused_ln_modulate import ( ) from sglang.kernels.ops.diffusion.modulate_scale_shift import modulate_scale_shift 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, + is_plain_layer_norm, +) from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( divide, @@ -92,6 +97,73 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name +_FLUX_FUSED_LN_MOD_DISABLED = False +# (shape, stride, eps) signatures whose fused output has been verified +# ``torch.equal`` against the live eager chain. +_FLUX_FUSED_LN_MOD_VERIFIED: set = set() + + +def _flux_fused_ln_modulate( + norm: nn.Module, + x: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> Optional[torch.Tensor]: + """Single-kernel ``LN(x) * (1 + scale) + shift``, bit-exact vs the eager + chain, or ``None`` when the fast path does not apply. + + The Triton kernel replicates the aten LayerNorm kernel this dispatch + selects for bf16 rows (PR #34008), but bit-exactness is a property of + the live dispatch: every distinct (shape, stride, eps) combination is + verified ``torch.equal`` against the eager chain on first sight, and any + mismatch disables the fast path permanently. + """ + global _FLUX_FUSED_LN_MOD_DISABLED + + if ( + _FLUX_FUSED_LN_MOD_DISABLED + or not is_plain_layer_norm(norm, x.shape[-1]) + or not can_use_fused_layernorm_modulate(x, scale, shift) + ): + return None + sig = ( + x.shape, + x.stride(), + scale.shape, + scale.stride(), + shift.shape, + shift.stride(), + norm.eps, + ) + verified = sig in _FLUX_FUSED_LN_MOD_VERIFIED + if not verified and ( + torch.compiler.is_compiling() or torch.cuda.is_current_stream_capturing() + ): + # The first-sight check needs the eager chain and a host sync; run + # neither inside compile tracing nor CUDA graph capture. + return None + try: + out = fused_layernorm_modulate(x, scale, shift, norm.eps) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling FLUX fused LN+modulate fast path: {exc}") + _FLUX_FUSED_LN_MOD_DISABLED = True + return None + if verified: + return out + ref = modulate_scale_shift(norm(x), scale, shift) + if torch.equal(out, ref): + _FLUX_FUSED_LN_MOD_VERIFIED.add(sig) + return out + logger.warning_once( + "FLUX fused LN+modulate fast path is not bit-exact against this " + "platform's LayerNorm dispatch; falling back to eager" + ) + _FLUX_FUSED_LN_MOD_DISABLED = True + return ref + + def _flux_norm_modulate( site: nn.Module, norm: nn.Module, @@ -101,10 +173,16 @@ def _flux_norm_modulate( ) -> torch.Tensor: """``norm(x) * (1 + scale) + shift`` for the FLUX adaLN sites. - Default: affine-free LayerNorm + the bit-exact fused modulate. When the - site is mounted (``quality="high"``) and the per-call guard passes, the - modulate is folded into the LN affine instead (one kernel; not bit-exact). + Priority: (1) the bit-exact single-kernel LN+modulate -- lossless, so it + needs no quality gate and also supersedes the ``quality="high"`` affine + fold wherever it verifies; (2) when the site is mounted + (``quality="high"``) and the bit-exact kernel is unavailable, the + modulate folded into the LN affine (one aten kernel; not bit-exact); + (3) affine-free LayerNorm + the bit-exact fused modulate. """ + out = _flux_fused_ln_modulate(norm, x, scale, shift) + if out is not None: + return out if fused_ln_modulate_active(site) and can_fuse_ln_modulate(x, scale, shift): return fused_ln_modulate(x, scale, shift, norm.eps) return modulate_scale_shift(norm(x), scale, shift) diff --git a/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py new file mode 100644 index 000000000..867591212 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py @@ -0,0 +1,75 @@ +"""FLUX.1 fused LN+modulate fast path must stay bit-exact vs eager.""" + +import pytest +import torch + +import sglang.multimodal_gen.runtime.models.dits.flux as flux +from sglang.kernels.ops.diffusion.fused_ln_modulate import ( + mark_fused_ln_modulate_site, + mount_fused_ln_modulate, +) +from sglang.multimodal_gen.runtime.models.dits.flux import ( + _flux_fused_ln_modulate, + _flux_norm_modulate, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +def _eager(norm, x, scale, shift): + return norm(x) * (1 + scale[:, None]) + shift[:, None] + + +def _make_site_inputs(shape, chunks, seed): + torch.manual_seed(seed) + batch, seq, hidden = shape + norm = torch.nn.LayerNorm(hidden, eps=1e-6, elementwise_affine=False).cuda() + x = (torch.randn(batch, seq, hidden, device="cuda") * 8).bfloat16() + emb = torch.randn(batch, chunks * hidden, device="cuda").bfloat16() + parts = emb.chunk(chunks, dim=1) # strided adaLN projection views + return norm, x, parts[0], parts[1] + + +@pytest.mark.parametrize( + "shape,chunks", + [ + ((1, 4096, 3072), 6), # dual-stream image tokens (1024^2), chunk(6) + ((1, 512, 3072), 6), # dual-stream text tokens + ((1, 4608, 3072), 3), # single-stream concat, chunk(3) + ((2, 300, 3072), 6), # CFG batch, odd seq + ], +) +def test_flux_fused_ln_modulate_is_bit_exact(shape, chunks): + # Every distinct (shape, stride, eps) signature the FLUX.1 sites emit + # must verify torch.equal on first sight and stay enabled. + norm, x, shift, scale = _make_site_inputs(shape, chunks, seed=0) + out = _flux_fused_ln_modulate(norm, x, scale, shift) + assert out is not None + assert torch.equal(out, _eager(norm, x, scale, shift)) + assert not flux._FLUX_FUSED_LN_MOD_DISABLED + assert flux._FLUX_FUSED_LN_MOD_VERIFIED + + +def test_flux_norm_modulate_bitexact_supersedes_high_fold(): + # With the quality="high" affine fold mounted, the bit-exact kernel + # still takes priority, so the site output stays lossless. + site = torch.nn.Module() + mark_fused_ln_modulate_site(site) + assert mount_fused_ln_modulate(site) + norm, x, shift, scale = _make_site_inputs((1, 128, 3072), 6, seed=1) + out = _flux_norm_modulate(site, norm, x, scale, shift) + assert torch.equal(out, _eager(norm, x, scale, shift)) + + +def test_flux_fused_ln_modulate_rejects_unsupported_hidden(): + # hidden % 4 != 0 is outside the kernel contract and must bail out. + norm, x, shift, scale = _make_site_inputs((1, 64, 3070), 6, seed=2) + assert _flux_fused_ln_modulate(norm, x, scale, shift) is None + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__]))