From 1c3ad92438ac6c1064f0ef1b8fdb675ecb514993 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Tue, 1 Sep 2026 16:14:29 +0800 Subject: [PATCH] [Diffusion] Fuse FLUX.2 ModelOpt FP8 producers and QKV packing (#37162) Co-authored-by: Cursor --- .../jit/csrc/diffusion/flux2_qkv_epilogue.cuh | 261 +++++++++ .../sglang/kernels/ops/diffusion/__init__.py | 18 + .../layout/flux2_token_cat_fp8_triton.py | 143 +++++ .../norm/layernorm_modulate_triton.py | 75 ++- .../diffusion/rope/flux2_qkv_epilogue_jit.py | 134 +++++ .../runtime/models/dits/flux_2.py | 497 ++++++++++++++---- .../component_accuracy/engine.py | 4 +- ...test_component_accuracy_weight_transfer.py | 36 ++ ...est_modelopt_fp8_layerwise_offload_load.py | 5 +- .../test/unit/test_transformer_quant.py | 52 ++ .../diffusion/bench_flux2_fp8_norm_quant.py | 48 ++ .../diffusion/bench_flux2_token_cat_fp8.py | 46 ++ .../diffusion/test_flux2_fp8_norm_quant.py | 52 ++ .../ops/diffusion/test_flux2_qkv_epilogue.py | 158 ++++++ .../ops/diffusion/test_flux2_token_cat_fp8.py | 56 ++ .../models/test_flux2_fp8_norm_quant_gate.py | 90 ++++ 16 files changed, 1571 insertions(+), 104 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/diffusion/flux2_qkv_epilogue.cuh create mode 100644 python/sglang/kernels/ops/diffusion/layout/flux2_token_cat_fp8_triton.py create mode 100644 python/sglang/kernels/ops/diffusion/rope/flux2_qkv_epilogue_jit.py create mode 100644 test/registered/kernels/benchmark/diffusion/bench_flux2_fp8_norm_quant.py create mode 100644 test/registered/kernels/benchmark/diffusion/bench_flux2_token_cat_fp8.py create mode 100644 test/registered/kernels/ops/diffusion/test_flux2_fp8_norm_quant.py create mode 100644 test/registered/kernels/ops/diffusion/test_flux2_qkv_epilogue.py create mode 100644 test/registered/kernels/ops/diffusion/test_flux2_token_cat_fp8.py create mode 100644 test/registered/unit/models/test_flux2_fp8_norm_quant_gate.py diff --git a/python/sglang/kernels/jit/csrc/diffusion/flux2_qkv_epilogue.cuh b/python/sglang/kernels/jit/csrc/diffusion/flux2_qkv_epilogue.cuh new file mode 100644 index 000000000..cee132a02 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/diffusion/flux2_qkv_epilogue.cuh @@ -0,0 +1,261 @@ +#pragma once + +#include + +#include +#include +#include +#include +#include + +#include +#include + +namespace sglang { + +namespace flux2_qkv_epilogue { + +constexpr int kHeadDim = 128; +constexpr int kThreads = 256; +constexpr int kWarps = kThreads / device::kWarpThreads; +constexpr int kElemsPerThread = kHeadDim / device::kWarpThreads; +constexpr int kVecSize = kElemsPerThread / 2; + +struct Params { + void* joint_q; + void* joint_k; + void* joint_v; + const void* img_q; + const void* img_k; + const void* img_v; + const void* txt_q; + const void* txt_k; + const void* txt_v; + const void* img_q_weight; + const void* img_k_weight; + const void* txt_q_weight; + const void* txt_k_weight; + const void* cos_sin_cache; + int64_t input_token_stride_bytes; + int64_t output_token_stride_bytes; + int64_t head_stride_bytes; + uint32_t img_tokens; + uint32_t txt_tokens; + uint32_t num_heads; + float img_eps; + float txt_eps; +}; + +__global__ void flux2_qkv_epilogue_kernel(const Params __grid_constant__ params) { + using namespace device; + using Packed = packed_t; + using Storage = AlignedVector; + + const uint32_t lane = threadIdx.x % kWarpThreads; + const uint32_t warp = threadIdx.x / kWarpThreads; + const uint32_t start = blockIdx.x * kWarps + warp; + const uint32_t workers = gridDim.x * kWarps; + const uint32_t total_tokens = params.txt_tokens + params.img_tokens; + const uint32_t token_head_works = total_tokens * params.num_heads; + const uint32_t total_works = 3 * token_head_works; + + for (uint32_t work = start; work < total_works; work += workers) { + const uint32_t kind = work / token_head_works; // 0: Q, 1: K, 2: V. + const uint32_t token_head = work % token_head_works; + const uint32_t joint_token = token_head / params.num_heads; + const uint32_t head = token_head % params.num_heads; + const bool is_text = joint_token < params.txt_tokens; + const uint32_t source_token = is_text ? joint_token : joint_token - params.txt_tokens; + + const void* input_base; + void* output_base; + if (kind == 0) { + input_base = is_text ? params.txt_q : params.img_q; + output_base = params.joint_q; + } else if (kind == 1) { + input_base = is_text ? params.txt_k : params.img_k; + output_base = params.joint_k; + } else { + input_base = is_text ? params.txt_v : params.img_v; + output_base = params.joint_v; + } + + const void* input = + pointer::offset(input_base, source_token * params.input_token_stride_bytes, head * params.head_stride_bytes); + void* output = + pointer::offset(output_base, joint_token * params.output_token_stride_bytes, head * params.head_stride_bytes); + + auto input_vec = load_as(input, lane); + if (kind == 2) { + store_as(output, input_vec, lane); + continue; + } + + const void* weight_base; + if (kind == 0) { + weight_base = is_text ? params.txt_q_weight : params.img_q_weight; + } else { + weight_base = is_text ? params.txt_k_weight : params.img_k_weight; + } + const auto weight_vec = load_as(weight_base, lane); + + float elems[kElemsPerThread]; + float sum_of_squares = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const auto [x0, x1] = cast(input_vec[j]); + elems[2 * j] = x0; + elems[2 * j + 1] = x1; + sum_of_squares += x0 * x0 + x1 * x1; + } + sum_of_squares = warp::reduce_sum(sum_of_squares); + const float eps = is_text ? params.txt_eps : params.img_eps; + const float norm_factor = math::rsqrt(sum_of_squares / static_cast(kHeadDim) + eps); + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const auto [w0, w1] = cast(weight_vec[j]); + elems[2 * j] *= norm_factor * w0; + elems[2 * j + 1] *= norm_factor * w1; + } + + const auto* cache = static_cast(params.cos_sin_cache); + const auto* cos_ptr = cache + joint_token * kHeadDim; + const auto* sin_ptr = cos_ptr + kHeadDim / 2; +#pragma unroll + for (uint32_t i = 0; i < kElemsPerThread; i += 2) { + const float x = elems[i]; + const float y = elems[i + 1]; + const uint32_t cache_idx = (lane * kElemsPerThread + i) / 2; + const float cos = __ldg(cos_ptr + cache_idx); + const float sin = __ldg(sin_ptr + cache_idx); + elems[i] = x * cos - y * sin; + elems[i + 1] = y * cos + x * sin; + } + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + input_vec[j] = cast({elems[2 * j], elems[2 * j + 1]}); + } + store_as(output, input_vec, lane); + } +} + +struct Flux2QKVEpilogueKernel { + static void + run(tvm::ffi::TensorView joint_q, + tvm::ffi::TensorView joint_k, + tvm::ffi::TensorView joint_v, + tvm::ffi::TensorView img_q, + tvm::ffi::TensorView img_k, + tvm::ffi::TensorView img_v, + tvm::ffi::TensorView txt_q, + tvm::ffi::TensorView txt_k, + tvm::ffi::TensorView txt_v, + tvm::ffi::TensorView img_q_weight, + tvm::ffi::TensorView img_k_weight, + tvm::ffi::TensorView txt_q_weight, + tvm::ffi::TensorView txt_k_weight, + tvm::ffi::TensorView cos_sin_cache, + double img_eps, + double txt_eps) { + using namespace host; + + auto NI = SymbolicSize{"img_tokens"}; + auto NT = SymbolicSize{"txt_tokens"}; + auto N = SymbolicSize{"joint_tokens"}; + auto H = SymbolicSize{"num_heads"}; + auto D = SymbolicSize{"head_dim"}; + auto device = SymbolicDevice{}; + D.set_value(kHeadDim); + device.set_options(); + + TensorMatcher({NI, H, D}) + .with_strides({-1, D, 1}) + .with_dtype() + .with_device(device) + .verify(img_q) + .verify(img_k) + .verify(img_v); + TensorMatcher({NT, H, D}) + .with_strides({-1, D, 1}) + .with_dtype() + .with_device(device) + .verify(txt_q) + .verify(txt_k) + .verify(txt_v); + N.set_value(NI.unwrap() + NT.unwrap()); + TensorMatcher({N, H, D}).with_dtype().with_device(device).verify(joint_q).verify(joint_k).verify(joint_v); + TensorMatcher({D}) + .with_dtype() + .with_device(device) + .verify(img_q_weight) + .verify(img_k_weight) + .verify(txt_q_weight) + .verify(txt_k_weight); + TensorMatcher({-1, D}).with_dtype().with_device(device).verify(cos_sin_cache); + + RuntimeCheck( + img_q.stride(0) == img_k.stride(0) && img_q.stride(0) == img_v.stride(0), + "image QKV inputs must use the same token stride"); + RuntimeCheck( + txt_q.stride(0) == txt_k.stride(0) && txt_q.stride(0) == txt_v.stride(0), + "text QKV inputs must use the same token stride"); + RuntimeCheck(img_q.stride(0) == txt_q.stride(0), "image/text QKV token strides must match"); + RuntimeCheck( + img_q.stride(1) == kHeadDim && img_k.stride(1) == kHeadDim && img_v.stride(1) == kHeadDim, + "image QKV heads must be contiguous"); + RuntimeCheck( + txt_q.stride(1) == kHeadDim && txt_k.stride(1) == kHeadDim && txt_v.stride(1) == kHeadDim, + "text QKV heads must be contiguous"); + RuntimeCheck(joint_q.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(joint_k.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(joint_v.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(cos_sin_cache.is_contiguous(), "cos/sin cache must be contiguous"); + RuntimeCheck(cos_sin_cache.size(0) >= N.unwrap(), "cos/sin cache does not cover all joint tokens"); + + const uint32_t img_tokens = static_cast(NI.unwrap()); + const uint32_t txt_tokens = static_cast(NT.unwrap()); + const uint32_t num_heads = static_cast(H.unwrap()); + const uint32_t total_works = 3 * (img_tokens + txt_tokens) * num_heads; + if (total_works == 0) return; + + const int64_t head_stride_bytes = kHeadDim * sizeof(bf16_t); + const int64_t input_token_stride_bytes = img_q.stride(0) * sizeof(bf16_t); + const int64_t output_token_stride_bytes = num_heads * head_stride_bytes; + const auto params = Params{ + .joint_q = joint_q.data_ptr(), + .joint_k = joint_k.data_ptr(), + .joint_v = joint_v.data_ptr(), + .img_q = img_q.data_ptr(), + .img_k = img_k.data_ptr(), + .img_v = img_v.data_ptr(), + .txt_q = txt_q.data_ptr(), + .txt_k = txt_k.data_ptr(), + .txt_v = txt_v.data_ptr(), + .img_q_weight = img_q_weight.data_ptr(), + .img_k_weight = img_k_weight.data_ptr(), + .txt_q_weight = txt_q_weight.data_ptr(), + .txt_k_weight = txt_k_weight.data_ptr(), + .cos_sin_cache = cos_sin_cache.data_ptr(), + .input_token_stride_bytes = input_token_stride_bytes, + .output_token_stride_bytes = output_token_stride_bytes, + .head_stride_bytes = head_stride_bytes, + .img_tokens = img_tokens, + .txt_tokens = txt_tokens, + .num_heads = num_heads, + .img_eps = static_cast(img_eps), + .txt_eps = static_cast(txt_eps), + }; + + const uint32_t sm_count = runtime::get_sm_count(device.unwrap().device_id); + static const uint32_t blocks_per_sm = runtime::get_blocks_per_sm(flux2_qkv_epilogue_kernel, kThreads); + const uint32_t needed_blocks = div_ceil(total_works, uint32_t(kWarps)); + const uint32_t blocks = std::min(blocks_per_sm * sm_count, needed_blocks); + LaunchKernel(blocks, kThreads, device.unwrap())(flux2_qkv_epilogue_kernel, params); + } +}; + +} // namespace flux2_qkv_epilogue + +} // namespace sglang diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index ba44c8aae..dd20a1522 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -201,6 +201,20 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Fused in-place QK RMS-norm + RoPE.", ), + ( + "diffusion.flux2_qkv_epilogue", + KernelBackend.JIT, + "rope.flux2_qkv_epilogue_jit:try_fused_flux2_qkv_epilogue", + _CUDA, + "FLUX.2 QK RMS-norm + RoPE + joint QKV packing.", + ), + ( + "diffusion.flux2_token_cat_fp8", + KernelBackend.TRITON, + "layout.flux2_token_cat_fp8_triton:try_flux2_token_cat_fp8", + _CUDA, + "FLUX.2 single-block token concatenation + static FP8 quantization.", + ), ( "diffusion.qwen_qkv_epilogue", KernelBackend.JIT, @@ -398,6 +412,7 @@ _EXPORTS: dict[str, str] = { "can_use_fused_layernorm_modulate": "norm.layernorm_modulate_triton", "can_use_fused_qk_head_layernorm": "norm.layernorm_modulate_triton", "fused_layernorm_modulate": "norm.layernorm_modulate_triton", + "fused_layernorm_modulate_fp8_quant_raw": "norm.layernorm_modulate_triton", "fused_layernorm_modulate_raw": "norm.layernorm_modulate_triton", "fused_qk_head_layernorm": "norm.layernorm_modulate_triton", "is_plain_layer_norm": "norm.layernorm_modulate_triton", @@ -443,6 +458,7 @@ _EXPORTS: dict[str, str] = { "can_use_fused_temb_table_slices": "modulate.wan_temb_table_slices_triton", "fused_temb_table_slices": "modulate.wan_temb_table_slices_triton", # Rotary embeddings and the QK-norm chains fused around them + "try_fused_flux2_qkv_epilogue": "rope.flux2_qkv_epilogue_jit", "hunyuan_qkv_rope_pack": "rope.hunyuan_qkv_pack_triton", "can_use_ltx2_qknorm_split_rope_cuda": "rope.ltx2_qknorm_split_rope_jit", "ltx2_qknorm_split_rope_cuda": "rope.ltx2_qknorm_split_rope_jit", @@ -460,6 +476,8 @@ _EXPORTS: dict[str, str] = { "can_use_helios_qk_rope": "rope.helios_qk_rope_jit", "fused_inplace_helios_qk_rope": "rope.helios_qk_rope_jit", "apply_rotary_embedding": "rope.rotary_triton", + # Tensor layout transformations fused with downstream quantization + "try_flux2_token_cat_fp8": "layout.flux2_token_cat_fp8_triton", # Activation-function fusions "can_use_fused_bias_glu": "activation.sana_conv_post_triton", "can_use_fused_bias_silu": "activation.sana_conv_post_triton", diff --git a/python/sglang/kernels/ops/diffusion/layout/flux2_token_cat_fp8_triton.py b/python/sglang/kernels/ops/diffusion/layout/flux2_token_cat_fp8_triton.py new file mode 100644 index 000000000..e854f7621 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/layout/flux2_token_cat_fp8_triton.py @@ -0,0 +1,143 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + +from sglang.kernels.jit.utils import is_arch_support_pdl +from sglang.kernels.ops.quantization.fp8_kernel import ( + fp8_dtype, + fp8_max, + fp8_min, +) +from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton + +_BLOCK = 4096 +_NUM_WARPS = 8 + + +@triton.jit +def _token_cat_fp8_kernel( + attention, + mlp, + output, + input_scale, + attention_hidden: tl.constexpr, + mlp_hidden: tl.constexpr, + output_hidden: tl.constexpr, + BLOCK: tl.constexpr, + FP8_DTYPE: tl.constexpr, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, + USE_PDL: tl.constexpr, +): + row = tl.program_id(0) + block = tl.program_id(1) + columns = block * BLOCK + tl.arange(0, BLOCK) + output_mask = columns < output_hidden + attention_mask = output_mask & (columns < attention_hidden) + mlp_columns = columns - attention_hidden + mlp_mask = output_mask & (columns >= attention_hidden) + + if USE_PDL: + tl.extra.cuda.gdc_wait() + + attention_values = tl.load( + attention + row * attention_hidden + columns, + mask=attention_mask, + other=0.0, + ).to(tl.float32) + mlp_values = tl.load( + mlp + row * mlp_hidden + mlp_columns, + mask=mlp_mask, + other=0.0, + ).to(tl.float32) + values = tl.where(columns < attention_hidden, attention_values, mlp_values) + scale = tl.load(input_scale).to(tl.float32) + quantized = tl.clamp(values * (1.0 / scale), FP8_MIN, FP8_MAX).to(FP8_DTYPE) + + if USE_PDL: + tl.extra.cuda.gdc_launch_dependents() + + tl.store( + output + row * output_hidden + columns, + quantized.to(tl.uint8, bitcast=True), + mask=output_mask, + ) + + +def try_flux2_token_cat_fp8( + attention: torch.Tensor, + mlp: torch.Tensor, + input_scale: torch.Tensor, +) -> torch.Tensor | None: + """Concatenate FLUX.2 single-block branches directly into static FP8.""" + if torch.compiler.is_compiling(): + return None + if not ( + isinstance(attention, torch.Tensor) + and isinstance(mlp, torch.Tensor) + and attention.is_cuda + and mlp.is_cuda + and attention.device == mlp.device + and attention.dtype == torch.bfloat16 + and mlp.dtype == torch.bfloat16 + and attention.ndim == 3 + and mlp.ndim == 3 + and attention.shape[:-1] == mlp.shape[:-1] + and attention.is_contiguous() + and mlp.is_contiguous() + and attention.numel() > 0 + and mlp.numel() > 0 + ): + return None + if ( + torch.cuda.is_current_stream_capturing() + or torch.cuda.get_device_capability(attention.device)[0] < 10 + ): + return None + if not ( + isinstance(input_scale, torch.Tensor) + and input_scale.is_cuda + and input_scale.device == attention.device + and input_scale.dtype == torch.float32 + and input_scale.numel() == 1 + and input_scale.is_contiguous() + ): + return None + + attention_hidden = attention.shape[-1] + mlp_hidden = mlp.shape[-1] + output_hidden = attention_hidden + mlp_hidden + rows = attention.numel() // attention_hidden + output = torch.empty( + (*attention.shape[:-1], output_hidden), + dtype=fp8_dtype, + device=attention.device, + ) + pdl_kwargs = ( + {"USE_PDL": True, "launch_pdl": True} + if is_arch_support_pdl() + else {"USE_PDL": False} + ) + with torch.cuda.device(attention.device): + _token_cat_fp8_kernel[(rows, triton.cdiv(output_hidden, _BLOCK))]( + attention, + mlp, + output.view(torch.uint8), + input_scale, + attention_hidden=attention_hidden, + mlp_hidden=mlp_hidden, + output_hidden=output_hidden, + BLOCK=_BLOCK, + FP8_DTYPE=fp8_dtype_to_triton(fp8_dtype), + FP8_MIN=fp8_min, + FP8_MAX=fp8_max, + num_warps=_NUM_WARPS, + num_stages=1, + **pdl_kwargs, + ) + return output + + +__all__ = ["try_flux2_token_cat_fp8"] diff --git a/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py b/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py index 000e993a8..09028dcbf 100644 --- a/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py +++ b/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py @@ -53,6 +53,8 @@ from sglang.kernels.ops.diffusion.common.numerics import ( round_bf16_to_fp32, ) from sglang.kernels.ops.diffusion.common.platform import is_cuda +from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype, fp8_max, fp8_min +from sglang.kernels.ops.quantization.fp8_utils import fp8_dtype_to_triton from sglang.srt.utils.custom_op import register_custom_op @@ -171,15 +173,22 @@ def _push_vec4( @triton.jit def _layernorm_modulate_kernel( y_ptr, + y_q_ptr, x_ptr, scale_ptr, shift_ptr, + input_scale_ptr, seq_len, n_rows, scale_row_stride, eps, D: tl.constexpr, ROWS: tl.constexpr, + FP8_DTYPE: tl.constexpr, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, + STORE_BF16: tl.constexpr, + QUANTIZE_FP8: tl.constexpr, ): pid = tl.program_id(0).to(tl.int64) row_offs = pid * ROWS + tl.arange(0, ROWS) @@ -258,7 +267,21 @@ def _layernorm_modulate_kernel( ).to(tl.float32) one_plus = round_bf16_to_fp32(1.0 + sc) y = round_bf16_to_fp32(y * one_plus) + sh - tl.store(y_ptr + row_base[:, None] + cols[None, :], y, mask=mask) + if STORE_BF16: + tl.store(y_ptr + row_base[:, None] + cols[None, :], y, mask=mask) + if QUANTIZE_FP8: + # The standalone static quantizer reads the just-written BF16 + # modulation output, so reproduce that final store/load rounding + # before applying its exact per-tensor scale expression. + y = round_bf16_to_fp32(y) + input_scale = tl.load(input_scale_ptr).to(tl.float32) + input_scale_inv = 1.0 / input_scale + y_q = tl.clamp(y * input_scale_inv, FP8_MIN, FP8_MAX).to(FP8_DTYPE) + tl.store( + y_q_ptr + row_base[:, None] + cols[None, :], + y_q.to(tl.uint8, bitcast=True), + mask=mask, + ) @triton.jit @@ -408,16 +431,23 @@ def fused_layernorm_modulate_raw( stride = _mod_row_stride(scale, batch, hidden) with torch.cuda.device(x.device): _layernorm_modulate_kernel[(triton.cdiv(n_rows, rows),)]( + out, out, x, scale, shift, + scale, seq_len, n_rows, stride, eps, D=hidden, ROWS=rows, + FP8_DTYPE=fp8_dtype_to_triton(fp8_dtype), + FP8_MIN=fp8_min, + FP8_MAX=fp8_max, + STORE_BF16=True, + QUANTIZE_FP8=False, # H200-tuned: 38.5us at (1, 4096, 4096) vs the 121.8us eager # chain, 14.3us at Sana's (2, 1024, 2240) vs 43.1us. ROWS=1 + # 4 warps triggers pathological Triton layout conversions in @@ -427,6 +457,49 @@ def fused_layernorm_modulate_raw( return out +def fused_layernorm_modulate_fp8_quant_raw( + x: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + input_scale: torch.Tensor, + eps: float, +) -> torch.Tensor: + """Fuse FLUX.2 LayerNorm, adaLN modulation, and static FP8 quantization. + + The returned tensor is byte-identical to running + :func:`fused_layernorm_modulate_raw` followed by ``static_quant_fp8``. + Unlike that two-op chain, this path does not materialize the intermediate + BF16 activation because FLUX.2 feeds it directly into an FP8 projection. + """ + batch, seq_len, hidden = x.shape + n_rows = batch * seq_len + rows = 2 + out = torch.empty_like(x, dtype=fp8_dtype) + stride = _mod_row_stride(scale, batch, hidden) + with torch.cuda.device(x.device): + _layernorm_modulate_kernel[(triton.cdiv(n_rows, rows),)]( + out, + out.view(torch.uint8), + x, + scale, + shift, + input_scale, + seq_len, + n_rows, + stride, + eps, + D=hidden, + ROWS=rows, + FP8_DTYPE=fp8_dtype_to_triton(fp8_dtype), + FP8_MIN=fp8_min, + FP8_MAX=fp8_max, + STORE_BF16=False, + QUANTIZE_FP8=True, + num_warps=4 if hidden >= 2048 else 2, + ) + return out + + fused_layernorm_modulate = register_custom_op( fused_layernorm_modulate_raw, op_name="triton_fused_layernorm_modulate", diff --git a/python/sglang/kernels/ops/diffusion/rope/flux2_qkv_epilogue_jit.py b/python/sglang/kernels/ops/diffusion/rope/flux2_qkv_epilogue_jit.py new file mode 100644 index 000000000..2018d6df5 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/rope/flux2_qkv_epilogue_jit.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import cache_once, load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +_HEAD_DIM = 128 +_ALIGN = 32 + + +@cache_once +def flux2_qkv_epilogue_module() -> Module: + return load_jit( + "flux2_qkv_epilogue_bf16", + cuda_files=["diffusion/flux2_qkv_epilogue.cuh"], + cuda_wrappers=[ + ( + "flux2_qkv_epilogue", + "flux2_qkv_epilogue::Flux2QKVEpilogueKernel::run", + ) + ], + ) + + +def _qkv_tensor(tensor: torch.Tensor, like: torch.Tensor | None = None) -> bool: + return ( + isinstance(tensor, torch.Tensor) + and tensor.is_cuda + and tensor.dtype == torch.bfloat16 + and tensor.ndim == 4 + and tensor.shape[0] == 1 + and tensor.shape[-1] == _HEAD_DIM + and tensor.numel() > 0 + and tensor.stride(-1) == 1 + and tensor.stride(-2) == _HEAD_DIM + and tensor.data_ptr() % _ALIGN == 0 + and ( + like is None + or ( + tensor.device == like.device + and tensor.shape == like.shape + and tensor.stride(1) == like.stride(1) + ) + ) + ) + + +def try_fused_flux2_qkv_epilogue( + img_q: torch.Tensor, + img_k: torch.Tensor, + img_v: torch.Tensor, + txt_q: torch.Tensor, + txt_k: torch.Tensor, + txt_v: torch.Tensor, + img_q_weight: torch.Tensor, + img_k_weight: torch.Tensor, + txt_q_weight: torch.Tensor, + txt_k_weight: torch.Tensor, + cos_sin_cache: torch.Tensor, + img_eps: float, + txt_eps: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None: + """Fuse FLUX.2 Q/K norm, RoPE, QKV packing, and text/image concat.""" + if torch.compiler.is_compiling(): + return None + if not ( + _qkv_tensor(img_q) + and _qkv_tensor(img_k, img_q) + and _qkv_tensor(img_v, img_q) + and _qkv_tensor(txt_q) + and _qkv_tensor(txt_k, txt_q) + and _qkv_tensor(txt_v, txt_q) + and img_q.shape[2] == txt_q.shape[2] + and torch.version.cuda is not None + and torch.cuda.get_device_capability(img_q.device)[0] >= 10 + ): + return None + if torch.cuda.is_current_stream_capturing(): + return None + + weights = [] + for tensor in (img_q_weight, img_k_weight, txt_q_weight, txt_k_weight): + if not ( + isinstance(tensor, torch.Tensor) + and tensor.is_cuda + and tensor.device == img_q.device + and tensor.dtype == torch.bfloat16 + and tensor.shape == (_HEAD_DIM,) + and tensor.is_contiguous() + and tensor.data_ptr() % _ALIGN == 0 + ): + return None + weights.append(tensor) + + total_tokens = txt_q.shape[1] + img_q.shape[1] + if not ( + isinstance(cos_sin_cache, torch.Tensor) + and cos_sin_cache.is_cuda + and cos_sin_cache.device == img_q.device + and cos_sin_cache.dtype == torch.float32 + and cos_sin_cache.ndim == 2 + and cos_sin_cache.shape[0] >= total_tokens + and cos_sin_cache.shape[1] == _HEAD_DIM + and cos_sin_cache.is_contiguous() + ): + return None + + heads = img_q.shape[2] + joint_shape = (1, total_tokens, heads, _HEAD_DIM) + joint_q = torch.empty(joint_shape, dtype=img_q.dtype, device=img_q.device) + joint_k = torch.empty_like(joint_q) + joint_v = torch.empty_like(joint_q) + flux2_qkv_epilogue_module().flux2_qkv_epilogue( + joint_q.view(-1, heads, _HEAD_DIM), + joint_k.view(-1, heads, _HEAD_DIM), + joint_v.view(-1, heads, _HEAD_DIM), + img_q.view(-1, heads, _HEAD_DIM), + img_k.view(-1, heads, _HEAD_DIM), + img_v.view(-1, heads, _HEAD_DIM), + txt_q.view(-1, heads, _HEAD_DIM), + txt_k.view(-1, heads, _HEAD_DIM), + txt_v.view(-1, heads, _HEAD_DIM), + *weights, + cos_sin_cache, + float(img_eps), + float(txt_eps), + ) + return joint_q, joint_k, joint_v diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 184fa8670..0a3c1b46d 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -27,12 +27,16 @@ from sglang.kernels.ops.diffusion import ( can_use_flux2_gated_resnorm, can_use_fused_layernorm_modulate, flux2_gated_resnorm_raw, + fused_layernorm_modulate_fp8_quant_raw, fused_layernorm_modulate_raw, fused_packed_silu_mul_bitexact, is_plain_layer_norm, residual_gate_add, + try_flux2_token_cat_fp8, try_flux2_token_cat_nvfp4, + try_fused_flux2_qkv_epilogue, ) +from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( divide, @@ -63,6 +67,8 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( ModelOptFp4Config, ModelOptFp4LinearMethod, + ModelOptFp8Config, + ModelOptFp8LinearMethod, apply_nvfp4_gemm_prequantized, ) from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( @@ -89,6 +95,93 @@ assert _FLUX2_LN_MOD_SIGS is not None _FLUX2_SWIGLU = BitExactFusionGate("FLUX.2 fused SwiGLU", per_signature=True) _FLUX2_SWIGLU_SIGS = _FLUX2_SWIGLU.verified_sigs assert _FLUX2_SWIGLU_SIGS is not None +_FLUX2_LN_FP8 = BitExactFusionGate( + "FLUX.2 fused LN+modulate+FP8 quant", per_signature=True +) +_FLUX2_LN_FP8_SIGS = _FLUX2_LN_FP8.verified_sigs +assert _FLUX2_LN_FP8_SIGS is not None + + +def _valid_modelopt_fp8_linear(linear: nn.Module) -> bool: + input_scale = getattr(linear, "input_scale", None) + return ( + isinstance(getattr(linear, "quant_method", None), ModelOptFp8LinearMethod) + and isinstance(input_scale, torch.Tensor) + and input_scale.is_cuda + and input_scale.dtype == torch.float32 + and input_scale.numel() == 1 + and input_scale.is_contiguous() + and bool(torch.isfinite(input_scale).all().item()) + and bool((input_scale > 0).all().item()) + ) + + +def _shared_modelopt_fp8_scale(linears: list[nn.Module]) -> bool: + if not all(_valid_modelopt_fp8_linear(linear) for linear in linears): + return False + reference = linears[0].input_scale + return all(torch.equal(reference, linear.input_scale) for linear in linears[1:]) + + +def _try_flux2_norm_modulate_fp8( + norm: nn.Module, + x: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, + input_scale: Optional[torch.Tensor], + *, + enabled: bool, +) -> Optional[torch.Tensor]: + """Return the exact prequantized FP8 projection input when eligible.""" + if not enabled or torch.compiler.is_compiling(): + return None + + scale_row = scale.squeeze(1) if scale.dim() == 3 and scale.shape[1] == 1 else scale + shift_row = shift.squeeze(1) if shift.dim() == 3 and shift.shape[1] == 1 else shift + if ( + _FLUX2_LN_FP8.disabled + or x.shape[-1] != 6144 + or not is_plain_layer_norm(norm, x.shape[-1]) + or not can_use_fused_layernorm_modulate(x, scale_row, shift_row) + ): + return None + + sig = ( + x.dtype, + x.device, + x.shape[0], + x.shape[-1], + x.stride(-1), + scale_row.stride(0) if scale_row.shape[0] > 1 else x.shape[-1], + shift_row.stride(0) if shift_row.shape[0] > 1 else x.shape[-1], + norm.eps, + ) + verified = sig in _FLUX2_LN_FP8_SIGS + if not verified and torch.cuda.is_current_stream_capturing(): + return None + try: + out = fused_layernorm_modulate_fp8_quant_raw( + x, scale_row, shift_row, input_scale, norm.eps + ) + except Exception as exc: + _FLUX2_LN_FP8.on_exception(exc, logger=logger) + return None + if verified: + return out + + reference_bf16 = _flux2_norm_modulate(norm, x, scale, shift) + reference, _ = static_quant_fp8(reference_bf16, input_scale) + return _FLUX2_LN_FP8.accept_or_fallback( + out, + reference, + sig=sig, + logger=logger, + mismatch_msg=( + "FLUX.2 fused LN+modulate+FP8 quant fast path is not bit-exact " + "on this platform; falling back to the split path" + ), + ) + PendingGatedResidual = Tuple[torch.Tensor, torch.Tensor, torch.Tensor] @@ -123,6 +216,45 @@ def _flux2_gated_resnorm( return _flux2_norm_modulate(norm, residual, scale, shift), residual +def _flux2_norm_maybe_fp8( + norm: nn.Module, + hidden_states: torch.Tensor | PendingGatedResidual, + scale: torch.Tensor, + shift: torch.Tensor, + input_scale: Optional[torch.Tensor], + *, + fp8_enabled: bool, + update: Optional[torch.Tensor] = None, + gate: Optional[torch.Tensor] = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Return ``(norm_hidden_states, residual_hidden_states)``. + + Uses fused gated residual-norm when a pending residual is present or when + ``update``/``gate`` are supplied. When the FP8 producer is enabled, the + residual is materialized first so LN+modulate+FP8 can write the GEMM input. + """ + if isinstance(hidden_states, tuple): + residual, pending_update, pending_gate = hidden_states + if fp8_enabled: + hidden_states = residual_gate_add(residual, pending_update, pending_gate) + else: + return _flux2_gated_resnorm( + norm, residual, pending_update, pending_gate, scale, shift + ) + elif update is not None and gate is not None: + if fp8_enabled: + hidden_states = residual_gate_add(hidden_states, update, gate) + else: + return _flux2_gated_resnorm(norm, hidden_states, update, gate, scale, shift) + + norm_hidden_states = _try_flux2_norm_modulate_fp8( + norm, hidden_states, scale, shift, input_scale, enabled=fp8_enabled + ) + if norm_hidden_states is None: + norm_hidden_states = _flux2_norm_modulate(norm, hidden_states, scale, shift) + return norm_hidden_states, hidden_states + + def _flux2_norm_modulate( norm: nn.Module, x: torch.Tensor, @@ -315,13 +447,22 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): self.added_kv_proj_dim = added_kv_proj_dim self.added_proj_bias = added_proj_bias - # Some FLUX.2 NVFP4 checkpoints store Q/K/V packed as a single tensor, while - # ModelOpt's standard diffusers export keeps the original to_q/to_k/to_v layout. - # Only enable the fused loader path for the packed checkpoint family. - self.use_fused_qkv = isinstance(quant_config, ModelOptFp4Config) and getattr( + # Packed NVFP4 checkpoints already serialize QKV together. ModelOpt + # FP8 exports separate Diffusers tensors, but the loader can merge + # those tensors losslessly and execute one channelwise-CUTLASS GEMM. + fp4_packed_qkv = isinstance(quant_config, ModelOptFp4Config) and getattr( quant_config, "checkpoint_uses_packed_qkv", False ) + capability = current_platform.get_device_capability() + fp8_merged_qkv = ( + isinstance(quant_config, ModelOptFp8Config) + and self.tp_size == 1 + and capability is not None + and capability.major >= 10 + ) + self.use_fused_qkv = fp4_packed_qkv or fp8_merged_qkv self.use_fused_added_qkv = self.use_fused_qkv + self.use_fused_qkv_epilogue = fp8_merged_qkv if self.use_fused_qkv: self.to_qkv = MergedColumnParallelLinear( @@ -379,13 +520,14 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): self.norm_added_q = RMSNorm(dim_head, eps=eps) self.norm_added_k = RMSNorm(dim_head, eps=eps) if self.use_fused_added_qkv: - # txt_attn.qkv is always BF16 in the NVFP4 checkpoint — no quant needed + # txt_attn.qkv is BF16 in the packed NVFP4 checkpoint, while + # ModelOpt FP8 keeps it quantized like the image projection. self.to_added_qkv = MergedColumnParallelLinear( added_kv_proj_dim, [self.inner_dim] * 3, bias=added_proj_bias, gather_output=False, - quant_config=None, + quant_config=None if fp4_packed_qkv else quant_config, prefix=f"{prefix}.to_added_qkv" if prefix else "to_added_qkv", ) else: @@ -447,7 +589,12 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): encoder_query, encoder_key, encoder_value, - ) = _get_qkv_projections(self, hidden_states, encoder_hidden_states) + ) = _get_qkv_projections( + self, + hidden_states, + encoder_hidden_states, + make_contiguous=not self.use_fused_qkv_epilogue, + ) query = query.unflatten(-1, (self.local_heads, -1)) key = key.unflatten(-1, (self.local_heads, -1)) @@ -464,41 +611,78 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): dim=-1, ) + joint_qkv = None + sp_txt_pad = 0 if self.added_kv_proj_dim is not None: encoder_query = encoder_query.unflatten(-1, (self.local_heads, -1)) encoder_key = encoder_key.unflatten(-1, (self.local_heads, -1)) encoder_value = encoder_value.unflatten(-1, (self.local_heads, -1)) text_seq_len = encoder_query.shape[1] - encoder_query, encoder_key = apply_qk_norm_with_optional_rope( - q=encoder_query, - k=encoder_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=self.head_dim, - cos_sin_cache=cos_sin_cache, - is_neox=False, - allow_inplace=True, - ) - query, key = apply_qk_norm_with_optional_rope( - q=query, - k=key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - cos_sin_cache=cos_sin_cache, - is_neox=False, - position_offset=text_seq_len, - allow_inplace=True, - ) - # join_seqs relocates any SP text tail-pad behind the image (see # sp_shard.join_seqs for why). sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0) - query = join_seqs(encoder_query, query, sp_txt_pad) - key = join_seqs(encoder_key, key, sp_txt_pad) - value = join_seqs(encoder_value, value, sp_txt_pad) + if ( + self.use_fused_qkv_epilogue + and cos_sin_cache is not None + and sp_txt_pad == 0 + ): + joint_qkv = try_fused_flux2_qkv_epilogue( + query, + key, + value, + encoder_query, + encoder_key, + encoder_value, + self.norm_q.weight, + self.norm_k.weight, + self.norm_added_q.weight, + self.norm_added_k.weight, + cos_sin_cache, + self.norm_q.variance_epsilon, + self.norm_added_q.variance_epsilon, + ) + + if joint_qkv is None: + if self.use_fused_qkv_epilogue: + query, key, value = [ + tensor.contiguous() for tensor in (query, key, value) + ] + encoder_query, encoder_key, encoder_value = [ + tensor.contiguous() + for tensor in (encoder_query, encoder_key, encoder_value) + ] + encoder_query, encoder_key = apply_qk_norm_with_optional_rope( + q=encoder_query, + k=encoder_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, + ) + query, key = apply_qk_norm_with_optional_rope( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + position_offset=text_seq_len, + allow_inplace=True, + ) + query = join_seqs(encoder_query, query, sp_txt_pad) + key = join_seqs(encoder_key, key, sp_txt_pad) + value = join_seqs(encoder_value, value, sp_txt_pad) + else: + query, key, value = joint_qkv else: + if self.use_fused_qkv_epilogue: + query, key, value = [ + tensor.contiguous() for tensor in (query, key, value) + ] query, key = apply_qk_norm_with_optional_rope( q=query, k=key, @@ -613,6 +797,9 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin): quant_config=quant_config, prefix=f"{prefix}.to_out" if prefix else "to_out", ) + self._enable_fp8_token_cat = self.tp_size == 1 and isinstance( + self.to_out.quant_method, ModelOptFp8LinearMethod + ) self._enable_nvfp4_token_cat = False capability = current_platform.get_device_capability() if ( @@ -723,25 +910,33 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin): # Handle the feedforward (FF) logic mlp_hidden_states = self.mlp_act_fn(mlp_hidden_states) - # Concatenate and parallel output projection. On SM103 NVFP4 the - # producer writes the concatenated packed values and swizzled scales - # directly, avoiding a full-width BF16 cat materialization. + # Concatenate and parallel output projection. FP8 writes a packed + # GEMM input; SM103 NVFP4 writes packed values and swizzled scales. + # Both avoid a full-width BF16 cat materialization. output_shape = (*hidden_states.shape[:-1], self.out_dim) + quantized = None packed = None - if self._enable_nvfp4_token_cat: + if self._enable_fp8_token_cat: + quantized = try_flux2_token_cat_fp8( + hidden_states, mlp_hidden_states, self.to_out.input_scale + ) + elif self._enable_nvfp4_token_cat: packed = try_flux2_token_cat_nvfp4( hidden_states, mlp_hidden_states, self.to_out.input_scale_inv ) - if packed is None: - hidden_states = torch.cat([hidden_states, mlp_hidden_states], dim=-1) + if quantized is not None: + hidden_states = quantized hidden_states, _ = self.to_out(hidden_states) - else: + elif packed is not None: hidden_states = apply_nvfp4_gemm_prequantized( self.to_out, *packed, output_dtype=hidden_states.dtype, bias=self.to_out.bias, ).view(*output_shape) + else: + hidden_states = torch.cat([hidden_states, mlp_hidden_states], dim=-1) + hidden_states, _ = self.to_out(hidden_states) return hidden_states @@ -780,6 +975,10 @@ class Flux2SingleTransformerBlock(nn.Module): quant_config=quant_config, prefix=f"{prefix}.attn" if prefix else "attn", ) + self._fp8_norm_quant = False + + def configure_fp8_norm_quant(self) -> None: + self._fp8_norm_quant = _valid_modelopt_fp8_linear(self.attn.to_qkv_mlp_proj) def forward( self, @@ -801,15 +1000,14 @@ class Flux2SingleTransformerBlock(nn.Module): mod_shift, mod_scale, mod_gate = temb_mod_params - if isinstance(hidden_states, tuple): - residual, update, gate = hidden_states - norm_hidden_states, hidden_states = _flux2_gated_resnorm( - self.norm, residual, update, gate, mod_scale, mod_shift - ) - else: - norm_hidden_states = _flux2_norm_modulate( - self.norm, hidden_states, mod_scale, mod_shift - ) + norm_hidden_states, hidden_states = _flux2_norm_maybe_fp8( + self.norm, + hidden_states, + mod_scale, + mod_shift, + (self.attn.to_qkv_mlp_proj.input_scale if self._fp8_norm_quant else None), + fp8_enabled=self._fp8_norm_quant, + ) joint_attention_kwargs = joint_attention_kwargs or {} attn_output = self.attn( @@ -891,6 +1089,36 @@ class Flux2TransformerBlock(nn.Module): quant_config=quant_config, prefix=f"{prefix}.ff_context" if prefix else "ff_context", ) + self._fp8_img_attn_norm_quant = False + self._fp8_txt_attn_norm_quant = False + self._fp8_img_ff_norm_quant = False + self._fp8_txt_ff_norm_quant = False + + def configure_fp8_norm_quant(self) -> None: + if self.attn.use_fused_qkv: + self._fp8_img_attn_norm_quant = _valid_modelopt_fp8_linear(self.attn.to_qkv) + else: + self._fp8_img_attn_norm_quant = _shared_modelopt_fp8_scale( + [self.attn.to_q, self.attn.to_k, self.attn.to_v] + ) + + if self.attn.use_fused_added_qkv: + self._fp8_txt_attn_norm_quant = _valid_modelopt_fp8_linear( + self.attn.to_added_qkv + ) + else: + self._fp8_txt_attn_norm_quant = _shared_modelopt_fp8_scale( + [ + self.attn.add_q_proj, + self.attn.add_k_proj, + self.attn.add_v_proj, + ] + ) + + self._fp8_img_ff_norm_quant = _valid_modelopt_fp8_linear(self.ff.linear_in) + self._fp8_txt_ff_norm_quant = _valid_modelopt_fp8_linear( + self.ff_context.linear_in + ) def forward( self, @@ -929,34 +1157,40 @@ class Flux2TransformerBlock(nn.Module): ) = temb_mod_params_txt # Img stream - if isinstance(hidden_states, tuple): - residual, update, gate = hidden_states - norm_hidden_states, hidden_states = _flux2_gated_resnorm( - self.norm1, residual, update, gate, scale_msa, shift_msa - ) - else: - norm_hidden_states = _flux2_norm_modulate( - self.norm1, hidden_states, scale_msa, shift_msa - ) + norm_hidden_states, hidden_states = _flux2_norm_maybe_fp8( + self.norm1, + hidden_states, + scale_msa, + shift_msa, + ( + ( + self.attn.to_qkv.input_scale + if self.attn.use_fused_qkv + else self.attn.to_q.input_scale + ) + if self._fp8_img_attn_norm_quant + else None + ), + fp8_enabled=self._fp8_img_attn_norm_quant, + ) # Conditioning txt stream - if isinstance(encoder_hidden_states, tuple): - residual, update, gate = encoder_hidden_states - norm_encoder_hidden_states, encoder_hidden_states = _flux2_gated_resnorm( - self.norm1_context, - residual, - update, - gate, - c_scale_msa, - c_shift_msa, - ) - else: - norm_encoder_hidden_states = _flux2_norm_modulate( - self.norm1_context, - encoder_hidden_states, - c_scale_msa, - c_shift_msa, - ) + norm_encoder_hidden_states, encoder_hidden_states = _flux2_norm_maybe_fp8( + self.norm1_context, + encoder_hidden_states, + c_scale_msa, + c_shift_msa, + ( + ( + self.attn.to_added_qkv.input_scale + if self.attn.use_fused_added_qkv + else self.attn.add_q_proj.input_scale + ) + if self._fp8_txt_attn_norm_quant + else None + ), + fp8_enabled=self._fp8_txt_attn_norm_quant, + ) # Attention on concatenated img + txt stream attention_outputs = self.attn( @@ -970,26 +1204,34 @@ class Flux2TransformerBlock(nn.Module): attn_output, context_attn_output = attention_outputs # Process attention outputs for the image stream (`hidden_states`). - norm_hidden_states, hidden_states = _flux2_gated_resnorm( + norm_hidden_states, hidden_states = _flux2_norm_maybe_fp8( self.norm2, hidden_states, - attn_output, - gate_msa, scale_mlp, shift_mlp, + (self.ff.linear_in.input_scale if self._fp8_img_ff_norm_quant else None), + fp8_enabled=self._fp8_img_ff_norm_quant, + update=attn_output, + gate=gate_msa, ) ff_output = self.ff(norm_hidden_states) hidden_states = _defer_gated_residual(hidden_states, ff_output, gate_mlp) # Process attention outputs for the text stream (`encoder_hidden_states`). - norm_encoder_hidden_states, encoder_hidden_states = _flux2_gated_resnorm( + norm_encoder_hidden_states, encoder_hidden_states = _flux2_norm_maybe_fp8( self.norm2_context, encoder_hidden_states, - context_attn_output, - c_gate_msa, c_scale_mlp, c_shift_mlp, + ( + self.ff_context.linear_in.input_scale + if self._fp8_txt_ff_norm_quant + else None + ), + fp8_enabled=self._fp8_txt_ff_norm_quant, + update=context_attn_output, + gate=c_gate_msa, ) context_ff_output = self.ff_context(norm_encoder_hidden_states) @@ -1110,7 +1352,45 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): """ - param_names_mapping = FluxConfig().arch_config.param_names_mapping + param_names_mapping = { + # ModelOpt FP8 exports separate Diffusers projections. Merge Q/K/V and + # their static scales into the runtime MergedColumnParallelLinear. + r"^(transformer_blocks\.\d+\.attn)\.to_q\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 0, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.to_k\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 1, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.to_v\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 2, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_q_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 0, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_k_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 1, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_v_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 2, + 3, + ), + **FluxConfig().arch_config.param_names_mapping, + } + packed_modules_mapping = { + "to_qkv": ["to_q", "to_k", "to_v"], + "to_added_qkv": ["add_q_proj", "add_k_proj", "add_v_proj"], + } scale_shift_swap_params = ("norm_out.linear.weight", "norm_out.linear.bias") # FLUX.2 stays closer to the official diffusers output with Torch SDPA. # The generic FA path still produces a measurable image-level drift here. @@ -1122,26 +1402,45 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): } def post_load_weights(self) -> None: - if not isinstance(getattr(self, "quant_config", None), ModelOptFp4Config): - return + super().post_load_weights() + if isinstance(getattr(self, "quant_config", None), ModelOptFp4Config): + # BFL/ComfyUI checkpoints store AdaLN modulation params as + # [scale, shift], while diffusers expects [shift, scale]. + for param_name in self.scale_shift_swap_params: + parts = param_name.split(".") + module = self + for part in parts[:-1]: + module = getattr(module, part) + param = getattr(module, parts[-1], None) + if param is None: + continue + half = param.shape[0] // 2 + with torch.no_grad(): + first_half = param[:half].clone() + param[:half] = param[half:] + param[half:] = first_half + logger.info( + "Swapped scale/shift order for %s (BFL → diffusers)", + param_name, + ) - # BFL/ComfyUI checkpoints store AdaLN modulation params as [scale, shift], - # while diffusers expects [shift, scale]. - for param_name in self.scale_shift_swap_params: - parts = param_name.split(".") - module = self - for part in parts[:-1]: - module = getattr(module, part) - param = getattr(module, parts[-1], None) - if param is None: - continue - half = param.shape[0] // 2 - with torch.no_grad(): - first_half = param[:half].clone() - param[:half] = param[half:] - param[half:] = first_half + for block in self.transformer_blocks: + block.configure_fp8_norm_quant() + for block in self.single_transformer_blocks: + block.configure_fp8_norm_quant() + enabled = sum( + block._fp8_img_attn_norm_quant + + block._fp8_txt_attn_norm_quant + + block._fp8_img_ff_norm_quant + + block._fp8_txt_ff_norm_quant + for block in self.transformer_blocks + ) + sum(block._fp8_norm_quant for block in self.single_transformer_blocks) + total = 4 * len(self.transformer_blocks) + len(self.single_transformer_blocks) + if enabled: logger.info( - "Swapped scale/shift order for %s (BFL → diffusers)", param_name + "Enabled FLUX.2 FP8 norm+quant fusion for %d/%d block paths", + enabled, + total, ) def __init__( diff --git a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py index 72749d513..4dfb4ade1 100644 --- a/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py +++ b/python/sglang/multimodal_gen/test/single_test_file/component_accuracy/engine.py @@ -452,7 +452,9 @@ class AccuracyEngine: ) if mapping: source_state, _ = hf_to_custom_state_dict( - source_state, get_param_names_mapping(mapping) + source_state, + get_param_names_mapping(mapping), + valid_target_names=set(target.state_dict()), ) lookup = build_state_lookup(source_state) diff --git a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py index 81ab4ae07..087b5e406 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_weight_transfer.py @@ -40,6 +40,20 @@ class _TargetProjectionSet(nn.Module): param.data[:, : source.shape[1]].copy_(source) +class _ConditionalQkvProjectionSet(nn.Module): + param_names_mapping = { + r"^to_q\.(weight)$": (r"to_qkv.\1", 0, 3), + r"^to_k\.(weight)$": (r"to_qkv.\1", 1, 3), + r"^to_v\.(weight)$": (r"to_qkv.\1", 2, 3), + } + + def __init__(self) -> None: + super().__init__() + self.to_q = nn.Linear(2, 2, bias=False) + self.to_k = nn.Linear(2, 2, bias=False) + self.to_v = nn.Linear(2, 2, bias=False) + + def test_transfer_weights_uses_loaders_for_fused_aliases_and_padding() -> None: source = _SourceProjectionSet().to(dtype=torch.bfloat16) target = _TargetProjectionSet() @@ -57,3 +71,25 @@ def test_transfer_weights_uses_loaders_for_fused_aliases_and_padding() -> None: assert torch.count_nonzero(target.gate_up_proj.weight[[3, 7]]) == 0 torch.testing.assert_close(target.down_proj.weight[:, :3], source.down_proj.weight) assert torch.count_nonzero(target.down_proj.weight[:, 3]) == 0 + + +def test_transfer_weights_preserves_unfused_targets_for_conditional_mapping() -> None: + source = _ConditionalQkvProjectionSet().to(dtype=torch.bfloat16) + target = _ConditionalQkvProjectionSet() + with torch.no_grad(): + for index, parameter in enumerate(source.parameters(), start=1): + parameter.fill_(index) + for parameter in target.parameters(): + parameter.zero_() + + AccuracyEngine.transfer_weights( + source, + target, + min_match_ratio=1.0, + target_device=torch.device("cpu"), + ) + + for source_parameter, target_parameter in zip( + source.parameters(), target.parameters(), strict=True + ): + torch.testing.assert_close(target_parameter, source_parameter) diff --git a/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py index 45a72c656..62c9f774e 100644 --- a/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py +++ b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py @@ -1,7 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 -"""Serialized ModelOpt FP8 checkpoints must postprocess on device even under -layerwise offload: requantize_with_max_scale() runs scaled_fp8_quant(), a -CUDA-only kernel, so a CPU-resident postprocess must never come back.""" +"""Serialized ModelOpt FP8 checkpoints must postprocess correctly even when +layerwise offload moves the component back to CPU after loading.""" import unittest from unittest.mock import patch diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index 98e63b8f7..3ff46b69f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -110,6 +110,9 @@ from sglang.multimodal_gen.runtime.loader.utils import ( ) from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock +from sglang.multimodal_gen.runtime.models.dits.flux_2 import ( + Flux2Transformer2DModel, +) from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import MiniMaxH3DiTModel from sglang.multimodal_gen.runtime.models.dits.qwen_image import ( QwenImageTransformer2DModel, @@ -320,6 +323,55 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertIsNone(merge_index) self.assertIsNone(total_shards) + def test_flux2_modelopt_fp8_qkv_checkpoint_tensors_are_merged(self): + mapping = get_param_names_mapping(Flux2Transformer2DModel.param_names_mapping) + prefix = "transformer_blocks.0.attn" + source = {} + for projection_prefix in ("to_", "add_"): + source_names = ( + ("q", "k", "v") + if projection_prefix == "to_" + else ("q_proj", "k_proj", "v_proj") + ) + for shard_id, shard_name in enumerate(source_names): + name = f"{prefix}.{projection_prefix}{shard_name}" + source[f"{name}.weight"] = torch.full( + (2, 3), shard_id + 1, dtype=torch.float8_e4m3fn + ) + source[f"{name}.weight_scale"] = torch.tensor( + [0.1 * (shard_id + 1)], dtype=torch.float32 + ) + source[f"{name}.input_scale"] = torch.tensor([0.2], dtype=torch.float32) + + merged, _ = hf_to_custom_state_dict(source, mapping) + + for target in ("to_qkv", "to_added_qkv"): + self.assertEqual(merged[f"{prefix}.{target}.weight"].shape, (6, 3)) + torch.testing.assert_close( + merged[f"{prefix}.{target}.weight_scale"], + torch.tensor([0.1, 0.2, 0.3], dtype=torch.float32), + ) + torch.testing.assert_close( + merged[f"{prefix}.{target}.input_scale"], + torch.tensor([0.2, 0.2, 0.2], dtype=torch.float32), + ) + self.assertEqual( + Flux2Transformer2DModel.packed_modules_mapping["to_qkv"], + ["to_q", "to_k", "to_v"], + ) + + # On an unfused model (BF16, Hopper FP8, or TP>1), the source + # projection names are valid model parameters and the loader must keep + # them separate rather than producing a nonexistent packed target. + unmerged, _ = hf_to_custom_state_dict( + source, + mapping, + valid_target_names=set(source), + ) + self.assertEqual(set(unmerged), set(source)) + for name, tensor in source.items(): + torch.testing.assert_close(unmerged[name], tensor) + @patch( "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", return_value=None, diff --git a/test/registered/kernels/benchmark/diffusion/bench_flux2_fp8_norm_quant.py b/test/registered/kernels/benchmark/diffusion/bench_flux2_fp8_norm_quant.py new file mode 100644 index 000000000..3d229725a --- /dev/null +++ b/test/registered/kernels/benchmark/diffusion/bench_flux2_fp8_norm_quant.py @@ -0,0 +1,48 @@ +import torch + +from sglang.kernels.jit.benchmark import marker +from sglang.kernels.ops.diffusion import ( + fused_layernorm_modulate_fp8_quant_raw, + fused_layernorm_modulate_raw, +) +from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=8, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + +DEVICE = "cuda" +DTYPE = torch.bfloat16 +HIDDEN = 6144 +EPS = 1e-6 + + +@marker.parametrize("rows", [512, 4096, 4608], [512]) +@marker.benchmark("impl", ["split", "fused"], unit="us") +def benchmark(rows: int, impl: str): + generator = torch.Generator(device=DEVICE) + generator.manual_seed(20260831 + rows) + x = torch.randn((1, rows, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + scale = torch.randn((1, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + shift = torch.randn((1, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + input_scale = torch.tensor(0.03125, dtype=torch.float32, device=DEVICE) + + if impl == "split": + + def fn(): + normalized = fused_layernorm_modulate_raw(x, scale, shift, EPS) + return static_quant_fp8(normalized, input_scale)[0] + + else: + + def fn(): + return fused_layernorm_modulate_fp8_quant_raw( + x, scale, shift, input_scale, EPS + ) + + return marker.do_bench(fn, disable_log_bandwidth=True) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/kernels/benchmark/diffusion/bench_flux2_token_cat_fp8.py b/test/registered/kernels/benchmark/diffusion/bench_flux2_token_cat_fp8.py new file mode 100644 index 000000000..6b36d9434 --- /dev/null +++ b/test/registered/kernels/benchmark/diffusion/bench_flux2_token_cat_fp8.py @@ -0,0 +1,46 @@ +import torch + +from sglang.kernels.jit.benchmark import marker +from sglang.kernels.ops.diffusion import try_flux2_token_cat_fp8 +from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=8, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + + +@marker.parametrize("tokens", [512, 4096], [4096]) +@marker.benchmark("impl", ["cat_then_quant", "fused"], unit="us") +def benchmark(tokens: int, impl: str): + generator = torch.Generator(device="cuda") + generator.manual_seed(20260831 + tokens) + attention = torch.randn( + (1, tokens, 6144), + dtype=torch.bfloat16, + device="cuda", + generator=generator, + ) + mlp = torch.randn( + (1, tokens, 18432), + dtype=torch.bfloat16, + device="cuda", + generator=generator, + ) + input_scale = torch.tensor([0.013], dtype=torch.float32, device="cuda") + + if impl == "cat_then_quant": + + def fn(): + return static_quant_fp8(torch.cat([attention, mlp], dim=-1), input_scale)[0] + + else: + + def fn(): + return try_flux2_token_cat_fp8(attention, mlp, input_scale) + + return marker.do_bench(fn, use_cuda_graph=False, disable_log_bandwidth=True) + + +if __name__ == "__main__": + benchmark.run() diff --git a/test/registered/kernels/ops/diffusion/test_flux2_fp8_norm_quant.py b/test/registered/kernels/ops/diffusion/test_flux2_fp8_norm_quant.py new file mode 100644 index 000000000..7b7cf08f6 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_flux2_fp8_norm_quant.py @@ -0,0 +1,52 @@ +import sys + +import pytest +import torch + +from sglang.kernels.ops.diffusion import ( + fused_layernorm_modulate_fp8_quant_raw, + fused_layernorm_modulate_raw, +) +from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=25, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +DEVICE = "cuda" +DTYPE = torch.bfloat16 +HIDDEN = 6144 +EPS = 1e-6 + + +def _make_inputs(rows: int): + generator = torch.Generator(device=DEVICE) + generator.manual_seed(20260831 + rows) + x = torch.randn((1, rows, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + scale = torch.randn((1, 1, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + shift = torch.randn((1, 1, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator) + return x, scale, shift + + +@pytest.mark.parametrize("rows", [1, 127, 512]) +@pytest.mark.parametrize( + "input_scale_value", [0.005, 0.03125, 0.25, 0.4754464328289032, 1.0] +) +def test_flux2_layernorm_modulate_fp8_is_bit_exact( + rows: int, input_scale_value: float +) -> None: + x, scale, shift = _make_inputs(rows) + input_scale = torch.tensor(input_scale_value, dtype=torch.float32, device=DEVICE) + + normalized = fused_layernorm_modulate_raw( + x, scale.squeeze(1), shift.squeeze(1), EPS + ) + expected, _ = static_quant_fp8(normalized, input_scale) + actual = fused_layernorm_modulate_fp8_quant_raw( + x, scale.squeeze(1), shift.squeeze(1), input_scale, EPS + ) + + assert torch.equal(actual.view(torch.uint8), expected.view(torch.uint8)) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_flux2_qkv_epilogue.py b/test/registered/kernels/ops/diffusion/test_flux2_qkv_epilogue.py new file mode 100644 index 000000000..ee4b01fef --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_flux2_qkv_epilogue.py @@ -0,0 +1,158 @@ +import sys +from unittest.mock import patch + +import pytest +import torch + +from sglang.kernels.ops.diffusion import ( + try_fused_flux2_qkv_epilogue, +) +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm_with_optional_rope, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=25, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +DEVICE = "cuda" +DTYPE = torch.bfloat16 +HEAD_DIM = 128 + + +def _packed_qkv(tokens: int, heads: int, generator: torch.Generator): + source = [ + torch.randn( + (1, tokens, heads, HEAD_DIM), + dtype=DTYPE, + device=DEVICE, + generator=generator, + ) + for _ in range(3) + ] + packed = torch.cat([tensor.flatten(2) for tensor in source], dim=-1) + views = [ + tensor.unflatten(-1, (heads, HEAD_DIM)) for tensor in packed.chunk(3, dim=-1) + ] + assert all(not tensor.is_contiguous() for tensor in views) + return views + + +@pytest.mark.parametrize("img_tokens,txt_tokens,heads", [(17, 7, 4), (256, 64, 8)]) +def test_flux2_qkv_epilogue_is_bit_exact( + img_tokens: int, txt_tokens: int, heads: int +) -> None: + generator = torch.Generator(device=DEVICE) + generator.manual_seed(20260831 + img_tokens) + img_qkv = _packed_qkv(img_tokens, heads, generator) + txt_qkv = _packed_qkv(txt_tokens, heads, generator) + norms = [ + RMSNorm(HEAD_DIM, eps=1e-6).to(device=DEVICE, dtype=DTYPE) for _ in range(4) + ] + for norm in norms: + norm.weight.data.normal_(generator=generator) + + angles = torch.randn( + (img_tokens + txt_tokens, HEAD_DIM // 2), + device=DEVICE, + generator=generator, + ) + cache = torch.cat([angles.cos(), angles.sin()], dim=-1).contiguous() + + img_reference = [tensor.contiguous() for tensor in img_qkv] + txt_reference = [tensor.contiguous() for tensor in txt_qkv] + txt_reference[0], txt_reference[1] = apply_qk_norm_with_optional_rope( + txt_reference[0], + txt_reference[1], + norms[2], + norms[3], + HEAD_DIM, + cache, + is_neox=False, + ) + img_reference[0], img_reference[1] = apply_qk_norm_with_optional_rope( + img_reference[0], + img_reference[1], + norms[0], + norms[1], + HEAD_DIM, + cache, + is_neox=False, + position_offset=txt_tokens, + ) + expected = tuple( + torch.cat([txt_reference[index], img_reference[index]], dim=1) + for index in range(3) + ) + + actual = try_fused_flux2_qkv_epilogue( + *img_qkv, + *txt_qkv, + norms[0].weight, + norms[1].weight, + norms[2].weight, + norms[3].weight, + cache, + 1e-6, + 1e-6, + ) + + assert actual is not None + assert all( + torch.equal(result, reference) + for result, reference in zip(actual, expected, strict=True) + ) + + +def test_flux2_qkv_epilogue_rejects_compile() -> None: + tensor = torch.empty((1, 1, 1, HEAD_DIM), device=DEVICE, dtype=DTYPE) + weight = torch.empty((HEAD_DIM,), device=DEVICE, dtype=DTYPE) + cache = torch.empty((2, HEAD_DIM), device=DEVICE, dtype=torch.float32) + with patch("torch.compiler.is_compiling", return_value=True): + assert ( + try_fused_flux2_qkv_epilogue( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + weight, + weight, + weight, + weight, + cache, + 1e-6, + 1e-6, + ) + is None + ) + + +def test_flux2_qkv_epilogue_rejects_cuda_graph_capture() -> None: + tensor = torch.empty((1, 1, 1, HEAD_DIM), device=DEVICE, dtype=DTYPE) + weight = torch.empty((HEAD_DIM,), device=DEVICE, dtype=DTYPE) + cache = torch.empty((2, HEAD_DIM), device=DEVICE, dtype=torch.float32) + with patch("torch.cuda.is_current_stream_capturing", return_value=True): + assert ( + try_fused_flux2_qkv_epilogue( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + weight, + weight, + weight, + weight, + cache, + 1e-6, + 1e-6, + ) + is None + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_flux2_token_cat_fp8.py b/test/registered/kernels/ops/diffusion/test_flux2_token_cat_fp8.py new file mode 100644 index 000000000..230bb7d8b --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_flux2_token_cat_fp8.py @@ -0,0 +1,56 @@ +import sys +from unittest.mock import patch + +import pytest +import torch + +from sglang.kernels.ops.diffusion import try_flux2_token_cat_fp8 +from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=25, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + + +@pytest.mark.parametrize("tokens", [1, 127, 4096]) +def test_flux2_token_cat_fp8_is_bit_exact(tokens: int) -> None: + generator = torch.Generator(device="cuda") + generator.manual_seed(20260831 + tokens) + attention = torch.randn( + (1, tokens, 6144), + dtype=torch.bfloat16, + device="cuda", + generator=generator, + ) + mlp = torch.randn( + (1, tokens, 18432), + dtype=torch.bfloat16, + device="cuda", + generator=generator, + ) + scale = torch.tensor([0.013], dtype=torch.float32, device="cuda") + + expected, _ = static_quant_fp8(torch.cat([attention, mlp], dim=-1), scale) + actual = try_flux2_token_cat_fp8(attention, mlp, scale) + + assert actual is not None + assert torch.equal(actual, expected) + + +def test_flux2_token_cat_fp8_rejects_compile() -> None: + attention = torch.empty((1, 1, 16), device="cuda", dtype=torch.bfloat16) + mlp = torch.empty((1, 1, 48), device="cuda", dtype=torch.bfloat16) + scale = torch.ones((1,), device="cuda", dtype=torch.float32) + with patch("torch.compiler.is_compiling", return_value=True): + assert try_flux2_token_cat_fp8(attention, mlp, scale) is None + + +def test_flux2_token_cat_fp8_rejects_cuda_graph_capture() -> None: + attention = torch.empty((1, 1, 16), device="cuda", dtype=torch.bfloat16) + mlp = torch.empty((1, 1, 48), device="cuda", dtype=torch.bfloat16) + scale = torch.ones((1,), device="cuda", dtype=torch.float32) + with patch("torch.cuda.is_current_stream_capturing", return_value=True): + assert try_flux2_token_cat_fp8(attention, mlp, scale) is None + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/unit/models/test_flux2_fp8_norm_quant_gate.py b/test/registered/unit/models/test_flux2_fp8_norm_quant_gate.py new file mode 100644 index 000000000..f3e8fcfe8 --- /dev/null +++ b/test/registered/unit/models/test_flux2_fp8_norm_quant_gate.py @@ -0,0 +1,90 @@ +"""Unit tests for FLUX.2 ModelOpt FP8 norm+quant activation gates.""" + +import unittest +from types import SimpleNamespace + +import torch +import torch.nn as nn + +from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( + ModelOptFp8LinearMethod, +) +from sglang.multimodal_gen.runtime.models.dits.flux_2 import ( + Flux2SingleTransformerBlock, + Flux2TransformerBlock, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small") + + +def _fp8_linear(input_scale: float) -> nn.Module: + linear = nn.Module() + linear.quant_method = object.__new__(ModelOptFp8LinearMethod) + linear.register_parameter( + "input_scale", + nn.Parameter( + torch.tensor(input_scale, dtype=torch.float32, device="cuda"), + requires_grad=False, + ), + ) + return linear + + +def _double_block(scales: tuple[float, float, float]) -> Flux2TransformerBlock: + block = object.__new__(Flux2TransformerBlock) + nn.Module.__init__(block) + block.attn = SimpleNamespace( + use_fused_qkv=False, + use_fused_added_qkv=False, + to_q=_fp8_linear(scales[0]), + to_k=_fp8_linear(scales[1]), + to_v=_fp8_linear(scales[2]), + add_q_proj=_fp8_linear(scales[0]), + add_k_proj=_fp8_linear(scales[1]), + add_v_proj=_fp8_linear(scales[2]), + ) + block.ff = SimpleNamespace(linear_in=_fp8_linear(scales[0])) + block.ff_context = SimpleNamespace(linear_in=_fp8_linear(scales[0])) + block._fp8_img_attn_norm_quant = False + block._fp8_txt_attn_norm_quant = False + block._fp8_img_ff_norm_quant = False + block._fp8_txt_ff_norm_quant = False + return block + + +class TestFlux2Fp8NormQuantGate(CustomTestCase): + def test_qkv_requires_identical_input_scales(self) -> None: + matching = _double_block((0.25, 0.25, 0.25)) + mismatched = _double_block((0.25, 0.5, 0.25)) + + matching.configure_fp8_norm_quant() + mismatched.configure_fp8_norm_quant() + + self.assertTrue(matching._fp8_img_attn_norm_quant) + self.assertTrue(matching._fp8_txt_attn_norm_quant) + self.assertFalse(mismatched._fp8_img_attn_norm_quant) + self.assertFalse(mismatched._fp8_txt_attn_norm_quant) + + def test_single_block_uses_merged_projection_scale(self) -> None: + block = object.__new__(Flux2SingleTransformerBlock) + nn.Module.__init__(block) + block.attn = SimpleNamespace(to_qkv_mlp_proj=_fp8_linear(0.25)) + block._fp8_norm_quant = False + + block.configure_fp8_norm_quant() + + self.assertTrue(block._fp8_norm_quant) + + def test_nonpositive_scale_keeps_fusion_disabled(self) -> None: + block = _double_block((0.0, 0.0, 0.0)) + + block.configure_fp8_norm_quant() + + self.assertFalse(block._fp8_img_attn_norm_quant) + self.assertFalse(block._fp8_img_ff_norm_quant) + + +if __name__ == "__main__": + unittest.main()