[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 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
295784723a
commit
3654740347
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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__]))
|
||||
Reference in New Issue
Block a user