[Diffusion] Fuse FLUX.2 ModelOpt FP8 producers and QKV packing (#37162)

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Xiaoyu Zhang
2026-09-01 16:14:29 +08:00
committed by GitHub
co-authored by Cursor
parent 379e33d87e
commit 1c3ad92438
16 changed files with 1571 additions and 104 deletions
@@ -0,0 +1,261 @@
#pragma once
#include <sgl_kernel/tensor.h>
#include <sgl_kernel/runtime.cuh>
#include <sgl_kernel/type.cuh>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <sgl_kernel/warp.cuh>
#include <algorithm>
#include <cstdint>
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<bf16_t>;
using Storage = AlignedVector<Packed, kVecSize>;
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<Storage>(input, lane);
if (kind == 2) {
store_as<Storage>(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<Storage>(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<fp32x2_t>(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<float>(kHeadDim) + eps);
#pragma unroll
for (uint32_t j = 0; j < kVecSize; ++j) {
const auto [w0, w1] = cast<fp32x2_t>(weight_vec[j]);
elems[2 * j] *= norm_factor * w0;
elems[2 * j + 1] *= norm_factor * w1;
}
const auto* cache = static_cast<const float*>(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<Packed, fp32x2_t>({elems[2 * j], elems[2 * j + 1]});
}
store_as<Storage>(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<kDLCUDA>();
TensorMatcher({NI, H, D})
.with_strides({-1, D, 1})
.with_dtype<bf16_t>()
.with_device(device)
.verify(img_q)
.verify(img_k)
.verify(img_v);
TensorMatcher({NT, H, D})
.with_strides({-1, D, 1})
.with_dtype<bf16_t>()
.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<bf16_t>().with_device(device).verify(joint_q).verify(joint_k).verify(joint_v);
TensorMatcher({D})
.with_dtype<bf16_t>()
.with_device(device)
.verify(img_q_weight)
.verify(img_k_weight)
.verify(txt_q_weight)
.verify(txt_k_weight);
TensorMatcher({-1, D}).with_dtype<fp32_t>().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<uint32_t>(NI.unwrap());
const uint32_t txt_tokens = static_cast<uint32_t>(NT.unwrap());
const uint32_t num_heads = static_cast<uint32_t>(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<float>(img_eps),
.txt_eps = static_cast<float>(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
@@ -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",
@@ -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"]
@@ -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",
@@ -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
@@ -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__(
@@ -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)
@@ -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)
@@ -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
@@ -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,