From 148f15b0af3e89e2361c83266446593ae153553c Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Sat, 8 Aug 2026 13:07:42 +0800 Subject: [PATCH] [diffusion] FLUX.1 fused adaLN modulate (bit-exact) + RoPE cache hoist, LN-affine folding behind quality=high (H200 e2e -3.5% lossless / -6.9% high) (#34004) --- .../csrc/diffusion/modulate_scale_shift.cuh | 221 ++++++++++++++++++ .../ops/diffusion/fused_ln_modulate.py | 85 +++++++ .../ops/diffusion/modulate_scale_shift.py | 95 ++++++++ .../runtime/models/dits/flux.py | 171 ++++++++++++-- .../pipelines_core/stages/denoising.py | 9 + .../ops/diffusion/test_fused_ln_modulate.py | 56 +++++ .../diffusion/test_modulate_scale_shift.py | 56 +++++ 7 files changed, 671 insertions(+), 22 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/diffusion/modulate_scale_shift.cuh create mode 100644 python/sglang/kernels/ops/diffusion/fused_ln_modulate.py create mode 100644 python/sglang/kernels/ops/diffusion/modulate_scale_shift.py create mode 100644 test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py create mode 100644 test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py diff --git a/python/sglang/kernels/jit/csrc/diffusion/modulate_scale_shift.cuh b/python/sglang/kernels/jit/csrc/diffusion/modulate_scale_shift.cuh new file mode 100644 index 000000000..1ad52063d --- /dev/null +++ b/python/sglang/kernels/jit/csrc/diffusion/modulate_scale_shift.cuh @@ -0,0 +1,221 @@ +// CUDA fast path for diffusion adaLN modulate chains. +// +// Implements, with each intermediate computed in fp32 and rounded to the +// storage dtype (the per-op kernel boundaries of the eager aten chain): +// +// out = (x * (1 + scale)) + shift +// = round(round(x * round(1 + scale)) + shift) +// +// so the fused kernel is bit-exact vs eager for fp16/bf16. x is a +// contiguous [B, L, D] activation; scale/shift are contiguous [B, D] +// modulation rows. +// +// Intentionally narrow: 16-byte aligned tensors, D % kVec == 0 (the Python +// guard enforces this). + +#pragma once + +#include // For host dtype helpers and TensorView metadata +#include // For RuntimeCheck and div_ceil + +#include // For DTypeTrait conversions +#include // For LaunchKernel and CUDA dtype aliases +#include // For device::AlignedVector + +#include + +namespace sglang_modulate_scale_shift { + +namespace { + +constexpr int kRowsPerBlock = 4; +constexpr int kColsVecPerBlock = 256; +constexpr int64_t kMaxGrid = 65535; + +inline const char* data_ptr(const tvm::ffi::TensorView& t) { + return static_cast(t.data_ptr()) + t.byte_offset(); +} + +inline char* mutable_data_ptr(const tvm::ffi::TensorView& t) { + return static_cast(t.data_ptr()) + t.byte_offset(); +} + +inline bool aligned16(const void* p) { + return (reinterpret_cast(p) & 0xF) == 0; +} + +inline int64_t numel(const tvm::ffi::TensorView& t) { + int64_t n = 1; + for (int i = 0; i < t.ndim(); ++i) { + n *= t.size(i); + } + return n; +} + +inline bool is_dense_contiguous(const tvm::ffi::TensorView& t) { + int64_t expected = 1; + for (int i = t.ndim() - 1; i >= 0; --i) { + if (t.size(i) == 1) { + continue; + } + if (t.stride(i) != expected) { + return false; + } + expected *= t.size(i); + } + return true; +} + +template +inline void check_dtype(const tvm::ffi::TensorView& t) { + host::RuntimeCheck(host::is_type(t.dtype()), "unexpected dtype for modulate_scale_shift"); +} + +template +__device__ __forceinline__ float to_float(T v) { + return static_cast(v); +} + +template <> +__device__ __forceinline__ float to_float(fp16_t v) { + return __half2float(v); +} + +template <> +__device__ __forceinline__ float to_float(bf16_t v) { + return __bfloat162float(v); +} + +template +__device__ __forceinline__ T modulate_value(T x, T scale, T shift) { + // Round each intermediate back to T (the eager chain's kernel boundaries; + // also blocks fmul+fadd FMA contraction). + const T one_plus_scale = DTypeTrait::from(1.0f + to_float(scale)); + const T product = DTypeTrait::from(to_float(x) * to_float(one_plus_scale)); + return DTypeTrait::from(to_float(product) + to_float(shift)); +} + +template +__global__ void modulate_scale_shift_vec_kernel( + const T* __restrict__ x, + const T* __restrict__ scale, + const T* __restrict__ shift, + T* __restrict__ out, + int64_t rows, + int64_t rows_per_batch, + int64_t row_vec) { + using Vec = device::AlignedVector; + const int64_t col_vec = static_cast(blockIdx.x) * kColsVecPerBlock + threadIdx.x; + if (col_vec >= row_vec) { + return; + } + + // Grid-stride: the row-tile count can exceed the gridDim.y hardware limit. + const int64_t row_tile_stride = static_cast(gridDim.y) * kRowsPerBlock; + for (int64_t row_base = static_cast(blockIdx.y) * kRowsPerBlock; row_base < rows; + row_base += row_tile_stride) { +#pragma unroll + for (int row_offset = 0; row_offset < kRowsPerBlock; ++row_offset) { + const int64_t row = row_base + row_offset; + if (row < rows) { + const int64_t batch = row / rows_per_batch; + const int64_t mod_v = batch * row_vec + col_vec; + const int64_t v = row * row_vec + col_vec; + Vec xv, s, b, o; + s.load(scale, mod_v); + b.load(shift, mod_v); + xv.load(x, v); +#pragma unroll + for (int i = 0; i < kVec; ++i) { + o[i] = modulate_value(xv[i], s[i], b[i]); + } + o.store(out, v); + } + } + } +} + +template +inline void launch_modulate_scale_shift( + const tvm::ffi::TensorView& out, + const tvm::ffi::TensorView& x, + const tvm::ffi::TensorView& scale, + const tvm::ffi::TensorView& shift) { + const int64_t total = numel(x); + if (total == 0) { + return; + } + + const int64_t D = x.size(x.ndim() - 1); + const int64_t rows = total / D; + const int64_t batches = scale.size(0); + const int64_t rows_per_batch = rows / batches; + const T* x_ptr = reinterpret_cast(data_ptr(x)); + const T* scale_ptr = reinterpret_cast(data_ptr(scale)); + const T* shift_ptr = reinterpret_cast(data_ptr(shift)); + T* out_ptr = reinterpret_cast(mutable_data_ptr(out)); + constexpr int kVec = 16 / sizeof(T); + + host::RuntimeCheck( + aligned16(x_ptr) && aligned16(scale_ptr) && aligned16(shift_ptr) && aligned16(out_ptr), + "modulate_scale_shift requires 16-byte aligned tensors"); + host::RuntimeCheck(D % kVec == 0, "modulate_scale_shift requires D to be a multiple of the vector width"); + + const int64_t row_vec = D / kVec; + const int64_t col_blocks = host::div_ceil(row_vec, static_cast(kColsVecPerBlock)); + const int64_t row_tiles = host::div_ceil(rows, static_cast(kRowsPerBlock)); + const int64_t row_blocks = row_tiles > kMaxGrid ? kMaxGrid : row_tiles; + host::LaunchKernel( + dim3(static_cast(col_blocks), static_cast(row_blocks)), dim3(kColsVecPerBlock), out.device())( + modulate_scale_shift_vec_kernel, x_ptr, scale_ptr, shift_ptr, out_ptr, rows, rows_per_batch, row_vec); +} + +template +inline void validate_modulate_scale_shift( + const tvm::ffi::TensorView& out, + const tvm::ffi::TensorView& x, + const tvm::ffi::TensorView& scale, + const tvm::ffi::TensorView& shift) { + check_dtype(out); + check_dtype(x); + check_dtype(scale); + check_dtype(shift); + host::RuntimeCheck(x.device().device_type == kDLCUDA, "x must be CUDA"); + host::RuntimeCheck(scale.device().device_type == kDLCUDA, "scale must be CUDA"); + host::RuntimeCheck(shift.device().device_type == kDLCUDA, "shift must be CUDA"); + host::RuntimeCheck(out.device().device_type == kDLCUDA, "out must be CUDA"); + host::RuntimeCheck( + x.device().device_id == scale.device().device_id && x.device().device_id == shift.device().device_id && + x.device().device_id == out.device().device_id, + "x/scale/shift/out must be on the same CUDA device"); + host::RuntimeCheck(x.ndim() == 3, "x must be [B, L, D]"); + host::RuntimeCheck(scale.ndim() == 2, "scale must be [B, D]"); + host::RuntimeCheck(shift.ndim() == 2, "shift must be [B, D]"); + host::RuntimeCheck(out.ndim() == x.ndim(), "out rank must match x"); + for (int i = 0; i < x.ndim(); ++i) { + host::RuntimeCheck(out.size(i) == x.size(i), "out shape must match x"); + } + host::RuntimeCheck(scale.size(0) == x.size(0), "scale batch dim must match x"); + host::RuntimeCheck(scale.size(1) == x.size(2), "scale last dim must match x"); + host::RuntimeCheck(shift.size(0) == scale.size(0) && shift.size(1) == scale.size(1), "shift shape must match scale"); + host::RuntimeCheck(is_dense_contiguous(x), "x must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(scale), "scale must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(shift), "shift must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(out), "out must be contiguous"); + host::RuntimeCheck(data_ptr(out) != data_ptr(x), "out must not alias x"); + host::RuntimeCheck(data_ptr(out) != data_ptr(scale), "out must not alias scale"); + host::RuntimeCheck(data_ptr(out) != data_ptr(shift), "out must not alias shift"); +} + +} // namespace + +template +struct ModulateScaleShiftKernel { + static void + run(tvm::ffi::TensorView out, tvm::ffi::TensorView x, tvm::ffi::TensorView scale, tvm::ffi::TensorView shift) { + validate_modulate_scale_shift(out, x, scale, shift); + launch_modulate_scale_shift(out, x, scale, shift); + } +}; + +} // namespace sglang_modulate_scale_shift diff --git a/python/sglang/kernels/ops/diffusion/fused_ln_modulate.py b/python/sglang/kernels/ops/diffusion/fused_ln_modulate.py new file mode 100644 index 000000000..7383596cd --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/fused_ln_modulate.py @@ -0,0 +1,85 @@ +"""LayerNorm + adaLN modulate folded into one affine LN call. + +``layer_norm(x, weight=(1 + scale), bias=shift)`` replaces the affine-free +LayerNorm + modulate pair: one kernel and one HBM pass per site instead of +two. ``1 + scale`` keeps the eager rounding of the [1, D] modulation row, +but scale/shift then apply in fp32 to the *unrounded* normalized value, so +the result is not bit-exact vs the reference (half-precision rounding-order +differences only). + +Because it is not bit-exact the fold is opt-in per batch: model code marks +its LN+modulate sites with :func:`mark_fused_ln_modulate_site` (default off, +reference path), and the denoising stage calls +:func:`mount_fused_ln_modulate` / :func:`unmount_fused_ln_modulate` at batch +boundaries for ``quality="high"`` requests. +""" + +from __future__ import annotations + +from typing import Iterator + +import torch +import torch.nn.functional as F +from torch import nn + +_SITE_ENABLED_ATTR = "_sgl_fused_ln_modulate_enabled" + +_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16) + + +def mark_fused_ln_modulate_site(module: nn.Module) -> None: + """Mark ``module`` as an LN+modulate fusion site (mounted off).""" + setattr(module, _SITE_ENABLED_ATTR, False) + + +def fused_ln_modulate_active(module: nn.Module) -> bool: + return getattr(module, _SITE_ENABLED_ATTR, False) + + +def iter_fused_ln_modulate_sites(root: nn.Module) -> Iterator[nn.Module]: + for module in root.modules(): + if hasattr(module, _SITE_ENABLED_ATTR): + yield module + + +def mount_fused_ln_modulate(root: nn.Module) -> bool: + sites = list(iter_fused_ln_modulate_sites(root)) + for site in sites: + setattr(site, _SITE_ENABLED_ATTR, True) + return bool(sites) + + +def unmount_fused_ln_modulate(root: nn.Module) -> None: + for site in iter_fused_ln_modulate_sites(root): + setattr(site, _SITE_ENABLED_ATTR, False) + + +def can_fuse_ln_modulate( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> bool: + """Per-call guard: the folded affine is a [D] row, so batch must be 1.""" + return ( + x.is_cuda + and x.dtype in _SUPPORTED_DTYPES + and scale.dtype == x.dtype + and shift.dtype == x.dtype + and x.dim() == 3 + and x.shape[0] == 1 + and scale.dim() == 2 + and scale.shape == shift.shape + and scale.shape == (1, x.shape[-1]) + and x.numel() > 0 + ) + + +def fused_ln_modulate( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor, eps: float +) -> torch.Tensor: + """``layer_norm(x) * (1 + scale) + shift`` as one affine-folded LN kernel.""" + return F.layer_norm( + x, + (x.shape[-1],), + weight=(1 + scale).reshape(-1), + bias=shift.reshape(-1), + eps=eps, + ) diff --git a/python/sglang/kernels/ops/diffusion/modulate_scale_shift.py b/python/sglang/kernels/ops/diffusion/modulate_scale_shift.py new file mode 100644 index 000000000..8a635e79c --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/modulate_scale_shift.py @@ -0,0 +1,95 @@ +"""Fused adaLN modulate: ``x * (1 + scale) + shift`` in one CUDA kernel. + +Numerical contract: the kernel reproduces each eager op's +fp32-opmath/round-to-storage-dtype boundary (fp16/bf16), so its output is +bit-exact vs the eager chain (``torch.equal``) and needs no quality gate. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import cache_once, load_jit, make_cpp_args +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16) +_ALIGN_BYTES = 16 + + +@cache_once +def _jit_modulate_scale_shift_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "diffusion_modulate_scale_shift", + *args, + cuda_files=["diffusion/modulate_scale_shift.cuh"], + cuda_wrappers=[ + ( + "modulate_scale_shift", + "sglang_modulate_scale_shift::" + f"ModulateScaleShiftKernel<{args}>::run", + ), + ], + ) + + +def _fake_impl( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + return torch.empty_like(x) + + +@register_custom_op( + op_name="diffusion_modulate_scale_shift", + mutates_args=[], + fake_impl=_fake_impl, +) +def _modulate_scale_shift_custom_op( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + out = torch.empty_like(x) + module = _jit_modulate_scale_shift_module(x.dtype) + module.modulate_scale_shift(out, x, scale, shift) + return out + + +def _aligned(t: torch.Tensor) -> bool: + return t.data_ptr() % _ALIGN_BYTES == 0 + + +def can_use_modulate_scale_shift_cuda( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> bool: + if ( + x.dtype not in _SUPPORTED_DTYPES + or scale.dtype != x.dtype + or shift.dtype != x.dtype + or not (x.is_cuda and scale.is_cuda and shift.is_cuda) + or not (x.device == scale.device == shift.device) + or x.dim() != 3 + or scale.dim() != 2 + or shift.shape != scale.shape + or scale.shape != (x.shape[0], x.shape[-1]) + or not (x.is_contiguous() and scale.is_contiguous() and shift.is_contiguous()) + or x.numel() == 0 + ): + return False + vec = _ALIGN_BYTES // x.element_size() + return ( + x.shape[-1] % vec == 0 and _aligned(x) and _aligned(scale) and _aligned(shift) + ) + + +def modulate_scale_shift_cuda( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + """Fused ``x * (1 + scale[:, None]) + shift[:, None]`` (bit-exact vs eager).""" + if not can_use_modulate_scale_shift_cuda(x, scale, shift): + raise RuntimeError("unsupported input for modulate_scale_shift CUDA") + return _modulate_scale_shift_custom_op(x, scale, shift) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index 5aa3158a1..4ebf7ac40 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -33,6 +33,16 @@ from sglang.kernels.ops.diffusion.fused_linear_gelu import ( fused_linear_gelu_tanh, mark_fused_gelu_site, ) +from sglang.kernels.ops.diffusion.fused_ln_modulate import ( + can_fuse_ln_modulate, + fused_ln_modulate, + fused_ln_modulate_active, + mark_fused_ln_modulate_site, +) +from sglang.kernels.ops.diffusion.modulate_scale_shift import ( + can_use_modulate_scale_shift_cuda, + modulate_scale_shift_cuda, +) from sglang.kernels.ops.diffusion.residual_gate_add import ( can_use_residual_gate_add_cuda, residual_gate_add_cuda, @@ -121,6 +131,115 @@ def _flux_residual_gate_add( return residual + gate * update +_FLUX_MODULATE_CUDA_DISABLED = False + + +def _flux_modulate( + x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor +) -> torch.Tensor: + """``x * (1 + scale[:, None]) + shift[:, None]`` in one CUDA kernel. + + The kernel keeps the eager chain's per-op fp32-opmath/round-to-storage + boundaries, so it is bit-exact vs eager and needs no quality gate. + Guarded inputs fall back to the eager expression. + """ + global _FLUX_MODULATE_CUDA_DISABLED + + if not _FLUX_MODULATE_CUDA_DISABLED and can_use_modulate_scale_shift_cuda( + x, scale, shift + ): + try: + return modulate_scale_shift_cuda(x, scale, shift) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling FLUX modulate CUDA fast path: {exc}") + _FLUX_MODULATE_CUDA_DISABLED = True + + return x * (1 + scale[:, None]) + shift[:, None] + + +def _flux_norm_modulate( + site: nn.Module, + norm: nn.Module, + x: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> 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). + """ + if fused_ln_modulate_active(site) and can_fuse_ln_modulate(x, scale, shift): + return fused_ln_modulate(x, scale, shift, norm.eps) + return _flux_modulate(norm(x), scale, shift) + + +class FluxAdaLayerNormZero(AdaLayerNormZero): + """diffusers ``AdaLayerNormZero`` with the modulate routed through + :func:`_flux_norm_modulate`; parameters match the parent.""" + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + mark_fused_ln_modulate_site(self) + + def forward( + self, + x: torch.Tensor, + timestep: Optional[torch.Tensor] = None, + class_labels: Optional[torch.LongTensor] = None, + hidden_dtype: Optional[torch.dtype] = None, + emb: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if self.emb is not None: + emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk( + 6, dim=1 + ) + x = _flux_norm_modulate(self, self.norm, x, scale_msa, shift_msa) + return x, gate_msa, shift_mlp, scale_mlp, gate_mlp + + +class FluxAdaLayerNormZeroSingle(AdaLayerNormZeroSingle): + """diffusers ``AdaLayerNormZeroSingle`` with the modulate routed through + :func:`_flux_norm_modulate`; parameters match the parent.""" + + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + mark_fused_ln_modulate_site(self) + + def forward( + self, + x: torch.Tensor, + emb: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1) + x = _flux_norm_modulate(self, self.norm, x, scale_msa, shift_msa) + return x, gate_msa + + +def _rope_cos_sin_cache( + freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None], +) -> Optional[torch.Tensor]: + """Concatenate a ``(cos, sin)`` RoPE tuple into the fp32 cache layout that + ``apply_qk_norm_with_optional_rope`` consumes; a prebuilt cache tensor + passes through unchanged.""" + if freqs_cis is None or isinstance(freqs_cis, torch.Tensor): + return freqs_cis + cos, sin = freqs_cis + return torch.cat( + [ + cos.to(dtype=torch.float32).contiguous(), + sin.to(dtype=torch.float32).contiguous(), + ], + dim=-1, + ) + + try: from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import] from nunchaku.models.normalization import ( # type: ignore[import] @@ -554,16 +673,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): query = query.unflatten(-1, (num_heads, -1)) key = key.unflatten(-1, (num_heads, -1)) value = value.unflatten(-1, (num_heads, -1)) - cos_sin_cache = None - if freqs_cis is not None: - cos, sin = freqs_cis - cos_sin_cache = torch.cat( - [ - cos.to(dtype=torch.float32).contiguous(), - sin.to(dtype=torch.float32).contiguous(), - ], - dim=-1, - ) + # Raw (cos, sin) tuple, or the cache prebuilt by the transformer forward. + cos_sin_cache = _rope_cos_sin_cache(freqs_cis) if self.added_kv_proj_dim is not None: encoder_query = encoder_query.unflatten(-1, (num_heads, -1)) @@ -658,7 +769,7 @@ class FluxSingleTransformerBlock(nn.Module): self.local_mlp_hidden_dim = divide(self.mlp_hidden_dim, self.tp_size) self.local_dim = divide(dim, self.tp_size) - self.norm = AdaLayerNormZeroSingle(dim) + self.norm = FluxAdaLayerNormZeroSingle(dim) if self.use_nunchaku_structure: self.mlp_fc1 = ColumnParallelLinear( @@ -765,7 +876,7 @@ class FluxSingleTransformerBlock(nn.Module): hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor, - freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, num_replicated_prefix: int = 0, ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -849,8 +960,8 @@ class FluxTransformerBlock(nn.Module): ): super().__init__() - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) + self.norm1 = FluxAdaLayerNormZero(dim) + self.norm1_context = FluxAdaLayerNormZero(dim) self.attn = FluxAttention( query_dim=dim, @@ -867,6 +978,9 @@ class FluxTransformerBlock(nn.Module): self.norm2 = LayerNorm(dim, eps=1e-6, elementwise_affine=False) self.norm2_context = LayerNorm(dim, eps=1e-6, elementwise_affine=False) + # quality="high" site: the norm2/norm2_context modulate folds into the + # LN affine when mounted. + mark_fused_ln_modulate_site(self) nunchaku_enabled = ( quant_config is not None @@ -928,7 +1042,7 @@ class FluxTransformerBlock(nn.Module): hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, temb: torch.Tensor, - freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, num_replicated_prefix: int = 0, ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -963,14 +1077,14 @@ class FluxTransformerBlock(nn.Module): hidden_states = _flux_residual_gate_add( hidden_states, attn_output, gate_msa.unsqueeze(1) ) - norm_hidden_states = self.norm2(hidden_states) if self.use_nunchaku_structure: + norm_hidden_states = self.norm2(hidden_states) norm_hidden_states = ( norm_hidden_states * scale_mlp[:, None] + shift_mlp[:, None] ) else: - norm_hidden_states = ( - norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + norm_hidden_states = _flux_norm_modulate( + self, self.norm2, hidden_states, scale_mlp, shift_mlp ) ff_output = self.ff(norm_hidden_states) @@ -985,15 +1099,18 @@ class FluxTransformerBlock(nn.Module): encoder_hidden_states, context_attn_output, c_gate_msa.unsqueeze(1) ) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) if self.use_nunchaku_structure: + norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) norm_encoder_hidden_states = ( norm_encoder_hidden_states * c_scale_mlp[:, None] + c_shift_mlp[:, None] ) else: - norm_encoder_hidden_states = ( - norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) - + c_shift_mlp[:, None] + norm_encoder_hidden_states = _flux_norm_modulate( + self, + self.norm2_context, + encoder_hidden_states, + c_scale_mlp, + c_shift_mlp, ) context_ff_output = self.ff_context(norm_encoder_hidden_states) @@ -1233,6 +1350,16 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): join_seqs(sin[:t_loc], sin[t_loc:], pad, dim=0), ) + # Build the RoPE cos/sin cache once per step; every attention call + # below reuses the same tensor. + hoisted_freqs_cis = _rope_cos_sin_cache(freqs_cis) + singles_freqs_cis = ( + hoisted_freqs_cis + if singles_freqs_cis is freqs_cis + else _rope_cos_sin_cache(singles_freqs_cis) + ) + freqs_cis = hoisted_freqs_cis + if ( joint_attention_kwargs is not None and "ip_adapter_image_embeds" in joint_attention_kwargs diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 36196c959..a53ed33e2 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -28,6 +28,10 @@ from sglang.kernels.ops.diffusion.fused_linear_gelu import ( mount_fused_linear_gelu, unmount_fused_linear_gelu, ) +from sglang.kernels.ops.diffusion.fused_ln_modulate import ( + mount_fused_ln_modulate, + unmount_fused_ln_modulate, +) from sglang.multimodal_gen import envs from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode from sglang.multimodal_gen.configs.pipeline_configs.flux import ( @@ -478,18 +482,23 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): return mounted_gelu = False mounted_gate_norm = False + mounted_ln_modulate = False for transformer in filter(None, [self.transformer, self.transformer_2]): if want: mounted_gelu |= mount_fused_linear_gelu(transformer) mounted_gate_norm |= mount_fused_gate_rmsnorm(transformer) + mounted_ln_modulate |= mount_fused_ln_modulate(transformer) else: unmount_fused_linear_gelu(transformer) unmount_fused_gate_rmsnorm(transformer) + unmount_fused_ln_modulate(transformer) self._quality_fusions_mounted = want if want and mounted_gelu: logger.info( "Mounted fused linear+GELU (cublasLt epilogue) for quality=high" ) + if want and mounted_ln_modulate: + logger.info("Mounted fused LN+modulate (affine folding) for quality=high") if want and mounted_gate_norm: logger.info( "Mounted fused gate RMSNorm (Z-Image Triton suite) for quality=high" diff --git a/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py new file mode 100644 index 000000000..049350ddf --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py @@ -0,0 +1,56 @@ +import pytest +import torch +import torch.nn as nn + +from sglang.kernels.ops.diffusion.fused_ln_modulate import ( + can_fuse_ln_modulate, + fused_ln_modulate, + fused_ln_modulate_active, + mark_fused_ln_modulate_site, + mount_fused_ln_modulate, + unmount_fused_ln_modulate, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +@pytest.fixture(autouse=True) +def cuda_setup(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + torch.cuda.manual_seed(0) + + +@pytest.mark.parametrize("seq_len", [4096, 512]) +def test_fused_ln_modulate_matches_reference(seq_len): + x = torch.randn((1, seq_len, 3072), device="cuda", dtype=torch.bfloat16) + scale = torch.randn((1, 3072), device="cuda", dtype=torch.bfloat16) + shift = torch.randn_like(scale) + assert can_fuse_ln_modulate(x, scale, shift) + out = fused_ln_modulate(x, scale, shift, eps=1e-6) + norm = nn.LayerNorm(3072, eps=1e-6, elementwise_affine=False).cuda() + ref = norm(x) * (1 + scale[:, None]) + shift[:, None] + # Contract: bf16 rounding-order-level difference only, not bit-exact. + torch.testing.assert_close(out, ref, atol=0.0625, rtol=0.05) + + +def test_fused_ln_modulate_guards_and_mount_protocol(): + x = torch.randn((2, 64, 3072), device="cuda", dtype=torch.bfloat16) + row = torch.randn((2, 3072), device="cuda", dtype=torch.bfloat16) + assert not can_fuse_ln_modulate(x, row, row) # folded affine needs B == 1 + root = nn.Module() + root.child = nn.Module() + mark_fused_ln_modulate_site(root.child) + assert not fused_ln_modulate_active(root.child) + assert mount_fused_ln_modulate(root) + assert fused_ln_modulate_active(root.child) + unmount_fused_ln_modulate(root) + assert not fused_ln_modulate_active(root.child) + assert not mount_fused_ln_modulate(nn.Module()) # no marked sites + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py b/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py new file mode 100644 index 000000000..a0aa6d8e8 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py @@ -0,0 +1,56 @@ +import pytest +import torch + +from sglang.kernels.ops.diffusion.modulate_scale_shift import ( + can_use_modulate_scale_shift_cuda, + modulate_scale_shift_cuda, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +# FLUX.1 1024^2 adaLN shapes (D=3072) plus batched and odd-length coverage. +CASES = [(1, 4096, 3072), (1, 512, 3072), (1, 4608, 3072), (2, 1024, 3072), (1, 17, 64)] + + +@pytest.fixture(autouse=True) +def cuda_setup(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + torch.cuda.manual_seed(0) + + +def _eager(x, scale, shift): + return x * (1 + scale[:, None]) + shift[:, None] + + +@pytest.mark.parametrize("shape", CASES) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_modulate_scale_shift_matches_eager(shape, dtype): + x = torch.randn(shape, device="cuda", dtype=dtype) + scale = torch.randn((shape[0], shape[-1]), device="cuda", dtype=dtype) + shift = torch.randn_like(scale) + out = modulate_scale_shift_cuda(x, scale, shift) + assert torch.equal(out, _eager(x, scale, shift)) # bitwise contract + + +def test_modulate_scale_shift_adaln_chunk_views(): + x = torch.randn((1, 4096, 3072), device="cuda", dtype=torch.bfloat16) + emb = torch.randn((1, 6 * 3072), device="cuda", dtype=torch.bfloat16) + shift, scale = emb.chunk(6, dim=1)[:2] + assert can_use_modulate_scale_shift_cuda(x, scale, shift) + out = modulate_scale_shift_cuda(x, scale, shift) + assert torch.equal(out, _eager(x, scale, shift)) + + +def test_modulate_scale_shift_guards_reject_fp32(): + x = torch.randn((1, 64, 64), device="cuda", dtype=torch.float32) + row = torch.randn((1, 64), device="cuda", dtype=torch.float32) + assert not can_use_modulate_scale_shift_cuda(x, row, row) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__]))