[diffusion] FLUX.2 VAE decoder fast path behind quality=high (H200: 1024^2 97.6->29.2 ms, 2048^2 437.2->168.5 ms) (#33451)
This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Channels-last two-pass GroupNorm(+SiLU) Triton kernels.
|
||||
|
||||
Relationship to ``group_norm_silu.py`` (``triton_group_norm_silu``): that
|
||||
kernel serves the general NCHW-contiguous case (any channels-per-group, any
|
||||
ndim, always applies SiLU) and keeps backing ``apply_group_norm_silu`` for
|
||||
the HunyuanVAE / latent-upsampler paths. This module is a complementary
|
||||
kernel for the channels_last VAE decoder fast path with a different contract:
|
||||
|
||||
- activations whose channel dim is innermost ((N, H, W, C) channels_last
|
||||
views or (N, L, C) rows) are normalized without any layout round-trip,
|
||||
which is what makes a channels_last decoder run end-to-end without
|
||||
nchwToNhwc transposes;
|
||||
- fp32 statistics, with the affine transform folded into per-(batch, channel)
|
||||
``scale``/``shift`` in a separate finalize kernel, so the apply pass is a
|
||||
pure elementwise kernel;
|
||||
- optional SiLU epilogue (``apply_silu=False`` gives plain GroupNorm);
|
||||
- restricted static shapes: power-of-two ``C <= 2048`` that ``num_groups``
|
||||
divides. Callers must treat a ``None`` return as "unsupported" and fall
|
||||
back to their reference path.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import triton # type: ignore
|
||||
import triton.language as tl # type: ignore
|
||||
|
||||
_SUPPORTED_DTYPES = {torch.float16, torch.bfloat16, torch.float32}
|
||||
_MAX_CHANNELS = 2048
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _gn_partial_rows_kernel(
|
||||
x_ptr,
|
||||
psum_ptr,
|
||||
psq_ptr,
|
||||
rows,
|
||||
rows_per_prog,
|
||||
C: tl.constexpr,
|
||||
BLOCK_R: tl.constexpr,
|
||||
):
|
||||
# x is (N, rows, C) with C innermost-contiguous (channels_last view).
|
||||
chunk = tl.program_id(0).to(tl.int64)
|
||||
n = tl.program_id(1).to(tl.int64)
|
||||
nchunks = tl.num_programs(0)
|
||||
row0 = chunk * rows_per_prog
|
||||
cols = tl.arange(0, C)
|
||||
acc_s = tl.zeros((C,), tl.float32)
|
||||
acc_q = tl.zeros((C,), tl.float32)
|
||||
x_base = x_ptr + n * rows * C
|
||||
for r_off in range(0, rows_per_prog, BLOCK_R):
|
||||
rs = row0 + r_off + tl.arange(0, BLOCK_R)
|
||||
m = rs < rows
|
||||
x = tl.load(
|
||||
x_base + rs[:, None] * C + cols[None, :],
|
||||
mask=m[:, None],
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
acc_s += tl.sum(x, 0)
|
||||
acc_q += tl.sum(x * x, 0)
|
||||
out_off = (n * nchunks + chunk) * C + cols
|
||||
tl.store(psum_ptr + out_off, acc_s)
|
||||
tl.store(psq_ptr + out_off, acc_q)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _gn_finalize_kernel(
|
||||
psum_ptr,
|
||||
psq_ptr,
|
||||
w_ptr,
|
||||
b_ptr,
|
||||
ss_ptr,
|
||||
nchunks,
|
||||
group_numel,
|
||||
eps,
|
||||
C,
|
||||
CPG: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
):
|
||||
g = tl.program_id(0).to(tl.int64)
|
||||
n = tl.program_id(1).to(tl.int64)
|
||||
cols = g * CPG + tl.arange(0, CPG)
|
||||
s = tl.zeros((), tl.float32)
|
||||
q = tl.zeros((), tl.float32)
|
||||
for k0 in range(0, nchunks, BLOCK_K):
|
||||
ks = k0 + tl.arange(0, BLOCK_K)
|
||||
m = ks < nchunks
|
||||
offs = (n * nchunks + ks)[:, None] * C + cols[None, :]
|
||||
s += tl.sum(tl.load(psum_ptr + offs, mask=m[:, None], other=0.0))
|
||||
q += tl.sum(tl.load(psq_ptr + offs, mask=m[:, None], other=0.0))
|
||||
mean = s / group_numel
|
||||
var = q / group_numel - mean * mean
|
||||
var = tl.maximum(var, 0.0)
|
||||
rstd = tl.rsqrt(var + eps)
|
||||
w = tl.load(w_ptr + cols).to(tl.float32)
|
||||
b = tl.load(b_ptr + cols).to(tl.float32)
|
||||
scale = w * rstd
|
||||
shift = b - mean * scale
|
||||
tl.store(ss_ptr + n * 2 * C + cols, scale)
|
||||
tl.store(ss_ptr + (n * 2 + 1) * C + cols, shift)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _gn_apply_rows_kernel(
|
||||
x_ptr,
|
||||
ss_ptr,
|
||||
y_ptr,
|
||||
rows,
|
||||
C: tl.constexpr,
|
||||
BLOCK_R: tl.constexpr,
|
||||
SILU: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0).to(tl.int64)
|
||||
n = tl.program_id(1).to(tl.int64)
|
||||
cols = tl.arange(0, C)
|
||||
scale = tl.load(ss_ptr + n * 2 * C + cols)
|
||||
shift = tl.load(ss_ptr + (n * 2 + 1) * C + cols)
|
||||
rs = pid * BLOCK_R + tl.arange(0, BLOCK_R)
|
||||
m = rs < rows
|
||||
offs = n * rows * C + rs[:, None] * C + cols[None, :]
|
||||
x = tl.load(x_ptr + offs, mask=m[:, None], other=0.0).to(tl.float32)
|
||||
y = x * scale[None, :] + shift[None, :]
|
||||
if SILU:
|
||||
y = y * tl.sigmoid(y)
|
||||
tl.store(y_ptr + offs, y.to(y_ptr.dtype.element_ty), mask=m[:, None])
|
||||
|
||||
|
||||
def _gn_silu_rows(x3, weight, bias, num_groups, eps, apply_silu):
|
||||
"""x3: (N, R, C) contiguous with C innermost. Returns same-shape tensor."""
|
||||
n_batch, rows, c = x3.shape
|
||||
cpg = c // num_groups
|
||||
block_r = max(1, 8192 // c)
|
||||
rows_per_prog = block_r * 32
|
||||
nchunks = triton.cdiv(rows, rows_per_prog)
|
||||
psum = torch.empty((n_batch, nchunks, c), device=x3.device, dtype=torch.float32)
|
||||
psq = torch.empty_like(psum)
|
||||
_gn_partial_rows_kernel[(nchunks, n_batch)](
|
||||
x3, psum, psq, rows, rows_per_prog, C=c, BLOCK_R=block_r, num_warps=4
|
||||
)
|
||||
ss = torch.empty((n_batch, 2, c), device=x3.device, dtype=torch.float32)
|
||||
block_k = max(1, min(4096 // cpg, triton.next_power_of_2(nchunks)))
|
||||
_gn_finalize_kernel[(num_groups, n_batch)](
|
||||
psum,
|
||||
psq,
|
||||
weight,
|
||||
bias,
|
||||
ss,
|
||||
nchunks,
|
||||
rows * cpg,
|
||||
eps,
|
||||
c,
|
||||
CPG=cpg,
|
||||
BLOCK_K=block_k,
|
||||
num_warps=4,
|
||||
)
|
||||
y3 = torch.empty_like(x3)
|
||||
_gn_apply_rows_kernel[(triton.cdiv(rows, block_r), n_batch)](
|
||||
x3, ss, y3, rows, C=c, BLOCK_R=block_r, SILU=apply_silu, num_warps=4
|
||||
)
|
||||
return y3
|
||||
|
||||
|
||||
def _twopass_supported(x, weight, bias, num_groups) -> bool:
|
||||
"""Tensor-level support check shared by the 4D and rows entry points."""
|
||||
if not (x.is_cuda and not torch.is_grad_enabled()):
|
||||
return False
|
||||
if x.requires_grad or x.dtype not in _SUPPORTED_DTYPES:
|
||||
return False
|
||||
if weight is None or bias is None:
|
||||
return False
|
||||
c = x.shape[1] if x.dim() == 4 else x.shape[-1]
|
||||
if weight.shape != (c,) or bias.shape != (c,):
|
||||
return False
|
||||
if num_groups < 1 or c % num_groups != 0:
|
||||
return False
|
||||
# tl.arange needs a power-of-two C; num_groups divides it, so the
|
||||
# channels-per-group finalize block is a power of two as well.
|
||||
return triton.next_power_of_2(c) == c and c <= _MAX_CHANNELS
|
||||
|
||||
|
||||
def group_norm_silu_4d(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
num_groups: int,
|
||||
eps: float,
|
||||
apply_silu: bool = True,
|
||||
) -> torch.Tensor | None:
|
||||
"""Fused GroupNorm(+SiLU) for a channels_last 4D (N, C, H, W) activation.
|
||||
|
||||
Runs the rows kernel on the free (N, H*W, C) view (no layout copy) and
|
||||
preserves the channels_last output layout. Returns ``None`` when the
|
||||
input is unsupported; callers must fall back to their reference path.
|
||||
"""
|
||||
if x.dim() != 4 or not _twopass_supported(x, weight, bias, num_groups):
|
||||
return None
|
||||
n_batch, c, h, w = x.shape
|
||||
# c > 1 and a non-trivial spatial extent make the channels_last check
|
||||
# unambiguous (degenerate shapes are contiguous in both formats).
|
||||
if not (
|
||||
c > 1
|
||||
and (h > 1 or w > 1)
|
||||
and x.is_contiguous(memory_format=torch.channels_last)
|
||||
):
|
||||
return None
|
||||
x3 = x.permute(0, 2, 3, 1).reshape(n_batch, h * w, c)
|
||||
y3 = _gn_silu_rows(x3, weight, bias, num_groups, eps, apply_silu)
|
||||
return y3.reshape(n_batch, h, w, c).permute(0, 3, 1, 2)
|
||||
|
||||
|
||||
def group_norm_silu_rows(
|
||||
x3: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
num_groups: int,
|
||||
eps: float,
|
||||
apply_silu: bool = True,
|
||||
) -> torch.Tensor | None:
|
||||
"""Fused GroupNorm(+SiLU) over (N, L, C) rows (C = channels, innermost).
|
||||
|
||||
Returns ``None`` when the input is unsupported; callers must fall back.
|
||||
"""
|
||||
if x3.dim() != 3 or not x3.is_contiguous():
|
||||
return None
|
||||
if not _twopass_supported(x3, weight, bias, num_groups):
|
||||
return None
|
||||
return _gn_silu_rows(x3, weight, bias, num_groups, eps, apply_silu)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"group_norm_silu_4d",
|
||||
"group_norm_silu_rows",
|
||||
]
|
||||
@@ -0,0 +1,435 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CUDA fast paths for the FLUX.2 VAE decoder (AutoencoderKLFlux2, diffusers
|
||||
``Decoder``) on image workloads.
|
||||
|
||||
All rewrites are mathematically exact re-associations of the original
|
||||
operators (weight folding is done lazily in fp32 on first fast-path use and
|
||||
written back to the model compute dtype). The wrappers are installed once at
|
||||
VAE load and stay in place; each forward dispatches on a shared
|
||||
request-scoped :class:`VaeFastPathGate`: requests with ``quality == "high"``
|
||||
run the fast paths, the ``"lossless"`` default runs the original module path
|
||||
bit-for-bit.
|
||||
|
||||
- channels_last: run the decoder in channels_last so cuDNN convolutions run
|
||||
natively in NHWC (removes the nchwToNhwc/nhwcToNchw transpose kernels
|
||||
around every conv). The parameter layout is swapped at decode entry to
|
||||
match the gate, so lossless decodes always run the NCHW baseline kernels
|
||||
bit-for-bit. The mid-block attention needs a layout-safe forward because
|
||||
diffusers' ``AttnProcessor2_0`` calls ``.view`` on the 4D activation,
|
||||
which is illegal for channels_last tensors.
|
||||
- norm+SiLU: two-pass channels_last GroupNorm(+SiLU) Triton fusion (fp32
|
||||
statistics) for the ResnetBlock2D norm1/norm2 + SiLU chains and the
|
||||
decoder ``conv_norm_out``/``conv_act`` tail, which upcast to fp32 under
|
||||
autocast and dominate the decode profile.
|
||||
- fused upsample: nearest-2x upsample + Conv2d(3x3, p1) ==
|
||||
ConvTranspose2d(k4, s2, p1) with a lazily-summed kernel. Removes the 4x
|
||||
upsampled intermediate materialization.
|
||||
- attention V/proj fold: fold the attention output projection into the V
|
||||
projection of the single-head mid-block attention (softmax rows sum to 1,
|
||||
so ``A @ (V W_v^T + b_v) W_o^T + b_o == A @ (V W_v'^T + b')``).
|
||||
|
||||
Install is all-or-nothing and fail-closed: without Triton, or with any
|
||||
attention block lacking the layout-safe rewrite, no wrapper is installed and
|
||||
every request runs the unmodified decoder.
|
||||
"""
|
||||
|
||||
from types import MethodType
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
try:
|
||||
from sglang.kernels.ops.diffusion.triton.group_norm_silu_twopass import (
|
||||
group_norm_silu_4d,
|
||||
group_norm_silu_rows,
|
||||
)
|
||||
|
||||
_HAS_TRITON = True
|
||||
except ImportError: # pragma: no cover
|
||||
_HAS_TRITON = False
|
||||
|
||||
|
||||
class VaeFastPathGate:
|
||||
"""Mutable fast-path flag shared by every wrapper of one VAE.
|
||||
|
||||
Published on the VAE as ``_sgl_vae_fast_path_gate``; ``DecodingStage``
|
||||
enables it for the duration of a decode when the request's ``quality``
|
||||
sampling param is ``"high"``.
|
||||
"""
|
||||
|
||||
__slots__ = ("enabled",)
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.enabled = False
|
||||
|
||||
|
||||
GATE_ATTR = "_sgl_vae_fast_path_gate"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fuse A: two-pass GroupNorm(+SiLU) fusion (channels_last Triton kernel)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FusedGroupNormSiLU(nn.Module):
|
||||
"""GroupNorm + SiLU fused with the two-pass channels_last Triton kernel.
|
||||
|
||||
fp32 statistics and affine/SiLU application, output in the input dtype.
|
||||
Falls back to the original module chain (norm + F.silu, bit-identical to
|
||||
the original norm + nn.SiLU pair) for unsupported inputs and whenever
|
||||
the fast-path gate is disabled.
|
||||
"""
|
||||
|
||||
def __init__(self, norm: nn.GroupNorm, gate: VaeFastPathGate) -> None:
|
||||
super().__init__()
|
||||
# Keep the original GroupNorm state-dict layout (``weight``/``bias``)
|
||||
# so checkpoint loading and component-accuracy weight transfer do not
|
||||
# see wrapper-specific ``norm.*`` parameter names.
|
||||
self.num_groups = norm.num_groups
|
||||
self.num_channels = norm.num_channels
|
||||
self.eps = norm.eps
|
||||
self.affine = norm.affine
|
||||
self.weight = norm.weight
|
||||
self.bias = norm.bias
|
||||
self._sgl_gate = gate
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self._sgl_gate.enabled and x.dim() == 4:
|
||||
y = group_norm_silu_4d(
|
||||
x,
|
||||
self.weight,
|
||||
self.bias,
|
||||
self.num_groups,
|
||||
self.eps,
|
||||
apply_silu=True,
|
||||
)
|
||||
if y is not None:
|
||||
return y
|
||||
return F.silu(
|
||||
F.group_norm(x, self.num_groups, self.weight, self.bias, self.eps)
|
||||
)
|
||||
|
||||
|
||||
def _install_norm_silu(decoder, resnet_cls, gate: VaeFastPathGate) -> int:
|
||||
def _fusable(norm) -> bool:
|
||||
return (
|
||||
type(norm) is nn.GroupNorm
|
||||
and norm.affine
|
||||
and norm.weight is not None
|
||||
and norm.bias is not None
|
||||
)
|
||||
|
||||
count = 0
|
||||
for m in decoder.modules():
|
||||
if (
|
||||
type(m) is resnet_cls
|
||||
and m.time_emb_proj is None
|
||||
# "default"/"group" apply plain norm2 (+ SiLU); "scale_shift" and
|
||||
# "spatial" modify the activation between norm2 and SiLU.
|
||||
and m.time_embedding_norm in ("default", "group")
|
||||
and m.upsample is None
|
||||
and m.downsample is None
|
||||
and isinstance(m.nonlinearity, nn.SiLU)
|
||||
and _fusable(m.norm1)
|
||||
and _fusable(m.norm2)
|
||||
):
|
||||
m.norm1 = FusedGroupNormSiLU(m.norm1, gate)
|
||||
m.norm2 = FusedGroupNormSiLU(m.norm2, gate)
|
||||
m.nonlinearity = nn.Identity()
|
||||
count += 2
|
||||
if (
|
||||
type(getattr(decoder, "conv_norm_out", None)) is nn.GroupNorm
|
||||
and isinstance(getattr(decoder, "conv_act", None), nn.SiLU)
|
||||
and _fusable(decoder.conv_norm_out)
|
||||
):
|
||||
decoder.conv_norm_out = FusedGroupNormSiLU(decoder.conv_norm_out, gate)
|
||||
decoder.conv_act = nn.Identity()
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fuse B: nearest-2x upsample + Conv2d(3x3, p1) == ConvTranspose2d(k4, s2, p1)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Which 3x3 conv taps sum into each 4x4 transposed-conv tap (per spatial axis).
|
||||
_UPSAMPLE_TAP_MAP = {0: (2,), 1: (1, 2), 2: (0, 1), 3: (0,)}
|
||||
|
||||
|
||||
def _fold_upsample2x_conv2d_weight(conv: nn.Conv2d) -> torch.Tensor:
|
||||
"""Sum the 3x3 conv taps into the equivalent ConvTranspose2d(k4) kernel."""
|
||||
w = conv.weight.detach().float() # [Cout, Cin, 3, 3]
|
||||
cout, cin = w.shape[:2]
|
||||
wt = w.new_zeros(cin, cout, 4, 4) # ConvTranspose2d layout
|
||||
for a in range(4):
|
||||
for b in range(4):
|
||||
acc = w.new_zeros(cout, cin)
|
||||
for i in _UPSAMPLE_TAP_MAP[a]:
|
||||
for j in _UPSAMPLE_TAP_MAP[b]:
|
||||
acc += w[:, :, i, j]
|
||||
wt[:, :, a, b] = acc.t()
|
||||
wt = wt.to(conv.weight.dtype)
|
||||
if conv.weight.is_contiguous(memory_format=torch.channels_last):
|
||||
wt = wt.contiguous(memory_format=torch.channels_last)
|
||||
return wt.to(conv.weight.device)
|
||||
|
||||
|
||||
class FusedUpsample2xConv2d(nn.Module):
|
||||
"""ConvTranspose2d(k4, s2, p1) equivalent of diffusers Upsample2D
|
||||
(nearest-2x interpolate + Conv2d(3x3, p1)).
|
||||
|
||||
nearest 2x upsampling is pure pixel replication (no arithmetic), so the
|
||||
fusion only re-associates the conv taps; the kernel is summed lazily in
|
||||
fp32 on first fast-path use and written back to the conv dtype. With the
|
||||
fast-path gate disabled the original Upsample2D runs bit-for-bit.
|
||||
"""
|
||||
|
||||
def __init__(self, upsample: nn.Module, gate: VaeFastPathGate) -> None:
|
||||
super().__init__()
|
||||
# Keep ``conv`` registered directly on the wrapper so parameter names
|
||||
# remain ``...upsamplers.N.conv.*``. The unregistered original module
|
||||
# is retained only to run its exact lossless forward implementation.
|
||||
object.__setattr__(self, "_orig", upsample)
|
||||
self.conv = upsample.conv
|
||||
self.channels = upsample.channels
|
||||
self._sgl_gate = gate
|
||||
self._fused_weight = None
|
||||
|
||||
def forward(self, hidden_states, output_size=None, *args, **kwargs):
|
||||
if (
|
||||
not self._sgl_gate.enabled
|
||||
or output_size is not None
|
||||
or hidden_states.shape[1] != self.channels
|
||||
):
|
||||
return self._orig(hidden_states, output_size=output_size)
|
||||
conv = self.conv
|
||||
w = self._fused_weight
|
||||
if w is None:
|
||||
w = _fold_upsample2x_conv2d_weight(conv)
|
||||
self._fused_weight = w
|
||||
return F.conv_transpose2d(hidden_states, w, conv.bias, stride=2, padding=1)
|
||||
|
||||
|
||||
def _install_fused_upsample(decoder, upsample_cls, gate: VaeFastPathGate) -> int:
|
||||
count = 0
|
||||
for blk in decoder.up_blocks:
|
||||
upsamplers = getattr(blk, "upsamplers", None)
|
||||
if not upsamplers:
|
||||
continue
|
||||
for i, up in enumerate(upsamplers):
|
||||
if type(up) is not upsample_cls:
|
||||
continue
|
||||
conv = getattr(up, "conv", None)
|
||||
if (
|
||||
up.use_conv
|
||||
and not up.use_conv_transpose
|
||||
and up.interpolate
|
||||
and up.norm is None
|
||||
and up.name == "conv"
|
||||
and type(conv) is nn.Conv2d
|
||||
and conv.kernel_size == (3, 3)
|
||||
and conv.stride == (1, 1)
|
||||
and conv.padding == (1, 1)
|
||||
):
|
||||
upsamplers[i] = FusedUpsample2xConv2d(up, gate)
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fuse C: layout-safe single-head attention forward with the V/proj fold
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _fold_attn_vproj(m) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
w_v = m.to_v.weight.detach().float()
|
||||
b_v = m.to_v.bias.detach().float()
|
||||
w_o = m.to_out[0].weight.detach().float()
|
||||
b_o = m.to_out[0].bias.detach().float()
|
||||
dtype = m.to_v.weight.dtype
|
||||
# Softmax rows sum to 1, so the folded bias broadcasts exactly like the
|
||||
# original output-projection bias.
|
||||
return (w_o @ w_v).to(dtype).contiguous(), (w_o @ b_v + b_o).to(dtype)
|
||||
|
||||
|
||||
def _attn_fast_forward(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
temb=None,
|
||||
**cross_attention_kwargs,
|
||||
):
|
||||
if (
|
||||
not self._sgl_gate.enabled
|
||||
or encoder_hidden_states is not None
|
||||
or attention_mask is not None
|
||||
or temb is not None
|
||||
or hidden_states.ndim != 4
|
||||
):
|
||||
# Lossless requests (and anything unexpected) run the stock
|
||||
# diffusers path; the decoder layout is NCHW in that case.
|
||||
return type(self).forward(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
temb=temb,
|
||||
**cross_attention_kwargs,
|
||||
)
|
||||
|
||||
residual = hidden_states
|
||||
batch_size, channels, height, width = hidden_states.shape
|
||||
# Free view: with the gate enabled the decoder runs in channels_last.
|
||||
hs = hidden_states.permute(0, 2, 3, 1).reshape(batch_size, height * width, channels)
|
||||
|
||||
if self.group_norm is not None:
|
||||
gn = self.group_norm
|
||||
y = group_norm_silu_rows(
|
||||
hs, gn.weight, gn.bias, gn.num_groups, gn.eps, apply_silu=False
|
||||
)
|
||||
hs = y if y is not None else gn(hs.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
query = self.to_q(hs)
|
||||
key = self.to_k(hs)
|
||||
folded = self._sgl_folded_v
|
||||
if folded is None:
|
||||
folded = _fold_attn_vproj(self)
|
||||
self._sgl_folded_v = folded
|
||||
value = F.linear(hs, folded[0], folded[1])
|
||||
|
||||
out = F.scaled_dot_product_attention(
|
||||
query.unsqueeze(1), key.unsqueeze(1), value.unsqueeze(1)
|
||||
)
|
||||
out = out.squeeze(1).to(query.dtype)
|
||||
out = self.to_out[1](out) # dropout (identity in eval)
|
||||
|
||||
# The permuted view has channels_last strides, matching the residual.
|
||||
out = out.reshape(batch_size, height, width, channels).permute(0, 3, 1, 2)
|
||||
if self.residual_connection:
|
||||
out = out + residual
|
||||
return out / self.rescale_output_factor
|
||||
|
||||
|
||||
def _attn_fast_compatible(m, attn_cls, processor_cls) -> bool:
|
||||
return (
|
||||
type(m) is attn_cls
|
||||
and isinstance(m.processor, processor_cls)
|
||||
and m.heads == 1
|
||||
and m.scale_qk
|
||||
and m.spatial_norm is None
|
||||
and not m.norm_cross
|
||||
and getattr(m, "norm_q", None) is None
|
||||
and getattr(m, "norm_k", None) is None
|
||||
and getattr(m, "add_k_proj", None) is None
|
||||
and type(m.to_q) is nn.Linear
|
||||
and type(m.to_k) is nn.Linear
|
||||
and type(m.to_v) is nn.Linear
|
||||
and type(m.to_out[0]) is nn.Linear
|
||||
and isinstance(m.to_out[1], nn.Dropout)
|
||||
# The V/proj fold needs both biases present.
|
||||
and m.to_v.bias is not None
|
||||
and m.to_out[0].bias is not None
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# channels_last layout dispatch (swapped at decode entry to match the gate)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _decoder_layout_forward(self, *args, **kwargs):
|
||||
want_cl = self._sgl_gate.enabled
|
||||
if want_cl != self._sgl_channels_last:
|
||||
# Layout swaps are pure permutations of the parameter memory (values
|
||||
# are bit-identical), so flipping back to contiguous restores the
|
||||
# baseline NCHW cuDNN kernel selection exactly.
|
||||
self.to(
|
||||
memory_format=(torch.channels_last if want_cl else torch.contiguous_format)
|
||||
)
|
||||
self._sgl_channels_last = want_cl
|
||||
logger.info(
|
||||
"FLUX.2 VAE: decoder switched to %s layout.",
|
||||
"channels_last (NHWC)" if want_cl else "contiguous (NCHW)",
|
||||
)
|
||||
return type(self).forward(self, *args, **kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Entry point
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def maybe_optimize_flux2_vae(vae: nn.Module) -> nn.Module:
|
||||
"""Install the quality-gated CUDA FLUX.2 VAE decoder fast paths."""
|
||||
from diffusers.models.attention_processor import Attention, AttnProcessor2_0
|
||||
from diffusers.models.autoencoders.vae import Decoder
|
||||
from diffusers.models.resnet import ResnetBlock2D
|
||||
from diffusers.models.upsampling import Upsample2D
|
||||
|
||||
from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_flux2 import (
|
||||
AutoencoderKLFlux2,
|
||||
)
|
||||
|
||||
if not isinstance(vae, AutoencoderKLFlux2) or type(vae.decoder) is not Decoder:
|
||||
return vae
|
||||
if getattr(vae, "_spatial_parallel_decode_enabled", False):
|
||||
logger.info(
|
||||
"FLUX.2 VAE: spatial-parallel decode enabled; "
|
||||
"skipping CUDA decoder fast paths."
|
||||
)
|
||||
return vae
|
||||
if not _HAS_TRITON:
|
||||
# aten GroupNorm is ~2x slower on NHWC tensors than on NCHW, so the
|
||||
# channels_last fast path is only a net win with the Triton
|
||||
# GroupNorm+SiLU fuse (measured: 97 -> 141 ms at 1024^2 with
|
||||
# channels_last alone vs 97 -> 29 ms with both).
|
||||
logger.warning(
|
||||
"FLUX.2 VAE: Triton unavailable; skipping CUDA decoder fast paths."
|
||||
)
|
||||
return vae
|
||||
|
||||
decoder = vae.decoder
|
||||
attn_modules = [
|
||||
m
|
||||
for m in decoder.modules()
|
||||
if _attn_fast_compatible(m, Attention, AttnProcessor2_0)
|
||||
]
|
||||
n_attn_total = sum(1 for m in decoder.modules() if isinstance(m, Attention))
|
||||
if len(attn_modules) != n_attn_total:
|
||||
# AttnProcessor2_0 `.view`s the 4D activation, which is illegal on
|
||||
# channels_last tensors; without a layout-safe rewrite for every
|
||||
# attention block the layout switch cannot be applied (fail closed).
|
||||
logger.warning(
|
||||
"FLUX.2 VAE: %d/%d attention blocks lack a layout-safe rewrite; "
|
||||
"skipping CUDA decoder fast paths.",
|
||||
n_attn_total - len(attn_modules),
|
||||
n_attn_total,
|
||||
)
|
||||
return vae
|
||||
|
||||
gate = VaeFastPathGate()
|
||||
decoder._sgl_gate = gate
|
||||
decoder._sgl_channels_last = False
|
||||
decoder.forward = MethodType(_decoder_layout_forward, decoder)
|
||||
n_up = _install_fused_upsample(decoder, Upsample2D, gate)
|
||||
for m in attn_modules:
|
||||
m._sgl_gate = gate
|
||||
m._sgl_folded_v = None
|
||||
m.forward = MethodType(_attn_fast_forward, m)
|
||||
n_norm = _install_norm_silu(decoder, ResnetBlock2D, gate)
|
||||
setattr(vae, GATE_ATTR, gate)
|
||||
logger.info(
|
||||
"FLUX.2 VAE: installed quality-gated decoder fast paths "
|
||||
"(channels_last dispatch, %d fused upsamplers, %d fast attention "
|
||||
"blocks, %d GroupNorm+SiLU fusions).",
|
||||
n_up,
|
||||
len(attn_modules),
|
||||
n_norm,
|
||||
)
|
||||
return vae
|
||||
@@ -297,34 +297,45 @@ class DecodingStage(PipelineStage):
|
||||
assert vae is not None
|
||||
self.vae = vae
|
||||
|
||||
frames = self.decode(batch.latents, server_args, vae_dtype=vae_dtype)
|
||||
# Request-scoped VAE fast-path gate (see flux2_vae_cuda_opt):
|
||||
# quality == "high" opts this decode into the near-lossless fast
|
||||
# paths; the "lossless" default keeps the bit-exact original
|
||||
# module path. VAEs without installed wrappers have no gate.
|
||||
gate = getattr(vae, "_sgl_vae_fast_path_gate", None)
|
||||
if gate is not None:
|
||||
gate.enabled = getattr(batch.sampling_params, "quality", None) == "high"
|
||||
try:
|
||||
frames = self.decode(batch.latents, server_args, vae_dtype=vae_dtype)
|
||||
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
assert (
|
||||
batch.trajectory_latents is not None
|
||||
), "batch should have trajectory latents"
|
||||
# decode trajectory latents if needed
|
||||
if batch.return_trajectory_decoded:
|
||||
assert (
|
||||
batch.trajectory_latents is not None
|
||||
), "batch should have trajectory latents"
|
||||
|
||||
# 1. Batch trajectory decoding to improve GPU utilization
|
||||
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
B, T, C, F, H, W = batch.trajectory_latents.shape
|
||||
flat_latents = batch.trajectory_latents.view(B * T, C, F, H, W)
|
||||
# 1. Batch trajectory decoding to improve GPU utilization
|
||||
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
|
||||
B, T, C, F, H, W = batch.trajectory_latents.shape
|
||||
flat_latents = batch.trajectory_latents.view(B * T, C, F, H, W)
|
||||
|
||||
logger.info("decoding %s trajectory latents in batch", B * T)
|
||||
# Use the optimized batch decode
|
||||
all_decoded = self.decode(
|
||||
flat_latents, server_args, vae_dtype=vae_dtype
|
||||
)
|
||||
logger.info("decoding %s trajectory latents in batch", B * T)
|
||||
# Use the optimized batch decode
|
||||
all_decoded = self.decode(
|
||||
flat_latents, server_args, vae_dtype=vae_dtype
|
||||
)
|
||||
|
||||
# 2. Reshape back
|
||||
# Keep on GPU to allow faster vectorized post-processing
|
||||
decoded_tensor = all_decoded.view(B, T, *all_decoded.shape[1:])
|
||||
# 2. Reshape back
|
||||
# Keep on GPU to allow faster vectorized post-processing
|
||||
decoded_tensor = all_decoded.view(B, T, *all_decoded.shape[1:])
|
||||
|
||||
# Convert to list of tensors (per timestep) as expected by OutputBatch
|
||||
# Each element in list is [B, channels, frames, H_out, W_out]
|
||||
trajectory_decoded = [decoded_tensor[:, i] for i in range(T)]
|
||||
else:
|
||||
trajectory_decoded = None
|
||||
# Convert to list of tensors (per timestep) as expected by OutputBatch
|
||||
# Each element in list is [B, channels, frames, H_out, W_out]
|
||||
trajectory_decoded = [decoded_tensor[:, i] for i in range(T)]
|
||||
else:
|
||||
trajectory_decoded = None
|
||||
finally:
|
||||
if gate is not None:
|
||||
gate.enabled = False
|
||||
|
||||
frames = server_args.pipeline_config.post_decoding(frames, server_args)
|
||||
|
||||
|
||||
@@ -537,6 +537,28 @@ class CudaPlatformBase(Platform):
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
return "sglang.multimodal_gen.runtime.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
|
||||
|
||||
@classmethod
|
||||
def optimize_vae(cls, vae: torch.nn.Module) -> torch.nn.Module:
|
||||
"""Install the quality-gated FLUX.2 VAE decoder fast paths.
|
||||
|
||||
Requests with quality == "high" run the fast paths; the "lossless"
|
||||
default runs the original module path bit-for-bit. See
|
||||
flux2_vae_cuda_opt for details.
|
||||
"""
|
||||
try:
|
||||
from sglang.multimodal_gen.runtime.models.vaes.flux2_vae_cuda_opt import (
|
||||
maybe_optimize_flux2_vae,
|
||||
)
|
||||
|
||||
vae = maybe_optimize_flux2_vae(vae)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to apply CUDA FLUX.2 VAE optimizations; "
|
||||
"using the unmodified VAE.",
|
||||
exc_info=True,
|
||||
)
|
||||
return vae
|
||||
|
||||
|
||||
# NVML utils
|
||||
# Note that NVML is not affected by `CUDA_VISIBLE_DEVICES`,
|
||||
|
||||
Reference in New Issue
Block a user