From 3654740347f63edd3e8df78b2282ec79782d4f2c Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Thu, 6 Aug 2026 19:58:44 +0800 Subject: [PATCH] [diffusion] ERNIE-Image bit-exact fused RMSNorm+scale/shift (H200 1024^2 e2e 15.63 -> 15.00 s, denoise -3.3%) (#33854) Co-authored-by: Claude Fable 5 --- .../triton/rmsnorm_scale_shift_bitexact.py | 342 ++++++++++++++++++ .../runtime/models/dits/ernie_image.py | 129 ++++++- .../diffusion/test_ernie_norm_scale_shift.py | 57 +++ 3 files changed, 524 insertions(+), 4 deletions(-) create mode 100644 python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py create mode 100644 test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py diff --git a/python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py b/python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py new file mode 100644 index 000000000..59002f493 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py @@ -0,0 +1,342 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Bit-exact fused RMSNorm + adaLN scale/shift (optionally with a preceding +residual-gate add) for bf16 activations. + +Replaces the eager ERNIE-Image adaLN chain + + ``norm(x) * (1 + scale) + shift`` (4 kernels) + ``res = residual + gate * update`` before the norm (+1 kernel) + +with one Triton kernel per site while reproducing the eager chain's rounding +*bit for bit* (``torch.equal``), so it can be wired unconditionally (no +quality gate), unlike a plain fp32 single-pass fusion which perturbs the +50-step denoising trajectory (T8: PSNR 18.83 dB at quality=high). + +Numerics contract (each step matches the eager kernel boundary): + +- ``RMSNorm.forward_cuda`` dispatches to ``sgl_kernel.rmsnorm`` -> + flashinfer's CuTe-DSL ``RMSNormKernel``. For contiguous bf16 rows with + ``H == 64 * threads_per_row`` (threads_per_row 32 for H<=3072 else 64, + cluster_n == 1) that kernel computes, per row: + + * thread ``tx`` owns columns ``{8*tpr*b + 8*tx + v : b, v in [0,8)}``; + its fragment is ordered ``v`` fastest, then ``b``; + * ``x_sq = x * x`` in fp32 (each square rounded separately, no FMA), then + an *ordered* sequential fadd chain over the 64 fragment values + (MLIR ``vector.reduction`` without reassoc); + * warp reduction via ``shfl.bfly`` with offsets 1,2,4,8,16 == an + adjacent-pairs fold tree; two warp sums are then added (tpr == 64); + * ``rstd = rsqrt.approx.f32(sum_sq / H + eps)`` (``cute.math.rsqrt`` + with fastmath); + * ``y = (bf16)(float(x) * rstd * (w + 0.0))`` -- one final rounding. + +- The aten modulate chain rounds to bf16 after every op (fp32 opmath): + ``round(1 + scale)``, ``round(y * that)``, ``round(prod + shift)``. + +- The residual variant reproduces the eager pair ``round(gate * update)``, + ``round(residual + that)`` (identical to ``residual_gate_add_cuda``) and + feeds the rounded result into the same faithful norm. + +The per-fragment chain is expressed as 64 ordered adds over (TPR,)-wide +strided loads, the fold trees with ``tl.reshape``/``tl.split`` (order-exact, +single adds), bf16 boundaries with a bitcast round-to-nearest-even helper, +and the square with an opaque ``mul.rn.f32`` so the compiler cannot contract +it into an FMA. Verified ``torch.equal`` against the live eager chain on +(1,4216,4096)/(1,4096,4096)/(2,1140,4096)/(1,128,2048) bf16; callers should +still verify once at runtime and fall back if the platform's rmsnorm dispatch +ever changes (see ``ernie_image.py``). +""" + +from __future__ import annotations + +from typing import Tuple + +import torch +import triton # type: ignore +import triton.language as tl # type: ignore + +from sglang.srt.utils.custom_op import register_custom_op + + +@triton.jit +def _round_bf16_to_fp32(value): + # RNE round of an fp32 value to bf16 precision, staying in fp32 registers + # (also blocks any fmul+fadd contraction across the boundary). + bits = value.to(tl.int32, bitcast=True) + rounding_bias = 0x7FFF + ((bits >> 16) & 1) + rounded_bits = (bits + rounding_bias) & -65536 + return rounded_bits.to(tl.float32, bitcast=True) + + +@triton.jit +def _mul_rn_f32(x, y): + # opaque mul.rn.f32: keeps the square a separately rounded fp32 op and + # blocks contraction with the following add. + return tl.inline_asm_elementwise( + asm="mul.rn.f32 $0, $1, $2;", + constraints="=f,f,f", + args=[x, y], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _rsqrt_approx_f32(x): + return tl.inline_asm_elementwise( + asm="rsqrt.approx.f32 $0, $1;", + constraints="=f,f", + args=[x], + dtype=tl.float32, + is_pure=True, + pack=1, + ) + + +@triton.jit +def _fold_adjacent(p, rows: tl.constexpr, width: tl.constexpr): + # (rows, 2*width) -> (rows, width): add adjacent pairs (even + odd), the + # tree a shfl butterfly with the smallest offset first produces. + a, b = tl.split(tl.reshape(p, (rows, width, 2))) + return a + b + + +@triton.jit +def _rmsnorm_scale_shift_kernel( + out_ptr, + res_out_ptr, + x_ptr, # norm input (no gate) / update (with gate) + residual_ptr, + gate_ptr, + weight_ptr, + scale_ptr, + shift_ptr, + seq_len, + eps, + D: tl.constexpr, + TPR: tl.constexpr, + WPR: tl.constexpr, + HAS_GATE: tl.constexpr, +): + row = tl.program_id(0).to(tl.int64) + batch = row // seq_len + row_base = row * D + vec_base = batch * D + + # ----- pass 1: sum of squares in the exact CuTe reduction order ----- + # "thread" tx of the replicated kernel owns columns 8*TPR*b + 8*tx + v; + # its fragment is iterated v fastest, then b, as one ordered fadd chain + # (each square separately rounded, no FMA). + tx = tl.arange(0, TPR) * 8 + acc = tl.zeros((TPR,), dtype=tl.float32) + for b in tl.static_range(8): + for v in tl.static_range(8): + col = tx + (b * 8 * TPR + v) + if HAS_GATE: + rj = tl.load(residual_ptr + row_base + col).to(tl.float32) + uj = tl.load(x_ptr + row_base + col).to(tl.float32) + gj = tl.load(gate_ptr + vec_base + col).to(tl.float32) + # eager pair: bf16 round after gate*update and after the add + xj = _round_bf16_to_fp32(rj + _round_bf16_to_fp32(gj * uj)) + else: + xj = tl.load(x_ptr + row_base + col).to(tl.float32) + acc = acc + _mul_rn_f32(xj, xj) + + # warp butterfly (offsets 1,2,4,8,16) == adjacent-pairs fold tree, + # then the WPR warp sums are combined the same way. + p = tl.reshape(acc, (WPR, 32)) + p = _fold_adjacent(p, WPR, 16) + p = _fold_adjacent(p, WPR, 8) + p = _fold_adjacent(p, WPR, 4) + p = _fold_adjacent(p, WPR, 2) + p = _fold_adjacent(p, WPR, 1) + s = tl.reshape(p, (1, WPR)) + if WPR == 2: + s = _fold_adjacent(s, 1, 1) + rcp = tl.sum(_rsqrt_approx_f32(s / D + eps)) # single element, exact + + # ----- pass 2: normalize + modulate, contiguous chunks ----- + for i in tl.static_range(D // 1024): + cols = i * 1024 + tl.arange(0, 1024) + if HAS_GATE: + r = tl.load(residual_ptr + row_base + cols).to(tl.float32) + u = tl.load(x_ptr + row_base + cols).to(tl.float32) + g = tl.load(gate_ptr + vec_base + cols).to(tl.float32) + xin = _round_bf16_to_fp32(r + _round_bf16_to_fp32(g * u)) + tl.store(res_out_ptr + row_base + cols, xin) + else: + xin = tl.load(x_ptr + row_base + cols).to(tl.float32) + w = tl.load(weight_ptr + cols).to(tl.float32) + sc = tl.load(scale_ptr + vec_base + cols).to(tl.float32) + sh = tl.load(shift_ptr + vec_base + cols).to(tl.float32) + y = _round_bf16_to_fp32(xin * rcp * w) # (bf16)(x * rstd * w) + one_plus = _round_bf16_to_fp32(1.0 + sc) + prod = _round_bf16_to_fp32(y * one_plus) + tl.store(out_ptr + row_base + cols, prod + sh) # store rounds to bf16 + + +def _threads_per_row(hidden: int) -> int | None: + # mirror of flashinfer RMSNormKernel._compute_threads_per_row for the + # regime this kernel replicates (one 8-wide vector per (thread, block)) + tpr = 32 if hidden <= 3072 else 64 if hidden <= 6144 else None + if tpr is None or hidden != 64 * tpr: + return None + return tpr + + +def _is_row_broadcast(t: torch.Tensor, x: torch.Tensor) -> bool: + return ( + t.dtype is torch.bfloat16 + and t.is_cuda + and t.device == x.device + and t.shape == (x.shape[0], 1, x.shape[-1]) + and t.is_contiguous() + ) + + +def can_use_fused_rmsnorm_scale_shift( + x: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> bool: + return ( + x.dtype is torch.bfloat16 + and x.is_cuda + and x.dim() == 3 + and x.is_contiguous() + and _threads_per_row(x.shape[-1]) is not None + and weight.dtype is torch.bfloat16 + and weight.is_cuda + and weight.device == x.device + and weight.shape == (x.shape[-1],) + and weight.is_contiguous() + and _is_row_broadcast(scale, x) + and _is_row_broadcast(shift, x) + ) + + +def can_use_fused_scale_residual_rmsnorm_scale_shift( + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> bool: + return ( + can_use_fused_rmsnorm_scale_shift(residual, weight, scale, shift) + and update.dtype is torch.bfloat16 + and update.is_cuda + and update.device == residual.device + and update.shape == residual.shape + and update.is_contiguous() + and _is_row_broadcast(gate, residual) + ) + + +def _fake_norm_scale_shift( + x: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + eps: float, +) -> torch.Tensor: + return torch.empty_like(x) + + +@register_custom_op( + op_name="triton_fused_rmsnorm_scale_shift_bitexact", + mutates_args=[], + fake_impl=_fake_norm_scale_shift, +) +def fused_rmsnorm_scale_shift_bitexact( + x: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + eps: float, +) -> torch.Tensor: + """``norm(x) * (1 + scale) + shift``, bit-exact vs the eager chain.""" + batch, seq_len, hidden = x.shape + tpr = _threads_per_row(hidden) + out = torch.empty_like(x) + with torch.cuda.device(x.device): + _rmsnorm_scale_shift_kernel[(batch * seq_len,)]( + out, + out, + x, + x, + x, + weight, + scale, + shift, + seq_len, + eps, + D=hidden, + TPR=tpr, + WPR=tpr // 32, + HAS_GATE=False, + # num_warps must match the replicated kernel's warps-per-row: + # larger blocks trigger pathological Triton layout conversions + # in the fold stage (measured 25us -> 500us at num_warps=4). + num_warps=max(tpr // 32, 1), + ) + return out + + +def _fake_scale_residual_norm_scale_shift( + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + eps: float, +) -> Tuple[torch.Tensor, torch.Tensor]: + return torch.empty_like(residual), torch.empty_like(residual) + + +@register_custom_op( + op_name="triton_fused_scale_residual_rmsnorm_scale_shift_bitexact", + mutates_args=[], + fake_impl=_fake_scale_residual_norm_scale_shift, +) +def fused_scale_residual_rmsnorm_scale_shift_bitexact( + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + eps: float, +) -> Tuple[torch.Tensor, torch.Tensor]: + """``res = residual + gate * update; norm(res) * (1 + scale) + shift``. + + Returns ``(modulated, res)``, both bit-exact vs the eager pair + norm + chain (and therefore vs the ``residual_gate_add_cuda`` fast path). + """ + batch, seq_len, hidden = residual.shape + tpr = _threads_per_row(hidden) + out = torch.empty_like(residual) + res_out = torch.empty_like(residual) + with torch.cuda.device(residual.device): + _rmsnorm_scale_shift_kernel[(batch * seq_len,)]( + out, + res_out, + update, + residual, + gate, + weight, + scale, + shift, + seq_len, + eps, + D=hidden, + TPR=tpr, + WPR=tpr // 32, + HAS_GATE=True, + num_warps=max(tpr // 32, 1), + ) + return out, res_out diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py index def48ee48..c7deb1c28 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py @@ -23,6 +23,12 @@ from sglang.kernels.ops.diffusion.residual_gate_add import ( can_use_residual_gate_add_cuda, residual_gate_add_cuda, ) +from sglang.kernels.ops.diffusion.triton.rmsnorm_scale_shift_bitexact import ( + can_use_fused_rmsnorm_scale_shift, + can_use_fused_scale_residual_rmsnorm_scale_shift, + fused_rmsnorm_scale_shift_bitexact, + fused_scale_residual_rmsnorm_scale_shift_bitexact, +) from sglang.multimodal_gen.configs.models.dits.ernie_image import ( ErnieImageDitConfig, ) @@ -80,6 +86,121 @@ def _ernie_residual_gate_add( return residual + gate * update +_ERNIE_FUSED_NORM_DISABLED = False +_ERNIE_FUSED_NORM_VERIFIED = False +_ERNIE_FUSED_GATED_NORM_DISABLED = False +_ERNIE_FUSED_GATED_NORM_VERIFIED = False + + +def _eager_norm_scale_shift( + norm: RMSNorm, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + return norm(x) * (1 + scale) + shift + + +def _ernie_norm_scale_shift( + norm: RMSNorm, x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + """Single-kernel ``norm(x) * (1 + scale) + shift``, bit-exact vs eager. + + The Triton kernel replicates the flashinfer CuTe rmsnorm reduction order + and every aten bf16 rounding boundary. Because bit-exactness depends on + which rmsnorm implementation ``RMSNorm.forward_cuda`` dispatches to, the + first call verifies ``torch.equal`` against the eager chain and disables + the fast path permanently on any mismatch. + """ + global _ERNIE_FUSED_NORM_DISABLED, _ERNIE_FUSED_NORM_VERIFIED + + if ( + not _ERNIE_FUSED_NORM_DISABLED + and norm.variance_size_override is None + and can_use_fused_rmsnorm_scale_shift(x, norm.weight, scale, shift) + and (_ERNIE_FUSED_NORM_VERIFIED or not torch.compiler.is_compiling()) + ): + try: + out = fused_rmsnorm_scale_shift_bitexact( + x, norm.weight, scale, shift, norm.variance_epsilon + ) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling ERNIE fused-norm fast path: {exc}") + _ERNIE_FUSED_NORM_DISABLED = True + else: + if _ERNIE_FUSED_NORM_VERIFIED: + return out + ref = _eager_norm_scale_shift(norm, x, scale, shift) + if torch.equal(out, ref): + _ERNIE_FUSED_NORM_VERIFIED = True + return out + logger.warning_once( + "ERNIE fused-norm fast path is not bit-exact against this " + "platform's rmsnorm dispatch; falling back to eager" + ) + _ERNIE_FUSED_NORM_DISABLED = True + return ref + + return _eager_norm_scale_shift(norm, x, scale, shift) + + +def _ernie_gated_norm_scale_shift( + norm: RMSNorm, + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """``res = residual + gate * update`` then the fused norm/scale/shift. + + Returns ``(modulated, res)``. Single kernel, bit-exact vs the eager pair + (and the ``residual_gate_add_cuda`` fast path) + norm chain; first call + self-verifies like :func:`_ernie_norm_scale_shift`. + """ + global _ERNIE_FUSED_GATED_NORM_DISABLED, _ERNIE_FUSED_GATED_NORM_VERIFIED + + if ( + not _ERNIE_FUSED_GATED_NORM_DISABLED + and norm.variance_size_override is None + and can_use_fused_scale_residual_rmsnorm_scale_shift( + residual, update, gate, norm.weight, scale, shift + ) + and (_ERNIE_FUSED_GATED_NORM_VERIFIED or not torch.compiler.is_compiling()) + ): + try: + out, res = fused_scale_residual_rmsnorm_scale_shift_bitexact( + residual, + update, + gate, + norm.weight, + scale, + shift, + norm.variance_epsilon, + ) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling ERNIE fused gated-norm fast path: {exc}") + _ERNIE_FUSED_GATED_NORM_DISABLED = True + else: + if _ERNIE_FUSED_GATED_NORM_VERIFIED: + return out, res + res_ref = residual + gate * update + ref = _eager_norm_scale_shift(norm, res_ref, scale, shift) + if torch.equal(out, ref) and torch.equal(res, res_ref): + _ERNIE_FUSED_GATED_NORM_VERIFIED = True + return out, res + logger.warning_once( + "ERNIE fused gated-norm fast path is not bit-exact against " + "this platform's rmsnorm dispatch; falling back to eager" + ) + _ERNIE_FUSED_GATED_NORM_DISABLED = True + return ref, res_ref + + res = _ernie_residual_gate_add(residual, update, gate) + return _eager_norm_scale_shift(norm, res, scale, shift), res + + def _rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor: assert dim % 2 == 0 scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim @@ -293,13 +414,13 @@ class ErnieImageSharedAdaLNBlock(nn.Module): attn_mask_meta: dict | None = None, ) -> torch.Tensor: residual = x - x = self.adaLN_sa_ln(x) * (1 + scale_msa) + shift_msa + x = _ernie_norm_scale_shift(self.adaLN_sa_ln, x, scale_msa, shift_msa) attn_out = self.self_attention( x, rotary_pos_emb, attn_mask=attn_mask, attn_mask_meta=attn_mask_meta ) - residual = _ernie_residual_gate_add(residual, attn_out, gate_msa) - - x = self.adaLN_mlp_ln(residual) * (1 + scale_mlp) + shift_mlp + x, residual = _ernie_gated_norm_scale_shift( + self.adaLN_mlp_ln, residual, attn_out, gate_msa, scale_mlp, shift_mlp + ) x = _ernie_residual_gate_add(residual, self.mlp(x), gate_mlp) return x diff --git a/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py b/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py new file mode 100644 index 000000000..aeca75111 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py @@ -0,0 +1,57 @@ +"""ERNIE fused norm/scale/shift fast paths must stay bit-exact vs eager.""" + +import sys + +import pytest +import torch + +import sglang.multimodal_gen.runtime.models.dits.ernie_image as ernie_image +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm +from sglang.multimodal_gen.runtime.models.dits.ernie_image import ( + _ernie_gated_norm_scale_shift, + _ernie_norm_scale_shift, +) +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") + + +@pytest.mark.parametrize("shape", [(1, 4216, 4096), (2, 1140, 4096), (1, 128, 2048)]) +def test_fused_norm_scale_shift_is_bit_exact(shape): + # (1, 4216, 4096) is the real ERNIE-Image shape (1024^2 image + text + # tokens, hidden 4096); 2048 covers the threads_per_row=32 regime. + torch.manual_seed(0) + batch, seq, hidden = shape + norm = RMSNorm(hidden, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + with torch.no_grad(): + norm.weight.copy_(torch.randn(hidden)) + x = torch.randn(batch, seq, hidden, device="cuda", dtype=torch.bfloat16) + residual = torch.randn_like(x) + update = torch.randn_like(x) + scale = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + shift = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + gate = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) + + with torch.no_grad(): + out = _ernie_norm_scale_shift(norm, x, scale, shift) + ref = norm(x) * (1 + scale) + shift + assert torch.equal(out, ref) + + out2, res = _ernie_gated_norm_scale_shift( + norm, residual, update, gate, scale, shift + ) + res_ref = residual + gate * update + ref2 = norm(res_ref) * (1 + scale) + shift + assert torch.equal(res, res_ref) + assert torch.equal(out2, ref2) + + # the fast paths must actually be in use (not silently disabled) + assert ernie_image._ERNIE_FUSED_NORM_VERIFIED + assert ernie_image._ERNIE_FUSED_GATED_NORM_VERIFIED + assert not ernie_image._ERNIE_FUSED_NORM_DISABLED + assert not ernie_image._ERNIE_FUSED_GATED_NORM_DISABLED + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__]))