[Diffusion] Fuse Qwen-Image FP8 norm and activation quantization (#37156)
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -17,6 +17,7 @@
|
|||||||
#include <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
|
#include <sgl_kernel/tensor.h> // For TensorMatcher, SymbolicSize, SymbolicDevice
|
||||||
|
|
||||||
#include <sgl_kernel/math.cuh> // For device::math::rsqrt
|
#include <sgl_kernel/math.cuh> // For device::math::rsqrt
|
||||||
|
#include <sgl_kernel/type.cuh> // For DTypeTrait
|
||||||
#include <sgl_kernel/utils.cuh> // For SGL_DEVICE, bf16_t, LaunchKernel
|
#include <sgl_kernel/utils.cuh> // For SGL_DEVICE, bf16_t, LaunchKernel
|
||||||
#include <sgl_kernel/vec.cuh> // For AlignedVector
|
#include <sgl_kernel/vec.cuh> // For AlignedVector
|
||||||
#include <sgl_kernel/warp.cuh> // For warp::reduce_sum
|
#include <sgl_kernel/warp.cuh> // For warp::reduce_sum
|
||||||
@@ -39,12 +40,14 @@ static_assert(kWarps == 6);
|
|||||||
struct NormScaleShiftParams {
|
struct NormScaleShiftParams {
|
||||||
void* y;
|
void* y;
|
||||||
void* res_out;
|
void* res_out;
|
||||||
|
void* quantized;
|
||||||
const void* x;
|
const void* x;
|
||||||
const void* input_bias;
|
const void* input_bias;
|
||||||
const void* residual;
|
const void* residual;
|
||||||
const void* gate;
|
const void* gate;
|
||||||
const void* scale;
|
const void* scale;
|
||||||
const void* shift;
|
const void* shift;
|
||||||
|
const void* input_scale;
|
||||||
float eps;
|
float eps;
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -66,7 +69,16 @@ SGL_DEVICE float cta_reduce_sum(float v, int warp, int lane, float* scratch) {
|
|||||||
return scratch[kWarps];
|
return scratch[kWarps];
|
||||||
}
|
}
|
||||||
|
|
||||||
template <bool kHasResidual, bool kHasInputBias = false>
|
SGL_DEVICE float triton_scale_reciprocal(float scale) {
|
||||||
|
float reciprocal;
|
||||||
|
// Triton's static FP8 quantizer lowers `1.0 / scale` to div.full.f32.
|
||||||
|
// Match it exactly because a one-ULP difference at an E4M3 midpoint can
|
||||||
|
// change the quantized byte.
|
||||||
|
asm("div.full.f32 %0, %1, %2;" : "=f"(reciprocal) : "f"(1.0f), "f"(scale));
|
||||||
|
return reciprocal;
|
||||||
|
}
|
||||||
|
|
||||||
|
template <bool kHasResidual, bool kHasInputBias = false, bool kQuantizeFp8 = false>
|
||||||
__global__ void norm_scale_shift_kernel(const NormScaleShiftParams __grid_constant__ params) {
|
__global__ void norm_scale_shift_kernel(const NormScaleShiftParams __grid_constant__ params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
using Vec = AlignedVector<bf16_t, kVecElems>;
|
using Vec = AlignedVector<bf16_t, kVecElems>;
|
||||||
@@ -135,15 +147,31 @@ __global__ void norm_scale_shift_kernel(const NormScaleShiftParams __grid_consta
|
|||||||
Vec scv;
|
Vec scv;
|
||||||
Vec shv;
|
Vec shv;
|
||||||
Vec yv;
|
Vec yv;
|
||||||
|
AlignedVector<fp8_e4m3_t, kVecElems> qv;
|
||||||
scv.load(static_cast<const bf16_t*>(params.scale) + elem_offset);
|
scv.load(static_cast<const bf16_t*>(params.scale) + elem_offset);
|
||||||
shv.load(static_cast<const bf16_t*>(params.shift) + elem_offset);
|
shv.load(static_cast<const bf16_t*>(params.shift) + elem_offset);
|
||||||
|
|
||||||
|
float input_scale_inv = 0.0f;
|
||||||
|
if constexpr (kQuantizeFp8) {
|
||||||
|
input_scale_inv = triton_scale_reciprocal(*static_cast<const float*>(params.input_scale));
|
||||||
|
}
|
||||||
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int i = 0; i < kVecElems; ++i) {
|
for (int i = 0; i < kVecElems; ++i) {
|
||||||
const float norm = static_cast<float>(static_cast<bf16_t>((v[i] - mean) * factor));
|
const float norm = static_cast<float>(static_cast<bf16_t>((v[i] - mean) * factor));
|
||||||
yv[i] = static_cast<bf16_t>(norm * (1.0f + static_cast<float>(scv[i])) + static_cast<float>(shv[i]));
|
const bf16_t rounded = static_cast<bf16_t>(norm * (1.0f + static_cast<float>(scv[i])) + static_cast<float>(shv[i]));
|
||||||
|
yv[i] = rounded;
|
||||||
|
if constexpr (kQuantizeFp8) {
|
||||||
|
const float scaled = static_cast<float>(rounded) * input_scale_inv;
|
||||||
|
const float clamped =
|
||||||
|
math::min(math::max(scaled, -DTypeTrait<fp8_e4m3_t>::kFloatMax), DTypeTrait<fp8_e4m3_t>::kFloatMax);
|
||||||
|
qv[i] = static_cast<fp8_e4m3_t>(clamped);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
yv.store(static_cast<bf16_t*>(params.y) + row_offset + elem_offset);
|
yv.store(static_cast<bf16_t*>(params.y) + row_offset + elem_offset);
|
||||||
|
if constexpr (kQuantizeFp8) {
|
||||||
|
qv.store(static_cast<fp8_e4m3_t*>(params.quantized) + row_offset + elem_offset);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
__global__ void bias_mul_add_kernel(const NormScaleShiftParams __grid_constant__ params) {
|
__global__ void bias_mul_add_kernel(const NormScaleShiftParams __grid_constant__ params) {
|
||||||
@@ -197,12 +225,14 @@ struct NormScaleShiftKernel {
|
|||||||
const auto params = NormScaleShiftParams{
|
const auto params = NormScaleShiftParams{
|
||||||
.y = y.data_ptr(),
|
.y = y.data_ptr(),
|
||||||
.res_out = nullptr,
|
.res_out = nullptr,
|
||||||
|
.quantized = nullptr,
|
||||||
.x = x.data_ptr(),
|
.x = x.data_ptr(),
|
||||||
.input_bias = nullptr,
|
.input_bias = nullptr,
|
||||||
.residual = nullptr,
|
.residual = nullptr,
|
||||||
.gate = nullptr,
|
.gate = nullptr,
|
||||||
.scale = scale.data_ptr(),
|
.scale = scale.data_ptr(),
|
||||||
.shift = shift.data_ptr(),
|
.shift = shift.data_ptr(),
|
||||||
|
.input_scale = nullptr,
|
||||||
.eps = static_cast<float>(eps),
|
.eps = static_cast<float>(eps),
|
||||||
};
|
};
|
||||||
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<false>, params);
|
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<false>, params);
|
||||||
@@ -237,18 +267,105 @@ struct ScaleResidualNormScaleShiftKernel {
|
|||||||
const auto params = NormScaleShiftParams{
|
const auto params = NormScaleShiftParams{
|
||||||
.y = y.data_ptr(),
|
.y = y.data_ptr(),
|
||||||
.res_out = res_out.data_ptr(),
|
.res_out = res_out.data_ptr(),
|
||||||
|
.quantized = nullptr,
|
||||||
.x = x.data_ptr(),
|
.x = x.data_ptr(),
|
||||||
.input_bias = nullptr,
|
.input_bias = nullptr,
|
||||||
.residual = residual.data_ptr(),
|
.residual = residual.data_ptr(),
|
||||||
.gate = gate.data_ptr(),
|
.gate = gate.data_ptr(),
|
||||||
.scale = scale.data_ptr(),
|
.scale = scale.data_ptr(),
|
||||||
.shift = shift.data_ptr(),
|
.shift = shift.data_ptr(),
|
||||||
|
.input_scale = nullptr,
|
||||||
.eps = static_cast<float>(eps),
|
.eps = static_cast<float>(eps),
|
||||||
};
|
};
|
||||||
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<true>, params);
|
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<true>, params);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
/** \brief Fuse Qwen LayerNorm/modulation with static E4M3 activation quantization. */
|
||||||
|
struct NormScaleShiftFp8Kernel {
|
||||||
|
static void
|
||||||
|
run(tvm::ffi::TensorView y,
|
||||||
|
tvm::ffi::TensorView quantized,
|
||||||
|
tvm::ffi::TensorView x,
|
||||||
|
tvm::ffi::TensorView scale,
|
||||||
|
tvm::ffi::TensorView shift,
|
||||||
|
tvm::ffi::TensorView input_scale,
|
||||||
|
double eps) {
|
||||||
|
using namespace host;
|
||||||
|
auto N = SymbolicSize{"num_rows"};
|
||||||
|
auto device = SymbolicDevice{};
|
||||||
|
device.set_options<kDLCUDA>();
|
||||||
|
|
||||||
|
TensorMatcher({N, kHidden}).with_dtype<bf16_t>().with_device(device).verify(x).verify(y);
|
||||||
|
TensorMatcher({N, kHidden}).with_dtype<fp8_e4m3_t>().with_device(device).verify(quantized);
|
||||||
|
TensorMatcher({kHidden}).with_dtype<bf16_t>().with_device(device).verify(scale).verify(shift);
|
||||||
|
TensorMatcher({1}).with_dtype<fp32_t>().with_device(device).verify(input_scale);
|
||||||
|
|
||||||
|
const uint32_t grid = verify_nss_geometry(N);
|
||||||
|
const auto params = NormScaleShiftParams{
|
||||||
|
.y = y.data_ptr(),
|
||||||
|
.res_out = nullptr,
|
||||||
|
.quantized = quantized.data_ptr(),
|
||||||
|
.x = x.data_ptr(),
|
||||||
|
.input_bias = nullptr,
|
||||||
|
.residual = nullptr,
|
||||||
|
.gate = nullptr,
|
||||||
|
.scale = scale.data_ptr(),
|
||||||
|
.shift = shift.data_ptr(),
|
||||||
|
.input_scale = input_scale.data_ptr(),
|
||||||
|
.eps = static_cast<float>(eps),
|
||||||
|
};
|
||||||
|
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<false, false, true>, params);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
/** \brief Fuse Qwen residual LayerNorm/modulation with static E4M3 activation quantization. */
|
||||||
|
struct ScaleResidualNormScaleShiftFp8Kernel {
|
||||||
|
static void
|
||||||
|
run(tvm::ffi::TensorView y,
|
||||||
|
tvm::ffi::TensorView quantized,
|
||||||
|
tvm::ffi::TensorView res_out,
|
||||||
|
tvm::ffi::TensorView residual,
|
||||||
|
tvm::ffi::TensorView x,
|
||||||
|
tvm::ffi::TensorView gate,
|
||||||
|
tvm::ffi::TensorView scale,
|
||||||
|
tvm::ffi::TensorView shift,
|
||||||
|
tvm::ffi::TensorView input_scale,
|
||||||
|
double eps) {
|
||||||
|
using namespace host;
|
||||||
|
auto N = SymbolicSize{"num_rows"};
|
||||||
|
auto device = SymbolicDevice{};
|
||||||
|
device.set_options<kDLCUDA>();
|
||||||
|
|
||||||
|
TensorMatcher({N, kHidden})
|
||||||
|
.with_dtype<bf16_t>()
|
||||||
|
.with_device(device)
|
||||||
|
.verify(x)
|
||||||
|
.verify(residual)
|
||||||
|
.verify(y)
|
||||||
|
.verify(res_out);
|
||||||
|
TensorMatcher({N, kHidden}).with_dtype<fp8_e4m3_t>().with_device(device).verify(quantized);
|
||||||
|
TensorMatcher({kHidden}).with_dtype<bf16_t>().with_device(device).verify(gate).verify(scale).verify(shift);
|
||||||
|
TensorMatcher({1}).with_dtype<fp32_t>().with_device(device).verify(input_scale);
|
||||||
|
|
||||||
|
const uint32_t grid = verify_nss_geometry(N);
|
||||||
|
const auto params = NormScaleShiftParams{
|
||||||
|
.y = y.data_ptr(),
|
||||||
|
.res_out = res_out.data_ptr(),
|
||||||
|
.quantized = quantized.data_ptr(),
|
||||||
|
.x = x.data_ptr(),
|
||||||
|
.input_bias = nullptr,
|
||||||
|
.residual = residual.data_ptr(),
|
||||||
|
.gate = gate.data_ptr(),
|
||||||
|
.scale = scale.data_ptr(),
|
||||||
|
.shift = shift.data_ptr(),
|
||||||
|
.input_scale = input_scale.data_ptr(),
|
||||||
|
.eps = static_cast<float>(eps),
|
||||||
|
};
|
||||||
|
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<true, false, true>, params);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
struct BiasScaleResidualNormScaleShiftKernel {
|
struct BiasScaleResidualNormScaleShiftKernel {
|
||||||
static void
|
static void
|
||||||
run(tvm::ffi::TensorView y,
|
run(tvm::ffi::TensorView y,
|
||||||
@@ -284,12 +401,14 @@ struct BiasScaleResidualNormScaleShiftKernel {
|
|||||||
const auto params = NormScaleShiftParams{
|
const auto params = NormScaleShiftParams{
|
||||||
.y = y.data_ptr(),
|
.y = y.data_ptr(),
|
||||||
.res_out = res_out.data_ptr(),
|
.res_out = res_out.data_ptr(),
|
||||||
|
.quantized = nullptr,
|
||||||
.x = x.data_ptr(),
|
.x = x.data_ptr(),
|
||||||
.input_bias = input_bias.data_ptr(),
|
.input_bias = input_bias.data_ptr(),
|
||||||
.residual = residual.data_ptr(),
|
.residual = residual.data_ptr(),
|
||||||
.gate = gate.data_ptr(),
|
.gate = gate.data_ptr(),
|
||||||
.scale = scale.data_ptr(),
|
.scale = scale.data_ptr(),
|
||||||
.shift = shift.data_ptr(),
|
.shift = shift.data_ptr(),
|
||||||
|
.input_scale = nullptr,
|
||||||
.eps = static_cast<float>(eps),
|
.eps = static_cast<float>(eps),
|
||||||
};
|
};
|
||||||
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<true, true>, params);
|
LaunchKernel(grid, kThreads, device.unwrap())(norm_scale_shift_kernel<true, true>, params);
|
||||||
@@ -315,12 +434,14 @@ struct BiasMulAddKernel {
|
|||||||
const auto params = NormScaleShiftParams{
|
const auto params = NormScaleShiftParams{
|
||||||
.y = y.data_ptr(),
|
.y = y.data_ptr(),
|
||||||
.res_out = nullptr,
|
.res_out = nullptr,
|
||||||
|
.quantized = nullptr,
|
||||||
.x = x.data_ptr(),
|
.x = x.data_ptr(),
|
||||||
.input_bias = input_bias.data_ptr(),
|
.input_bias = input_bias.data_ptr(),
|
||||||
.residual = residual.data_ptr(),
|
.residual = residual.data_ptr(),
|
||||||
.gate = gate.data_ptr(),
|
.gate = gate.data_ptr(),
|
||||||
.scale = nullptr,
|
.scale = nullptr,
|
||||||
.shift = nullptr,
|
.shift = nullptr,
|
||||||
|
.input_scale = nullptr,
|
||||||
.eps = 0.0f,
|
.eps = 0.0f,
|
||||||
};
|
};
|
||||||
LaunchKernel(grid, kThreads, device.unwrap())(bias_mul_add_kernel, params);
|
LaunchKernel(grid, kThreads, device.unwrap())(bias_mul_add_kernel, params);
|
||||||
|
|||||||
@@ -382,6 +382,10 @@ _EXPORTS: dict[str, str] = {
|
|||||||
"fused_scale_residual_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact",
|
"fused_scale_residual_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact",
|
||||||
"fused_norm_scale_shift": "norm.scale_residual_norm_cutedsl",
|
"fused_norm_scale_shift": "norm.scale_residual_norm_cutedsl",
|
||||||
"fused_scale_residual_norm_scale_shift": "norm.scale_residual_norm_cutedsl",
|
"fused_scale_residual_norm_scale_shift": "norm.scale_residual_norm_cutedsl",
|
||||||
|
"fused_norm_scale_shift_fp8": "norm.norm_scale_shift_jit",
|
||||||
|
"fused_scale_residual_norm_scale_shift_fp8": "norm.norm_scale_shift_jit",
|
||||||
|
"try_fused_norm_scale_shift_fp8": "norm.norm_scale_shift_jit",
|
||||||
|
"try_fused_scale_residual_norm_scale_shift_fp8": "norm.norm_scale_shift_jit",
|
||||||
"validate_scale_shift": "norm.scale_residual_norm_cutedsl",
|
"validate_scale_shift": "norm.scale_residual_norm_cutedsl",
|
||||||
"can_use_wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton",
|
"can_use_wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton",
|
||||||
"wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton",
|
"wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton",
|
||||||
|
|||||||
@@ -67,6 +67,11 @@ def _row_bf16(t, device: torch.device):
|
|||||||
|
|
||||||
@cache_once
|
@cache_once
|
||||||
def norm_scale_shift_module() -> Module:
|
def norm_scale_shift_module() -> Module:
|
||||||
|
device = torch.device("cuda", torch.cuda.current_device())
|
||||||
|
if not _blackwell_or_newer(device):
|
||||||
|
raise RuntimeError(
|
||||||
|
"Qwen-Image norm-scale-shift JIT kernels require NVIDIA Blackwell or newer"
|
||||||
|
)
|
||||||
return load_jit(
|
return load_jit(
|
||||||
"norm_scale_shift_native",
|
"norm_scale_shift_native",
|
||||||
cuda_files=["diffusion/norm_scale_shift.cuh"],
|
cuda_files=["diffusion/norm_scale_shift.cuh"],
|
||||||
@@ -79,6 +84,14 @@ def norm_scale_shift_module() -> Module:
|
|||||||
"srnss_bf16_row",
|
"srnss_bf16_row",
|
||||||
"norm_scale_shift::ScaleResidualNormScaleShiftKernel::run",
|
"norm_scale_shift::ScaleResidualNormScaleShiftKernel::run",
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
"nss_fp8_row",
|
||||||
|
"norm_scale_shift::NormScaleShiftFp8Kernel::run",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"srnss_fp8_row",
|
||||||
|
"norm_scale_shift::ScaleResidualNormScaleShiftFp8Kernel::run",
|
||||||
|
),
|
||||||
(
|
(
|
||||||
"bias_srnss_bf16_row",
|
"bias_srnss_bf16_row",
|
||||||
"norm_scale_shift::BiasScaleResidualNormScaleShiftKernel::run",
|
"norm_scale_shift::BiasScaleResidualNormScaleShiftKernel::run",
|
||||||
@@ -94,6 +107,44 @@ def norm_scale_shift_module() -> Module:
|
|||||||
_module = norm_scale_shift_module
|
_module = norm_scale_shift_module
|
||||||
|
|
||||||
|
|
||||||
|
def fused_norm_scale_shift_fp8(x, scale, shift, input_scale, eps):
|
||||||
|
"""Return exact BF16 modulation output and its static E4M3 quantization."""
|
||||||
|
normalized = torch.empty_like(x)
|
||||||
|
quantized = torch.empty_like(x, dtype=torch.float8_e4m3fn)
|
||||||
|
_module().nss_fp8_row(
|
||||||
|
normalized.view(-1, _HIDDEN),
|
||||||
|
quantized.view(-1, _HIDDEN),
|
||||||
|
x.view(-1, _HIDDEN),
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
input_scale.reshape(1),
|
||||||
|
float(eps),
|
||||||
|
)
|
||||||
|
return normalized, quantized
|
||||||
|
|
||||||
|
|
||||||
|
def fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual, x, gate, scale, shift, input_scale, eps
|
||||||
|
):
|
||||||
|
"""Return exact BF16 residual/modulation outputs and E4M3 quantization."""
|
||||||
|
normalized = torch.empty_like(x)
|
||||||
|
quantized = torch.empty_like(x, dtype=torch.float8_e4m3fn)
|
||||||
|
residual_out = torch.empty_like(x)
|
||||||
|
_module().srnss_fp8_row(
|
||||||
|
normalized.view(-1, _HIDDEN),
|
||||||
|
quantized.view(-1, _HIDDEN),
|
||||||
|
residual_out.view(-1, _HIDDEN),
|
||||||
|
residual.view(-1, _HIDDEN),
|
||||||
|
x.view(-1, _HIDDEN),
|
||||||
|
gate,
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
input_scale.reshape(1),
|
||||||
|
float(eps),
|
||||||
|
)
|
||||||
|
return normalized, quantized, residual_out
|
||||||
|
|
||||||
|
|
||||||
def try_fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps):
|
def try_fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps):
|
||||||
if norm_type != "layer" or weight is not None or bias is not None:
|
if norm_type != "layer" or weight is not None or bias is not None:
|
||||||
return None
|
return None
|
||||||
@@ -145,6 +196,67 @@ def try_fused_scale_residual_norm_scale_shift(
|
|||||||
return y, residual_out
|
return y, residual_out
|
||||||
|
|
||||||
|
|
||||||
|
def try_fused_norm_scale_shift_fp8(
|
||||||
|
x, weight, bias, scale, shift, input_scale, norm_type, eps
|
||||||
|
):
|
||||||
|
if norm_type != "layer" or weight is not None or bias is not None:
|
||||||
|
return None
|
||||||
|
if not _nss_activation(x) or not _blackwell_or_newer(x.device):
|
||||||
|
return None
|
||||||
|
|
||||||
|
scale = _row_bf16(scale, x.device)
|
||||||
|
shift = _row_bf16(shift, x.device)
|
||||||
|
if scale is None or shift is None:
|
||||||
|
return None
|
||||||
|
if not _fp8_input_scale(input_scale, x.device):
|
||||||
|
return None
|
||||||
|
return fused_norm_scale_shift_fp8(x, scale, shift, input_scale, eps)
|
||||||
|
|
||||||
|
|
||||||
|
def try_fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual,
|
||||||
|
x,
|
||||||
|
gate,
|
||||||
|
weight,
|
||||||
|
bias,
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
input_scale,
|
||||||
|
norm_type,
|
||||||
|
eps,
|
||||||
|
):
|
||||||
|
if norm_type != "layer" or weight is not None or bias is not None:
|
||||||
|
return None
|
||||||
|
if not (
|
||||||
|
_nss_activation(x)
|
||||||
|
and _nss_activation(residual, x)
|
||||||
|
and _blackwell_or_newer(x.device)
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
|
||||||
|
gate = _row_bf16(gate, x.device)
|
||||||
|
scale = _row_bf16(scale, x.device)
|
||||||
|
shift = _row_bf16(shift, x.device)
|
||||||
|
if gate is None or scale is None or shift is None:
|
||||||
|
return None
|
||||||
|
if not _fp8_input_scale(input_scale, x.device):
|
||||||
|
return None
|
||||||
|
return fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual, x, gate, scale, shift, input_scale, eps
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _fp8_input_scale(t, device: torch.device) -> bool:
|
||||||
|
return (
|
||||||
|
isinstance(t, torch.Tensor)
|
||||||
|
and t.is_cuda
|
||||||
|
and t.device == device
|
||||||
|
and t.dtype == torch.float32
|
||||||
|
and t.numel() == 1
|
||||||
|
and t.is_contiguous()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def try_fused_bias_scale_residual_norm_scale_shift(
|
def try_fused_bias_scale_residual_norm_scale_shift(
|
||||||
residual, x, input_bias, gate, weight, bias, scale, shift, norm_type, eps
|
residual, x, input_bias, gate, weight, bias, scale, shift, norm_type, eps
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -22,6 +22,8 @@ from sglang.kernels.ops.diffusion import (
|
|||||||
mark_fused_gelu_site,
|
mark_fused_gelu_site,
|
||||||
try_fused_bias_mul_add,
|
try_fused_bias_mul_add,
|
||||||
try_fused_bias_scale_residual_norm_scale_shift,
|
try_fused_bias_scale_residual_norm_scale_shift,
|
||||||
|
try_fused_norm_scale_shift_fp8,
|
||||||
|
try_fused_scale_residual_norm_scale_shift_fp8,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
|
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
|
||||||
from sglang.multimodal_gen.configs.models.fsdp import is_transformer_block
|
from sglang.multimodal_gen.configs.models.fsdp import is_transformer_block
|
||||||
@@ -70,6 +72,9 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i
|
|||||||
NunchakuConfig,
|
NunchakuConfig,
|
||||||
is_nunchaku_available,
|
is_nunchaku_available,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
|
||||||
|
ModelOptFp8LinearMethod,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
apply_flashinfer_rope_qk_inplace,
|
apply_flashinfer_rope_qk_inplace,
|
||||||
)
|
)
|
||||||
@@ -1125,6 +1130,65 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
self.img_mlp = NunchakuFeedForward(self.img_mlp, **nunchaku_kwargs)
|
self.img_mlp = NunchakuFeedForward(self.img_mlp, **nunchaku_kwargs)
|
||||||
self.txt_mlp = NunchakuFeedForward(self.txt_mlp, **nunchaku_kwargs)
|
self.txt_mlp = NunchakuFeedForward(self.txt_mlp, **nunchaku_kwargs)
|
||||||
|
|
||||||
|
self._fp8_img_attn_norm_quant = False
|
||||||
|
self._fp8_txt_attn_norm_quant = False
|
||||||
|
self._fp8_img_mlp_norm_quant = False
|
||||||
|
self._fp8_txt_mlp_norm_quant = False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
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())
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _shared_modelopt_fp8_scale(cls, linears: list[nn.Module]) -> bool:
|
||||||
|
if not all(cls._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 configure_fp8_norm_quant(self) -> None:
|
||||||
|
"""Enable exact norm+quant paths after checkpoint scales are materialized."""
|
||||||
|
if not torch.cuda.is_available():
|
||||||
|
return
|
||||||
|
capability = torch.cuda.get_device_capability()
|
||||||
|
if self.dim != 3072 or capability[0] < 10 or self.zero_cond_t:
|
||||||
|
return
|
||||||
|
if self.attn.use_fused_qkv:
|
||||||
|
self._fp8_img_attn_norm_quant = self._valid_modelopt_fp8_linear(
|
||||||
|
self.attn.to_qkv
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._fp8_img_attn_norm_quant = self._shared_modelopt_fp8_scale(
|
||||||
|
[self.attn.to_q, self.attn.to_k, self.attn.to_v]
|
||||||
|
)
|
||||||
|
if self.attn.added_kv_proj_dim is not None:
|
||||||
|
if self.attn.use_fused_added_qkv:
|
||||||
|
self._fp8_txt_attn_norm_quant = self._valid_modelopt_fp8_linear(
|
||||||
|
self.attn.to_added_qkv
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._fp8_txt_attn_norm_quant = self._shared_modelopt_fp8_scale(
|
||||||
|
[self.attn.add_q_proj, self.attn.add_k_proj, self.attn.add_v_proj]
|
||||||
|
)
|
||||||
|
if isinstance(self.img_mlp, QwenImageFeedForward):
|
||||||
|
self._fp8_img_mlp_norm_quant = self._valid_modelopt_fp8_linear(
|
||||||
|
self.img_mlp.net[0].proj
|
||||||
|
)
|
||||||
|
if isinstance(self.txt_mlp, QwenImageFeedForward):
|
||||||
|
self._fp8_txt_mlp_norm_quant = self._valid_modelopt_fp8_linear(
|
||||||
|
self.txt_mlp.net[0].proj
|
||||||
|
)
|
||||||
|
|
||||||
def _norm_scale_shift(
|
def _norm_scale_shift(
|
||||||
self,
|
self,
|
||||||
norm_module: LayerNormScaleShift,
|
norm_module: LayerNormScaleShift,
|
||||||
@@ -1200,6 +1264,68 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
return img_mod_params, txt_mod_params
|
return img_mod_params, txt_mod_params
|
||||||
|
|
||||||
|
def _try_fp8_norm_quant(
|
||||||
|
self,
|
||||||
|
norm_module: LayerNormScaleShift,
|
||||||
|
*,
|
||||||
|
x: torch.Tensor,
|
||||||
|
mod_params: torch.Tensor,
|
||||||
|
input_scale: Optional[torch.Tensor],
|
||||||
|
enabled: bool,
|
||||||
|
modulate_index: Optional[torch.Tensor],
|
||||||
|
use_bcg_helpers: bool,
|
||||||
|
) -> Optional[tuple[torch.Tensor, torch.Tensor, torch.Tensor]]:
|
||||||
|
if not enabled or modulate_index is not None or use_bcg_helpers:
|
||||||
|
return None
|
||||||
|
shift, scale, gate = mod_params.chunk(3, dim=-1)
|
||||||
|
result = try_fused_norm_scale_shift_fp8(
|
||||||
|
x,
|
||||||
|
getattr(norm_module.norm, "weight", None),
|
||||||
|
getattr(norm_module.norm, "bias", None),
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
input_scale,
|
||||||
|
norm_module.norm_type,
|
||||||
|
norm_module.eps,
|
||||||
|
)
|
||||||
|
if result is None:
|
||||||
|
return None
|
||||||
|
normalized, quantized = result
|
||||||
|
return quantized, gate.unsqueeze(1), normalized
|
||||||
|
|
||||||
|
def _try_fp8_residual_norm_quant(
|
||||||
|
self,
|
||||||
|
norm_module: ScaleResidualLayerNormScaleShift,
|
||||||
|
*,
|
||||||
|
residual: torch.Tensor,
|
||||||
|
x: torch.Tensor,
|
||||||
|
residual_gate: torch.Tensor,
|
||||||
|
mod_params: torch.Tensor,
|
||||||
|
input_scale: Optional[torch.Tensor],
|
||||||
|
enabled: bool,
|
||||||
|
modulate_index: Optional[torch.Tensor],
|
||||||
|
use_bcg_helpers: bool,
|
||||||
|
) -> Optional[tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]]:
|
||||||
|
if not enabled or modulate_index is not None or use_bcg_helpers:
|
||||||
|
return None
|
||||||
|
shift, scale, gate = mod_params.chunk(3, dim=-1)
|
||||||
|
result = try_fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual,
|
||||||
|
x,
|
||||||
|
residual_gate,
|
||||||
|
getattr(norm_module.norm, "weight", None),
|
||||||
|
getattr(norm_module.norm, "bias", None),
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
input_scale,
|
||||||
|
norm_module.norm_type,
|
||||||
|
norm_module.eps,
|
||||||
|
)
|
||||||
|
if result is None:
|
||||||
|
return None
|
||||||
|
normalized, quantized, residual_out = result
|
||||||
|
return quantized, residual_out, gate.unsqueeze(1), normalized
|
||||||
|
|
||||||
def _modulate(
|
def _modulate(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
@@ -1361,16 +1487,57 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
use_bcg_helpers = is_in_breakable_cuda_graph()
|
use_bcg_helpers = is_in_breakable_cuda_graph()
|
||||||
|
|
||||||
# Process image stream - norm1 + modulation
|
# Process image stream - norm1 + modulation
|
||||||
img_modulated, img_gate1 = self._modulate(
|
img_fp8 = self._try_fp8_norm_quant(
|
||||||
hidden_states,
|
|
||||||
img_mod1,
|
|
||||||
self.img_norm1,
|
self.img_norm1,
|
||||||
modulate_index,
|
x=hidden_states,
|
||||||
|
mod_params=img_mod1,
|
||||||
|
input_scale=(
|
||||||
|
(
|
||||||
|
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
|
||||||
|
),
|
||||||
|
enabled=self._fp8_img_attn_norm_quant,
|
||||||
|
modulate_index=modulate_index,
|
||||||
use_bcg_helpers=use_bcg_helpers,
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
)
|
)
|
||||||
|
if img_fp8 is None:
|
||||||
|
img_modulated, img_gate1 = self._modulate(
|
||||||
|
hidden_states,
|
||||||
|
img_mod1,
|
||||||
|
self.img_norm1,
|
||||||
|
modulate_index,
|
||||||
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
|
)
|
||||||
|
img_modulated_bf16 = None
|
||||||
|
else:
|
||||||
|
img_modulated, img_gate1, img_modulated_bf16 = img_fp8
|
||||||
# Process text stream - norm1 + modulation
|
# Process text stream - norm1 + modulation
|
||||||
|
txt_fp8 = self._try_fp8_norm_quant(
|
||||||
|
self.txt_norm1,
|
||||||
|
x=encoder_hidden_states,
|
||||||
|
mod_params=txt_mod1,
|
||||||
|
input_scale=(
|
||||||
|
(
|
||||||
|
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
|
||||||
|
),
|
||||||
|
enabled=self._fp8_txt_attn_norm_quant,
|
||||||
|
modulate_index=modulate_index,
|
||||||
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
|
)
|
||||||
txt_shift1, txt_scale1, txt_gate1_raw = txt_mod1.chunk(3, dim=-1)
|
txt_shift1, txt_scale1, txt_gate1_raw = txt_mod1.chunk(3, dim=-1)
|
||||||
if use_bcg_helpers:
|
if txt_fp8 is not None:
|
||||||
|
txt_modulated, txt_gate1, txt_modulated_bf16 = txt_fp8
|
||||||
|
elif use_bcg_helpers:
|
||||||
|
txt_modulated_bf16 = None
|
||||||
txt_modulated = self._norm_scale_shift(
|
txt_modulated = self._norm_scale_shift(
|
||||||
self.txt_norm1,
|
self.txt_norm1,
|
||||||
encoder_hidden_states,
|
encoder_hidden_states,
|
||||||
@@ -1378,10 +1545,12 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
scale=txt_scale1,
|
scale=txt_scale1,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
txt_modulated_bf16 = None
|
||||||
txt_modulated = self.txt_norm1(
|
txt_modulated = self.txt_norm1(
|
||||||
encoder_hidden_states, shift=txt_shift1, scale=txt_scale1
|
encoder_hidden_states, shift=txt_shift1, scale=txt_scale1
|
||||||
)
|
)
|
||||||
txt_gate1 = txt_gate1_raw.unsqueeze(1)
|
if txt_fp8 is None:
|
||||||
|
txt_gate1 = txt_gate1_raw.unsqueeze(1)
|
||||||
|
|
||||||
# Use QwenAttnProcessor2_0 for joint attention computation
|
# Use QwenAttnProcessor2_0 for joint attention computation
|
||||||
# This directly implements the DoubleStreamLayerMegatron logic:
|
# This directly implements the DoubleStreamLayerMegatron logic:
|
||||||
@@ -1399,6 +1568,7 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
image_rotary_emb=image_rotary_emb,
|
image_rotary_emb=image_rotary_emb,
|
||||||
**joint_attention_kwargs,
|
**joint_attention_kwargs,
|
||||||
)
|
)
|
||||||
|
del img_modulated_bf16, txt_modulated_bf16
|
||||||
|
|
||||||
# QwenAttnProcessor2_0 returns (img_output, txt_output) when encoder_hidden_states is provided
|
# QwenAttnProcessor2_0 returns (img_output, txt_output) when encoder_hidden_states is provided
|
||||||
(
|
(
|
||||||
@@ -1408,16 +1578,40 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
txt_attn_bias,
|
txt_attn_bias,
|
||||||
) = attn_output
|
) = attn_output
|
||||||
# Process image stream - norm2 + MLP
|
# Process image stream - norm2 + MLP
|
||||||
img_modulated2, hidden_states, img_gate2 = self._modulate(
|
img_fp8_mlp = self._try_fp8_residual_norm_quant(
|
||||||
img_attn_output,
|
|
||||||
img_mod2,
|
|
||||||
self.img_norm2,
|
self.img_norm2,
|
||||||
modulate_index,
|
residual=hidden_states,
|
||||||
gate_x=img_gate1,
|
x=img_attn_output,
|
||||||
residual_x=hidden_states,
|
residual_gate=img_gate1,
|
||||||
x_bias=img_attn_bias,
|
mod_params=img_mod2,
|
||||||
|
input_scale=(
|
||||||
|
self.img_mlp.net[0].proj.input_scale
|
||||||
|
if self._fp8_img_mlp_norm_quant
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
enabled=self._fp8_img_mlp_norm_quant,
|
||||||
|
modulate_index=modulate_index,
|
||||||
use_bcg_helpers=use_bcg_helpers,
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
)
|
)
|
||||||
|
if img_fp8_mlp is None:
|
||||||
|
img_modulated2, hidden_states, img_gate2 = self._modulate(
|
||||||
|
img_attn_output,
|
||||||
|
img_mod2,
|
||||||
|
self.img_norm2,
|
||||||
|
modulate_index,
|
||||||
|
gate_x=img_gate1,
|
||||||
|
residual_x=hidden_states,
|
||||||
|
x_bias=img_attn_bias,
|
||||||
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
|
)
|
||||||
|
img_modulated2_bf16 = None
|
||||||
|
else:
|
||||||
|
(
|
||||||
|
img_modulated2,
|
||||||
|
hidden_states,
|
||||||
|
img_gate2,
|
||||||
|
img_modulated2_bf16,
|
||||||
|
) = img_fp8_mlp
|
||||||
if isinstance(self.img_mlp, QwenImageFeedForward):
|
if isinstance(self.img_mlp, QwenImageFeedForward):
|
||||||
img_mlp_output, img_mlp_bias = self.img_mlp.forward_with_bias(
|
img_mlp_output, img_mlp_bias = self.img_mlp.forward_with_bias(
|
||||||
img_modulated2
|
img_modulated2
|
||||||
@@ -1425,6 +1619,7 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
else:
|
else:
|
||||||
img_mlp_output = self.img_mlp(img_modulated2)
|
img_mlp_output = self.img_mlp(img_modulated2)
|
||||||
img_mlp_bias = None
|
img_mlp_bias = None
|
||||||
|
del img_modulated2_bf16
|
||||||
|
|
||||||
if img_mlp_output.dim() == 2:
|
if img_mlp_output.dim() == 2:
|
||||||
img_mlp_output = img_mlp_output.unsqueeze(0)
|
img_mlp_output = img_mlp_output.unsqueeze(0)
|
||||||
@@ -1438,7 +1633,30 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
|
|
||||||
# Process text stream - norm2 + MLP
|
# Process text stream - norm2 + MLP
|
||||||
txt_shift2, txt_scale2, txt_gate2_raw = txt_mod2.chunk(3, dim=-1)
|
txt_shift2, txt_scale2, txt_gate2_raw = txt_mod2.chunk(3, dim=-1)
|
||||||
if use_bcg_helpers:
|
txt_fp8_mlp = self._try_fp8_residual_norm_quant(
|
||||||
|
self.txt_norm2,
|
||||||
|
residual=encoder_hidden_states,
|
||||||
|
x=txt_attn_output,
|
||||||
|
residual_gate=txt_gate1,
|
||||||
|
mod_params=txt_mod2,
|
||||||
|
input_scale=(
|
||||||
|
self.txt_mlp.net[0].proj.input_scale
|
||||||
|
if self._fp8_txt_mlp_norm_quant
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
enabled=self._fp8_txt_mlp_norm_quant,
|
||||||
|
modulate_index=modulate_index,
|
||||||
|
use_bcg_helpers=use_bcg_helpers,
|
||||||
|
)
|
||||||
|
if txt_fp8_mlp is not None:
|
||||||
|
(
|
||||||
|
txt_modulated2,
|
||||||
|
encoder_hidden_states,
|
||||||
|
txt_gate2,
|
||||||
|
txt_modulated2_bf16,
|
||||||
|
) = txt_fp8_mlp
|
||||||
|
elif use_bcg_helpers:
|
||||||
|
txt_modulated2_bf16 = None
|
||||||
if txt_attn_bias is not None:
|
if txt_attn_bias is not None:
|
||||||
txt_attn_output = txt_attn_output + txt_attn_bias
|
txt_attn_output = txt_attn_output + txt_attn_bias
|
||||||
(
|
(
|
||||||
@@ -1453,6 +1671,7 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
scale=txt_scale2,
|
scale=txt_scale2,
|
||||||
)
|
)
|
||||||
elif txt_attn_bias is not None:
|
elif txt_attn_bias is not None:
|
||||||
|
txt_modulated2_bf16 = None
|
||||||
txt_modulated2, encoder_hidden_states, _ = self._modulate(
|
txt_modulated2, encoder_hidden_states, _ = self._modulate(
|
||||||
txt_attn_output,
|
txt_attn_output,
|
||||||
txt_mod2,
|
txt_mod2,
|
||||||
@@ -1462,6 +1681,7 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
x_bias=txt_attn_bias,
|
x_bias=txt_attn_bias,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
txt_modulated2_bf16 = None
|
||||||
txt_modulated2, encoder_hidden_states = self.txt_norm2(
|
txt_modulated2, encoder_hidden_states = self.txt_norm2(
|
||||||
residual=encoder_hidden_states,
|
residual=encoder_hidden_states,
|
||||||
x=txt_attn_output,
|
x=txt_attn_output,
|
||||||
@@ -1469,7 +1689,8 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
shift=txt_shift2,
|
shift=txt_shift2,
|
||||||
scale=txt_scale2,
|
scale=txt_scale2,
|
||||||
)
|
)
|
||||||
txt_gate2 = txt_gate2_raw.unsqueeze(1)
|
if txt_fp8_mlp is None:
|
||||||
|
txt_gate2 = txt_gate2_raw.unsqueeze(1)
|
||||||
if isinstance(self.txt_mlp, QwenImageFeedForward):
|
if isinstance(self.txt_mlp, QwenImageFeedForward):
|
||||||
txt_mlp_output, txt_mlp_bias = self.txt_mlp.forward_with_bias(
|
txt_mlp_output, txt_mlp_bias = self.txt_mlp.forward_with_bias(
|
||||||
txt_modulated2
|
txt_modulated2
|
||||||
@@ -1477,6 +1698,7 @@ class QwenImageTransformerBlock(nn.Module):
|
|||||||
else:
|
else:
|
||||||
txt_mlp_output = self.txt_mlp(txt_modulated2)
|
txt_mlp_output = self.txt_mlp(txt_modulated2)
|
||||||
txt_mlp_bias = None
|
txt_mlp_bias = None
|
||||||
|
del txt_modulated2_bf16
|
||||||
|
|
||||||
if txt_mlp_output.dim() == 2:
|
if txt_mlp_output.dim() == 2:
|
||||||
txt_mlp_output = txt_mlp_output.unsqueeze(0)
|
txt_mlp_output = txt_mlp_output.unsqueeze(0)
|
||||||
@@ -1632,6 +1854,24 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
|
|
||||||
self.layer_names = ["transformer_blocks"]
|
self.layer_names = ["transformer_blocks"]
|
||||||
|
|
||||||
|
def post_load_weights(self) -> None:
|
||||||
|
super().post_load_weights()
|
||||||
|
for block in self.transformer_blocks:
|
||||||
|
block.configure_fp8_norm_quant()
|
||||||
|
enabled = sum(
|
||||||
|
block._fp8_img_attn_norm_quant
|
||||||
|
+ block._fp8_txt_attn_norm_quant
|
||||||
|
+ block._fp8_img_mlp_norm_quant
|
||||||
|
+ block._fp8_txt_mlp_norm_quant
|
||||||
|
for block in self.transformer_blocks
|
||||||
|
)
|
||||||
|
if enabled:
|
||||||
|
logger.info(
|
||||||
|
"Enabled Qwen FP8 norm+quant fusion for %d/%d block paths",
|
||||||
|
enabled,
|
||||||
|
4 * len(self.transformer_blocks),
|
||||||
|
)
|
||||||
|
|
||||||
@functools.lru_cache(maxsize=50)
|
@functools.lru_cache(maxsize=50)
|
||||||
def build_modulate_index(self, img_shapes: tuple[int, int, int], device):
|
def build_modulate_index(self, img_shapes: tuple[int, int, int], device):
|
||||||
sp_world_size = get_sp_world_size()
|
sp_world_size = get_sp_world_size()
|
||||||
|
|||||||
@@ -0,0 +1,82 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.jit.benchmark import marker
|
||||||
|
from sglang.kernels.ops.diffusion import (
|
||||||
|
fused_norm_scale_shift_fp8,
|
||||||
|
fused_scale_residual_norm_scale_shift_fp8,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8
|
||||||
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
|
LayerNormScaleShift,
|
||||||
|
ScaleResidualLayerNormScaleShift,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(
|
||||||
|
est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
|
||||||
|
)
|
||||||
|
|
||||||
|
DEVICE = "cuda"
|
||||||
|
DTYPE = torch.bfloat16
|
||||||
|
HIDDEN = 3072
|
||||||
|
EPS = 1e-6
|
||||||
|
|
||||||
|
|
||||||
|
@marker.parametrize("rows", [128, 1024, 4096], [128])
|
||||||
|
@marker.parametrize("residual_path", [False, True], [False, True])
|
||||||
|
@marker.benchmark("impl", ["split", "fused"], unit="us")
|
||||||
|
def benchmark(rows: int, residual_path: bool, impl: str):
|
||||||
|
if impl == "fused" and torch.cuda.get_device_capability()[0] < 10:
|
||||||
|
marker.skip("Fused Qwen-Image norm+FP8 quant requires NVIDIA Blackwell")
|
||||||
|
|
||||||
|
generator = torch.Generator(device=DEVICE)
|
||||||
|
generator.manual_seed(20260831 + rows + int(residual_path))
|
||||||
|
x = torch.randn((1, rows, HIDDEN), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
residual = torch.randn_like(x)
|
||||||
|
gate = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
scale = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
shift = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
input_scale = torch.tensor(0.03125, dtype=torch.float32, device=DEVICE)
|
||||||
|
|
||||||
|
if residual_path:
|
||||||
|
layer = ScaleResidualLayerNormScaleShift(
|
||||||
|
HIDDEN, eps=EPS, elementwise_affine=False, dtype=DTYPE
|
||||||
|
).to(DEVICE)
|
||||||
|
|
||||||
|
if impl == "split":
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
normalized, residual_out = layer.forward_cuda(
|
||||||
|
residual, x, gate, shift, scale
|
||||||
|
)
|
||||||
|
quantized, _ = static_quant_fp8(normalized, input_scale)
|
||||||
|
return quantized, residual_out
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
return fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual, x, gate, scale, shift, input_scale, EPS
|
||||||
|
)
|
||||||
|
|
||||||
|
else:
|
||||||
|
layer = LayerNormScaleShift(
|
||||||
|
HIDDEN, eps=EPS, elementwise_affine=False, dtype=DTYPE
|
||||||
|
).to(DEVICE)
|
||||||
|
|
||||||
|
if impl == "split":
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
normalized = layer.forward_cuda(x, shift, scale)
|
||||||
|
return static_quant_fp8(normalized, input_scale)[0]
|
||||||
|
|
||||||
|
else:
|
||||||
|
|
||||||
|
def fn():
|
||||||
|
return fused_norm_scale_shift_fp8(x, scale, shift, input_scale, EPS)
|
||||||
|
|
||||||
|
return marker.do_bench(fn, disable_log_bandwidth=True)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
benchmark.run()
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.kernels.ops.diffusion import (
|
||||||
|
fused_norm_scale_shift_fp8,
|
||||||
|
fused_scale_residual_norm_scale_shift_fp8,
|
||||||
|
)
|
||||||
|
from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8
|
||||||
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
|
LayerNormScaleShift,
|
||||||
|
ScaleResidualLayerNormScaleShift,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200")
|
||||||
|
|
||||||
|
DEVICE = "cuda"
|
||||||
|
DTYPE = torch.bfloat16
|
||||||
|
HIDDEN = 3072
|
||||||
|
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)
|
||||||
|
residual = torch.randn_like(x)
|
||||||
|
gate = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
scale = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
shift = torch.randn((HIDDEN,), dtype=DTYPE, device=DEVICE, generator=generator)
|
||||||
|
return x, residual, gate, scale, shift
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("rows", [1, 127, 1024])
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"input_scale_value", [0.005, 0.03125, 0.4263392984867096, 0.4754464328289032, 1.0]
|
||||||
|
)
|
||||||
|
def test_norm_scale_shift_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)
|
||||||
|
layer = LayerNormScaleShift(
|
||||||
|
HIDDEN, eps=EPS, elementwise_affine=False, dtype=DTYPE
|
||||||
|
).to(DEVICE)
|
||||||
|
|
||||||
|
normalized = layer.forward_cuda(x, shift, scale)
|
||||||
|
expected, _ = static_quant_fp8(normalized, input_scale)
|
||||||
|
actual_normalized, actual = fused_norm_scale_shift_fp8(
|
||||||
|
x, scale, shift, input_scale, EPS
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(actual_normalized, normalized)
|
||||||
|
assert torch.equal(actual.view(torch.uint8), expected.view(torch.uint8))
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("rows", [1, 127, 1024])
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"input_scale_value", [0.005, 0.03125, 0.4263392984867096, 0.4754464328289032, 1.0]
|
||||||
|
)
|
||||||
|
def test_residual_norm_scale_shift_fp8_is_bit_exact(
|
||||||
|
rows: int, input_scale_value: float
|
||||||
|
) -> None:
|
||||||
|
x, residual, gate, scale, shift = _make_inputs(rows)
|
||||||
|
input_scale = torch.tensor(input_scale_value, dtype=torch.float32, device=DEVICE)
|
||||||
|
layer = ScaleResidualLayerNormScaleShift(
|
||||||
|
HIDDEN, eps=EPS, elementwise_affine=False, dtype=DTYPE
|
||||||
|
).to(DEVICE)
|
||||||
|
|
||||||
|
normalized, expected_residual = layer.forward_cuda(residual, x, gate, shift, scale)
|
||||||
|
expected, _ = static_quant_fp8(normalized, input_scale)
|
||||||
|
actual_normalized, actual, actual_residual = (
|
||||||
|
fused_scale_residual_norm_scale_shift_fp8(
|
||||||
|
residual, x, gate, scale, shift, input_scale, EPS
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert torch.equal(actual_normalized, normalized)
|
||||||
|
assert torch.equal(actual.view(torch.uint8), expected.view(torch.uint8))
|
||||||
|
assert torch.equal(actual_residual, expected_residual)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|
||||||
@@ -0,0 +1,101 @@
|
|||||||
|
"""Unit tests for Qwen-Image ModelOpt FP8 norm+quant activation gates."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
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.qwen_image import (
|
||||||
|
QwenImageTransformerBlock,
|
||||||
|
)
|
||||||
|
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 _attention(*, fused: bool, scales: tuple[float, ...]) -> SimpleNamespace:
|
||||||
|
if fused:
|
||||||
|
return SimpleNamespace(
|
||||||
|
use_fused_qkv=True,
|
||||||
|
to_qkv=_fp8_linear(scales[0]),
|
||||||
|
added_kv_proj_dim=None,
|
||||||
|
)
|
||||||
|
return SimpleNamespace(
|
||||||
|
use_fused_qkv=False,
|
||||||
|
to_q=_fp8_linear(scales[0]),
|
||||||
|
to_k=_fp8_linear(scales[1]),
|
||||||
|
to_v=_fp8_linear(scales[2]),
|
||||||
|
added_kv_proj_dim=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _block(attn: SimpleNamespace) -> QwenImageTransformerBlock:
|
||||||
|
block = object.__new__(QwenImageTransformerBlock)
|
||||||
|
nn.Module.__init__(block)
|
||||||
|
block.dim = 3072
|
||||||
|
block.zero_cond_t = False
|
||||||
|
block.attn = attn
|
||||||
|
block.img_mlp = None
|
||||||
|
block.txt_mlp = None
|
||||||
|
block._fp8_img_attn_norm_quant = False
|
||||||
|
block._fp8_txt_attn_norm_quant = False
|
||||||
|
block._fp8_img_mlp_norm_quant = False
|
||||||
|
block._fp8_txt_mlp_norm_quant = False
|
||||||
|
return block
|
||||||
|
|
||||||
|
|
||||||
|
@patch("torch.cuda.get_device_capability", return_value=(10, 0))
|
||||||
|
@patch("torch.cuda.is_available", return_value=True)
|
||||||
|
class TestQwenImageFp8NormQuantGate(CustomTestCase):
|
||||||
|
def test_separate_qkv_requires_identical_input_scales(
|
||||||
|
self, _is_available, _capability
|
||||||
|
) -> None:
|
||||||
|
matching = _block(_attention(fused=False, scales=(0.25, 0.25, 0.25)))
|
||||||
|
mismatched = _block(_attention(fused=False, scales=(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.assertFalse(mismatched._fp8_img_attn_norm_quant)
|
||||||
|
|
||||||
|
def test_merged_qkv_uses_its_materialized_input_scale(
|
||||||
|
self, _is_available, _capability
|
||||||
|
) -> None:
|
||||||
|
block = _block(_attention(fused=True, scales=(0.25,)))
|
||||||
|
|
||||||
|
block.configure_fp8_norm_quant()
|
||||||
|
|
||||||
|
self.assertTrue(block._fp8_img_attn_norm_quant)
|
||||||
|
|
||||||
|
def test_nonpositive_scale_keeps_fusion_disabled(
|
||||||
|
self, _is_available, _capability
|
||||||
|
) -> None:
|
||||||
|
block = _block(_attention(fused=True, scales=(0.0,)))
|
||||||
|
|
||||||
|
block.configure_fp8_norm_quant()
|
||||||
|
|
||||||
|
self.assertFalse(block._fp8_img_attn_norm_quant)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user