[CPU] [Diffusion] Add fused scale-shift and norm kernels for CPU (#33452)
Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
co-authored by
Ma Mingfei
parent
f920be4b09
commit
2cb51f5d22
@@ -0,0 +1,662 @@
|
|||||||
|
#include "common.h"
|
||||||
|
#include "vec.h"
|
||||||
|
|
||||||
|
/*
|
||||||
|
* [Note]: Fused norm kernels for diffusion models
|
||||||
|
*
|
||||||
|
* This file contains CPU kernels for fused normalization and modulation
|
||||||
|
* operations used by diffusion models:
|
||||||
|
*
|
||||||
|
* - fused_scale_shift_cpu:
|
||||||
|
* Applies scale-shift modulation:
|
||||||
|
* output = input * (scale_constant + scale) + shift.
|
||||||
|
*
|
||||||
|
* - fused_norm_scale_shift_cpu:
|
||||||
|
* Applies RMSNorm or LayerNorm followed by scale-shift modulation.
|
||||||
|
*
|
||||||
|
* - fused_scale_residual_norm_scale_shift_cpu:
|
||||||
|
* Fuses optional gated residual accumulation, normalization, and
|
||||||
|
* scale-shift modulation.
|
||||||
|
*/
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
enum class DiffusionNormMode {
|
||||||
|
RMSNorm,
|
||||||
|
LayerNorm,
|
||||||
|
};
|
||||||
|
|
||||||
|
#define DISPATCH_DIFFUSION_NORM_TYPE(norm_type, name, ...) \
|
||||||
|
[&] { \
|
||||||
|
if ((norm_type) == "rms") { \
|
||||||
|
using norm_mode_t = std::integral_constant<DiffusionNormMode, DiffusionNormMode::RMSNorm>; \
|
||||||
|
return __VA_ARGS__(norm_mode_t{}); \
|
||||||
|
} \
|
||||||
|
TORCH_CHECK((norm_type) == "layer", name, ": norm_type must be 'rms' or 'layer', got ", (norm_type)); \
|
||||||
|
using norm_mode_t = std::integral_constant<DiffusionNormMode, DiffusionNormMode::LayerNorm>; \
|
||||||
|
return __VA_ARGS__(norm_mode_t{}); \
|
||||||
|
}()
|
||||||
|
|
||||||
|
template <DiffusionNormMode M>
|
||||||
|
struct DiffusionNormTraits;
|
||||||
|
|
||||||
|
template <>
|
||||||
|
struct DiffusionNormTraits<DiffusionNormMode::RMSNorm> {
|
||||||
|
static constexpr bool has_mean = false;
|
||||||
|
static constexpr bool has_bias = false;
|
||||||
|
};
|
||||||
|
|
||||||
|
template <>
|
||||||
|
struct DiffusionNormTraits<DiffusionNormMode::LayerNorm> {
|
||||||
|
static constexpr bool has_mean = true;
|
||||||
|
static constexpr bool has_bias = true;
|
||||||
|
};
|
||||||
|
|
||||||
|
using fVec = at::vec::Vectorized<float>;
|
||||||
|
|
||||||
|
template <typename T>
|
||||||
|
struct ModulationParam {
|
||||||
|
const T* data{nullptr};
|
||||||
|
int64_t stride_b{0};
|
||||||
|
int64_t stride_s{0};
|
||||||
|
int64_t stride_c{0};
|
||||||
|
|
||||||
|
ModulationParam() = default;
|
||||||
|
|
||||||
|
explicit ModulationParam(const at::Tensor& tensor)
|
||||||
|
: data(tensor.data_ptr<T>()),
|
||||||
|
stride_b(tensor.stride(0)),
|
||||||
|
stride_s(tensor.stride(1)),
|
||||||
|
stride_c(tensor.stride(2)) {}
|
||||||
|
|
||||||
|
inline const T* row(int64_t b, int64_t s) const {
|
||||||
|
if (data == nullptr) {
|
||||||
|
return nullptr;
|
||||||
|
}
|
||||||
|
return data + b * stride_b + s * stride_s;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
template <typename RowFn>
|
||||||
|
inline void parallel_for_rows(int64_t B, int64_t S, int64_t D, RowFn&& row_fn) {
|
||||||
|
at::parallel_for(0, B * S, 0, [&](int64_t begin, int64_t end) {
|
||||||
|
for (int64_t row = begin; row < end; ++row) {
|
||||||
|
row_fn(row / S, row % S, row * D);
|
||||||
|
}
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
template <typename T>
|
||||||
|
inline void load_param_vec2(fVec& v0, fVec& v1, const T* __restrict__ p, int64_t stride_c, int64_t d) {
|
||||||
|
if (stride_c == 0) {
|
||||||
|
v0 = v1 = fVec(static_cast<float>(p[0]));
|
||||||
|
} else {
|
||||||
|
std::tie(v0, v1) = load_float_vec2(p + d);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <typename param_t>
|
||||||
|
inline void apply_scale_shift_vec(
|
||||||
|
fVec& x0,
|
||||||
|
fVec& x1,
|
||||||
|
const param_t* __restrict__ scale,
|
||||||
|
const param_t* __restrict__ shift,
|
||||||
|
int64_t scale_stride_c,
|
||||||
|
int64_t shift_stride_c,
|
||||||
|
int64_t d,
|
||||||
|
float scale_constant = 1.0f) {
|
||||||
|
fVec scale0, scale1;
|
||||||
|
fVec shift0, shift1;
|
||||||
|
|
||||||
|
load_param_vec2(scale0, scale1, scale, scale_stride_c, d);
|
||||||
|
load_param_vec2(shift0, shift1, shift, shift_stride_c, d);
|
||||||
|
|
||||||
|
x0 = x0 * (fVec(scale_constant) + scale0) + shift0;
|
||||||
|
x1 = x1 * (fVec(scale_constant) + scale1) + shift1;
|
||||||
|
}
|
||||||
|
template <typename scalar_t>
|
||||||
|
inline void apply_residual_gate_vec(
|
||||||
|
fVec& x0,
|
||||||
|
fVec& x1,
|
||||||
|
const fVec& r0,
|
||||||
|
const fVec& r1,
|
||||||
|
const scalar_t* __restrict__ gate,
|
||||||
|
const float* __restrict__ gate_fp32,
|
||||||
|
int64_t gate_stride_c,
|
||||||
|
int64_t d) {
|
||||||
|
fVec g0, g1;
|
||||||
|
|
||||||
|
if (gate_fp32 != nullptr) {
|
||||||
|
load_param_vec2(g0, g1, gate_fp32, gate_stride_c, d);
|
||||||
|
} else if (gate != nullptr) {
|
||||||
|
load_param_vec2(g0, g1, gate, gate_stride_c, d);
|
||||||
|
} else {
|
||||||
|
g0 = g1 = fVec(1.0f);
|
||||||
|
}
|
||||||
|
|
||||||
|
x0 = r0 + x0 * g0;
|
||||||
|
x1 = r1 + x1 * g1;
|
||||||
|
}
|
||||||
|
|
||||||
|
template <DiffusionNormMode M, typename scalar_t, typename param_t>
|
||||||
|
inline void apply_norm_modulate_row(
|
||||||
|
scalar_t* __restrict__ output,
|
||||||
|
const scalar_t* __restrict__ input,
|
||||||
|
const float* __restrict__ weight,
|
||||||
|
const float* __restrict__ bias,
|
||||||
|
const param_t* __restrict__ scale,
|
||||||
|
const param_t* __restrict__ shift,
|
||||||
|
int64_t D,
|
||||||
|
int64_t scale_stride_c,
|
||||||
|
int64_t shift_stride_c,
|
||||||
|
const fVec& sum_vec,
|
||||||
|
const fVec& sum_sq_vec,
|
||||||
|
float sum,
|
||||||
|
float sum_sq,
|
||||||
|
float eps) {
|
||||||
|
sum_sq += vec_reduce_sum(sum_sq_vec);
|
||||||
|
|
||||||
|
float mean = 0.0f;
|
||||||
|
float variance = sum_sq / static_cast<float>(D);
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
sum += vec_reduce_sum(sum_vec);
|
||||||
|
mean = sum / static_cast<float>(D);
|
||||||
|
variance -= mean * mean;
|
||||||
|
}
|
||||||
|
|
||||||
|
const float rstd = 1.0f / std::sqrt(variance + eps);
|
||||||
|
|
||||||
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
|
constexpr int64_t kVecSize = bVec::size();
|
||||||
|
|
||||||
|
const fVec mean_vec(mean);
|
||||||
|
const fVec rstd_vec(rstd);
|
||||||
|
|
||||||
|
int64_t d = 0;
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d <= D - kVecSize; d += kVecSize) {
|
||||||
|
auto [x0, x1] = load_float_vec2(input + d);
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
x0 -= mean_vec;
|
||||||
|
x1 -= mean_vec;
|
||||||
|
}
|
||||||
|
|
||||||
|
x0 *= rstd_vec;
|
||||||
|
x1 *= rstd_vec;
|
||||||
|
|
||||||
|
if (weight != nullptr) {
|
||||||
|
auto [w0, w1] = load_float_vec2(weight + d);
|
||||||
|
x0 *= w0;
|
||||||
|
x1 *= w1;
|
||||||
|
}
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_bias) {
|
||||||
|
if (bias != nullptr) {
|
||||||
|
auto [b0, b1] = load_float_vec2(bias + d);
|
||||||
|
x0 += b0;
|
||||||
|
x1 += b1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match CUDA/CuTe activation-dtype boundary:
|
||||||
|
// norm FP32 -> activation dtype -> scale/shift.
|
||||||
|
const bVec norm_value = convert_from_float_ext<scalar_t>(x0, x1);
|
||||||
|
std::tie(x0, x1) = at::vec::convert_to_float(norm_value);
|
||||||
|
|
||||||
|
apply_scale_shift_vec(x0, x1, scale, shift, scale_stride_c, shift_stride_c, d);
|
||||||
|
convert_from_float_ext<scalar_t>(x0, x1).store(output + d);
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d < D; ++d) {
|
||||||
|
float x = static_cast<float>(input[d]);
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
x -= mean;
|
||||||
|
}
|
||||||
|
|
||||||
|
x *= rstd;
|
||||||
|
|
||||||
|
if (weight != nullptr) {
|
||||||
|
x *= weight[d];
|
||||||
|
}
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_bias) {
|
||||||
|
if (bias != nullptr) {
|
||||||
|
x += bias[d];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match CUDA/CuTe activation-dtype boundary.
|
||||||
|
x = static_cast<float>(static_cast<scalar_t>(x));
|
||||||
|
|
||||||
|
x = x * (1.0f + static_cast<float>(scale[d * scale_stride_c])) + static_cast<float>(shift[d * shift_stride_c]);
|
||||||
|
output[d] = static_cast<scalar_t>(x);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
template <typename scalar_t, typename param_t>
|
||||||
|
inline void fused_scale_shift_row(
|
||||||
|
scalar_t* __restrict__ output,
|
||||||
|
const scalar_t* __restrict__ input,
|
||||||
|
const param_t* __restrict__ scale,
|
||||||
|
const param_t* __restrict__ shift,
|
||||||
|
int64_t D,
|
||||||
|
int64_t scale_stride_c,
|
||||||
|
int64_t shift_stride_c,
|
||||||
|
float scale_constant) {
|
||||||
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
|
constexpr int64_t kVecSize = bVec::size();
|
||||||
|
int64_t d = 0;
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d <= D - kVecSize; d += kVecSize) {
|
||||||
|
auto [x0, x1] = load_float_vec2(input + d);
|
||||||
|
apply_scale_shift_vec(x0, x1, scale, shift, scale_stride_c, shift_stride_c, d, scale_constant);
|
||||||
|
convert_from_float_ext<scalar_t>(x0, x1).store(output + d);
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d < D; ++d) {
|
||||||
|
const float x = static_cast<float>(input[d]);
|
||||||
|
const float scale_value = static_cast<float>(scale[d * scale_stride_c]);
|
||||||
|
const float shift_value = static_cast<float>(shift[d * shift_stride_c]);
|
||||||
|
output[d] = static_cast<scalar_t>(x * (scale_constant + scale_value) + shift_value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <DiffusionNormMode M, typename scalar_t, typename param_t>
|
||||||
|
inline void fused_norm_scale_shift_row(
|
||||||
|
scalar_t* __restrict__ output,
|
||||||
|
const scalar_t* __restrict__ input,
|
||||||
|
const float* __restrict__ weight,
|
||||||
|
const float* __restrict__ bias,
|
||||||
|
const param_t* __restrict__ scale,
|
||||||
|
const param_t* __restrict__ shift,
|
||||||
|
int64_t D,
|
||||||
|
int64_t scale_stride_c,
|
||||||
|
int64_t shift_stride_c,
|
||||||
|
float eps) {
|
||||||
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
|
constexpr int64_t kVecSize = bVec::size();
|
||||||
|
|
||||||
|
fVec sum_vec{0.0f};
|
||||||
|
fVec sum_sq_vec{0.0f};
|
||||||
|
float sum = 0.0f;
|
||||||
|
float sum_sq = 0.0f;
|
||||||
|
|
||||||
|
int64_t d = 0;
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d <= D - kVecSize; d += kVecSize) {
|
||||||
|
auto [x0, x1] = load_float_vec2(input + d);
|
||||||
|
sum_sq_vec += x0 * x0 + x1 * x1;
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
sum_vec += x0 + x1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d < D; ++d) {
|
||||||
|
const float x = static_cast<float>(input[d]);
|
||||||
|
sum_sq += x * x;
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
sum += x;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
apply_norm_modulate_row<M>(
|
||||||
|
output,
|
||||||
|
input,
|
||||||
|
weight,
|
||||||
|
bias,
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
D,
|
||||||
|
scale_stride_c,
|
||||||
|
shift_stride_c,
|
||||||
|
sum_vec,
|
||||||
|
sum_sq_vec,
|
||||||
|
sum,
|
||||||
|
sum_sq,
|
||||||
|
eps);
|
||||||
|
}
|
||||||
|
|
||||||
|
template <DiffusionNormMode M, typename scalar_t, typename param_t>
|
||||||
|
inline void fused_scale_residual_norm_scale_shift_row(
|
||||||
|
scalar_t* __restrict__ output,
|
||||||
|
scalar_t* __restrict__ residual_output,
|
||||||
|
const scalar_t* __restrict__ residual,
|
||||||
|
const scalar_t* __restrict__ input,
|
||||||
|
const scalar_t* __restrict__ residual_gate,
|
||||||
|
const float* __restrict__ residual_gate_fp32,
|
||||||
|
const float* __restrict__ weight,
|
||||||
|
const float* __restrict__ bias,
|
||||||
|
const param_t* __restrict__ scale,
|
||||||
|
const param_t* __restrict__ shift,
|
||||||
|
int64_t D,
|
||||||
|
int64_t gate_stride_c,
|
||||||
|
int64_t scale_stride_c,
|
||||||
|
int64_t shift_stride_c,
|
||||||
|
float eps) {
|
||||||
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
|
constexpr int64_t kVecSize = bVec::size();
|
||||||
|
|
||||||
|
fVec sum_vec{0.0f};
|
||||||
|
fVec sum_sq_vec{0.0f};
|
||||||
|
float sum = 0.0f;
|
||||||
|
float sum_sq = 0.0f;
|
||||||
|
|
||||||
|
int64_t d = 0;
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d <= D - kVecSize; d += kVecSize) {
|
||||||
|
auto [x0, x1] = load_float_vec2(input + d);
|
||||||
|
auto [r0, r1] = load_float_vec2(residual + d);
|
||||||
|
|
||||||
|
apply_residual_gate_vec(x0, x1, r0, r1, residual_gate, residual_gate_fp32, gate_stride_c, d);
|
||||||
|
|
||||||
|
// Match CUDA: residual + gate * input is rounded to activation dtype
|
||||||
|
// before normalization.
|
||||||
|
const bVec residual_value = convert_from_float_ext<scalar_t>(x0, x1);
|
||||||
|
|
||||||
|
residual_value.store(residual_output + d);
|
||||||
|
|
||||||
|
std::tie(x0, x1) = at::vec::convert_to_float(residual_value);
|
||||||
|
|
||||||
|
sum_sq_vec += x0 * x0 + x1 * x1;
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
sum_vec += x0 + x1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#pragma GCC unroll 4
|
||||||
|
for (; d < D; ++d) {
|
||||||
|
float x = static_cast<float>(input[d]);
|
||||||
|
if (residual_gate_fp32 != nullptr) {
|
||||||
|
x *= residual_gate_fp32[d * gate_stride_c];
|
||||||
|
} else if (residual_gate != nullptr) {
|
||||||
|
x *= static_cast<float>(residual_gate[d * gate_stride_c]);
|
||||||
|
}
|
||||||
|
|
||||||
|
x += static_cast<float>(residual[d]);
|
||||||
|
|
||||||
|
const scalar_t residual_value = static_cast<scalar_t>(x);
|
||||||
|
|
||||||
|
residual_output[d] = residual_value;
|
||||||
|
|
||||||
|
x = static_cast<float>(residual_value);
|
||||||
|
|
||||||
|
sum_sq += x * x;
|
||||||
|
|
||||||
|
if constexpr (DiffusionNormTraits<M>::has_mean) {
|
||||||
|
sum += x;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
apply_norm_modulate_row<M>(
|
||||||
|
output,
|
||||||
|
residual_output,
|
||||||
|
weight,
|
||||||
|
bias,
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
D,
|
||||||
|
scale_stride_c,
|
||||||
|
shift_stride_c,
|
||||||
|
sum_vec,
|
||||||
|
sum_sq_vec,
|
||||||
|
sum,
|
||||||
|
sum_sq,
|
||||||
|
eps);
|
||||||
|
}
|
||||||
|
|
||||||
|
inline void check_modulation_param(const at::Tensor& param, const at::Tensor& input, const char* name) {
|
||||||
|
CHECK_CPU(param);
|
||||||
|
CHECK_DIM(3, param);
|
||||||
|
CHECK_EQ(param.sizes(), input.sizes());
|
||||||
|
TORCH_CHECK(param.stride(2) == 0 || param.stride(2) == 1, name, " hidden-dimension stride must be 0 or 1.");
|
||||||
|
}
|
||||||
|
|
||||||
|
inline const float* get_norm_param_ptr(const std::optional<at::Tensor>& param, int64_t D, const char* name) {
|
||||||
|
if (!param.has_value()) {
|
||||||
|
return nullptr;
|
||||||
|
}
|
||||||
|
|
||||||
|
const auto& tensor = param.value();
|
||||||
|
|
||||||
|
CHECK_INPUT(tensor);
|
||||||
|
CHECK_DIM(1, tensor);
|
||||||
|
CHECK_EQ(tensor.size(0), D);
|
||||||
|
|
||||||
|
TORCH_CHECK(
|
||||||
|
tensor.scalar_type() == at::ScalarType::Float, "CPU fused diffusion norm only supports FP32 norm ", name, ".");
|
||||||
|
return tensor.data_ptr<float>();
|
||||||
|
}
|
||||||
|
} // anonymous namespace
|
||||||
|
at::Tensor fused_scale_shift_cpu(
|
||||||
|
const at::Tensor& input, const at::Tensor& scale, const at::Tensor& shift, double scale_constant) {
|
||||||
|
CHECK_LAST_DIM_CONTIGUOUS_INPUT(input);
|
||||||
|
CHECK_DIM(3, input);
|
||||||
|
|
||||||
|
check_modulation_param(scale, input, "scale");
|
||||||
|
check_modulation_param(shift, input, "shift");
|
||||||
|
|
||||||
|
CHECK_EQ(scale.scalar_type(), shift.scalar_type());
|
||||||
|
|
||||||
|
const int64_t B = input.size(0);
|
||||||
|
const int64_t S = input.size(1);
|
||||||
|
const int64_t D = input.size(2);
|
||||||
|
|
||||||
|
// Output is contiguous even if input is only last-dim contiguous.
|
||||||
|
at::Tensor output = at::empty(input.sizes(), input.options());
|
||||||
|
|
||||||
|
if (input.numel() == 0) {
|
||||||
|
return output;
|
||||||
|
}
|
||||||
|
|
||||||
|
CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(input.scalar_type(), scale.scalar_type(), "fused_scale_shift_cpu", [&] {
|
||||||
|
const ModulationParam<param_t> scale_param(scale);
|
||||||
|
const ModulationParam<param_t> shift_param(shift);
|
||||||
|
|
||||||
|
const scalar_t* input_ptr = input.data_ptr<scalar_t>();
|
||||||
|
scalar_t* output_ptr = output.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
const int64_t input_stride_b = input.stride(0);
|
||||||
|
const int64_t input_stride_s = input.stride(1);
|
||||||
|
|
||||||
|
parallel_for_rows(B, S, D, [&](int64_t b, int64_t s, int64_t offset) {
|
||||||
|
const scalar_t* input_row = input_ptr + b * input_stride_b + s * input_stride_s;
|
||||||
|
fused_scale_shift_row<scalar_t, param_t>(
|
||||||
|
output_ptr + offset,
|
||||||
|
input_row,
|
||||||
|
scale_param.row(b, s),
|
||||||
|
shift_param.row(b, s),
|
||||||
|
D,
|
||||||
|
scale_param.stride_c,
|
||||||
|
shift_param.stride_c,
|
||||||
|
static_cast<float>(scale_constant));
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
return output;
|
||||||
|
}
|
||||||
|
at::Tensor fused_norm_scale_shift_cpu(
|
||||||
|
const at::Tensor& input,
|
||||||
|
const std::optional<at::Tensor>& weight,
|
||||||
|
const std::optional<at::Tensor>& bias,
|
||||||
|
const at::Tensor& scale,
|
||||||
|
const at::Tensor& shift,
|
||||||
|
const std::string& norm_type,
|
||||||
|
double eps) {
|
||||||
|
CHECK_LAST_DIM_CONTIGUOUS_INPUT(input);
|
||||||
|
CHECK_DIM(3, input);
|
||||||
|
|
||||||
|
check_modulation_param(scale, input, "scale");
|
||||||
|
check_modulation_param(shift, input, "shift");
|
||||||
|
|
||||||
|
CHECK_EQ(scale.scalar_type(), shift.scalar_type());
|
||||||
|
|
||||||
|
const int64_t B = input.size(0);
|
||||||
|
const int64_t S = input.size(1);
|
||||||
|
const int64_t D = input.size(2);
|
||||||
|
|
||||||
|
const float* weight_ptr = get_norm_param_ptr(weight, D, "weight");
|
||||||
|
const float* bias_ptr = get_norm_param_ptr(bias, D, "bias");
|
||||||
|
|
||||||
|
at::Tensor output = at::empty(input.sizes(), input.options());
|
||||||
|
|
||||||
|
if (input.numel() == 0) {
|
||||||
|
return output;
|
||||||
|
}
|
||||||
|
|
||||||
|
CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(input.scalar_type(), scale.scalar_type(), "fused_norm_scale_shift_cpu", [&] {
|
||||||
|
const ModulationParam<param_t> scale_param(scale);
|
||||||
|
const ModulationParam<param_t> shift_param(shift);
|
||||||
|
|
||||||
|
const scalar_t* input_ptr = input.data_ptr<scalar_t>();
|
||||||
|
scalar_t* output_ptr = output.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
const int64_t input_stride_b = input.stride(0);
|
||||||
|
const int64_t input_stride_s = input.stride(1);
|
||||||
|
|
||||||
|
DISPATCH_DIFFUSION_NORM_TYPE(norm_type, "fused_norm_scale_shift_cpu", [&](auto mode_tag) {
|
||||||
|
constexpr DiffusionNormMode M = decltype(mode_tag)::value;
|
||||||
|
|
||||||
|
if constexpr (!DiffusionNormTraits<M>::has_bias) {
|
||||||
|
TORCH_CHECK(!bias.has_value(), "bias is only supported for LayerNorm.");
|
||||||
|
}
|
||||||
|
|
||||||
|
parallel_for_rows(B, S, D, [&](int64_t b, int64_t s, int64_t offset) {
|
||||||
|
const scalar_t* input_row = input_ptr + b * input_stride_b + s * input_stride_s;
|
||||||
|
|
||||||
|
fused_norm_scale_shift_row<M, scalar_t, param_t>(
|
||||||
|
output_ptr + offset,
|
||||||
|
input_row,
|
||||||
|
weight_ptr,
|
||||||
|
bias_ptr,
|
||||||
|
scale_param.row(b, s),
|
||||||
|
shift_param.row(b, s),
|
||||||
|
D,
|
||||||
|
scale_param.stride_c,
|
||||||
|
shift_param.stride_c,
|
||||||
|
static_cast<float>(eps));
|
||||||
|
});
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
return output;
|
||||||
|
}
|
||||||
|
std::tuple<at::Tensor, at::Tensor> fused_scale_residual_norm_scale_shift_cpu(
|
||||||
|
const at::Tensor& residual,
|
||||||
|
const at::Tensor& input,
|
||||||
|
const std::optional<at::Tensor>& residual_gate,
|
||||||
|
const std::optional<at::Tensor>& weight,
|
||||||
|
const std::optional<at::Tensor>& bias,
|
||||||
|
const at::Tensor& scale,
|
||||||
|
const at::Tensor& shift,
|
||||||
|
const std::string& norm_type,
|
||||||
|
double eps) {
|
||||||
|
CHECK_LAST_DIM_CONTIGUOUS_INPUT(input);
|
||||||
|
CHECK_DIM(3, input);
|
||||||
|
|
||||||
|
CHECK_LAST_DIM_CONTIGUOUS_INPUT(residual);
|
||||||
|
CHECK_DIM(3, residual);
|
||||||
|
|
||||||
|
CHECK_EQ(residual.sizes(), input.sizes());
|
||||||
|
CHECK_EQ(residual.scalar_type(), input.scalar_type());
|
||||||
|
|
||||||
|
check_modulation_param(scale, input, "scale");
|
||||||
|
check_modulation_param(shift, input, "shift");
|
||||||
|
|
||||||
|
CHECK_EQ(scale.scalar_type(), shift.scalar_type());
|
||||||
|
|
||||||
|
if (residual_gate.has_value()) {
|
||||||
|
check_modulation_param(residual_gate.value(), input, "residual_gate");
|
||||||
|
|
||||||
|
TORCH_CHECK(
|
||||||
|
residual_gate->scalar_type() == input.scalar_type() || residual_gate->scalar_type() == at::ScalarType::Float,
|
||||||
|
"residual_gate must have the same dtype as "
|
||||||
|
"input or be FP32.");
|
||||||
|
}
|
||||||
|
|
||||||
|
const int64_t B = input.size(0);
|
||||||
|
const int64_t S = input.size(1);
|
||||||
|
const int64_t D = input.size(2);
|
||||||
|
|
||||||
|
const float* weight_ptr = get_norm_param_ptr(weight, D, "weight");
|
||||||
|
const float* bias_ptr = get_norm_param_ptr(bias, D, "bias");
|
||||||
|
|
||||||
|
at::Tensor output = at::empty(input.sizes(), input.options());
|
||||||
|
|
||||||
|
at::Tensor residual_output = at::empty(input.sizes(), input.options());
|
||||||
|
|
||||||
|
if (input.numel() == 0) {
|
||||||
|
return {output, residual_output};
|
||||||
|
}
|
||||||
|
|
||||||
|
CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(
|
||||||
|
input.scalar_type(), scale.scalar_type(), "fused_scale_residual_norm_scale_shift_cpu", [&] {
|
||||||
|
const ModulationParam<param_t> scale_param(scale);
|
||||||
|
const ModulationParam<param_t> shift_param(shift);
|
||||||
|
|
||||||
|
ModulationParam<scalar_t> gate_param{};
|
||||||
|
ModulationParam<float> gate_fp32_param{};
|
||||||
|
|
||||||
|
if (residual_gate.has_value()) {
|
||||||
|
if (residual_gate->scalar_type() == at::ScalarType::Float) {
|
||||||
|
gate_fp32_param = ModulationParam<float>(residual_gate.value());
|
||||||
|
} else {
|
||||||
|
gate_param = ModulationParam<scalar_t>(residual_gate.value());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
const scalar_t* input_ptr = input.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
const scalar_t* residual_ptr = residual.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
scalar_t* output_ptr = output.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
scalar_t* residual_output_ptr = residual_output.data_ptr<scalar_t>();
|
||||||
|
|
||||||
|
const int64_t input_stride_b = input.stride(0);
|
||||||
|
const int64_t input_stride_s = input.stride(1);
|
||||||
|
|
||||||
|
const int64_t residual_stride_b = residual.stride(0);
|
||||||
|
const int64_t residual_stride_s = residual.stride(1);
|
||||||
|
const int64_t gate_stride_c = residual_gate.has_value() ? residual_gate->stride(2) : 0;
|
||||||
|
DISPATCH_DIFFUSION_NORM_TYPE(norm_type, "fused_scale_residual_norm_scale_shift_cpu", [&](auto mode_tag) {
|
||||||
|
constexpr DiffusionNormMode M = decltype(mode_tag)::value;
|
||||||
|
if constexpr (!DiffusionNormTraits<M>::has_bias) {
|
||||||
|
TORCH_CHECK(!bias.has_value(), "bias is only supported for LayerNorm.");
|
||||||
|
}
|
||||||
|
parallel_for_rows(B, S, D, [&](int64_t b, int64_t s, int64_t offset) {
|
||||||
|
const scalar_t* input_row = input_ptr + b * input_stride_b + s * input_stride_s;
|
||||||
|
|
||||||
|
const scalar_t* residual_row = residual_ptr + b * residual_stride_b + s * residual_stride_s;
|
||||||
|
|
||||||
|
fused_scale_residual_norm_scale_shift_row<M, scalar_t, param_t>(
|
||||||
|
output_ptr + offset,
|
||||||
|
residual_output_ptr + offset,
|
||||||
|
residual_row,
|
||||||
|
input_row,
|
||||||
|
gate_param.row(b, s),
|
||||||
|
gate_fp32_param.row(b, s),
|
||||||
|
weight_ptr,
|
||||||
|
bias_ptr,
|
||||||
|
scale_param.row(b, s),
|
||||||
|
shift_param.row(b, s),
|
||||||
|
D,
|
||||||
|
gate_stride_c,
|
||||||
|
scale_param.stride_c,
|
||||||
|
shift_param.stride_c,
|
||||||
|
static_cast<float>(eps));
|
||||||
|
});
|
||||||
|
});
|
||||||
|
});
|
||||||
|
|
||||||
|
return {output, residual_output};
|
||||||
|
}
|
||||||
|
#undef DISPATCH_DIFFUSION_NORM_TYPE
|
||||||
@@ -43,6 +43,32 @@ at::Tensor gemma4_rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps,
|
|||||||
at::Tensor
|
at::Tensor
|
||||||
layernorm_cpu(const at::Tensor& input, const at::Tensor& weight, const std::optional<at::Tensor>& bias, double eps);
|
layernorm_cpu(const at::Tensor& input, const at::Tensor& weight, const std::optional<at::Tensor>& bias, double eps);
|
||||||
|
|
||||||
|
// fused_scale_shift
|
||||||
|
at::Tensor
|
||||||
|
fused_scale_shift_cpu(const at::Tensor& input, const at::Tensor& scale, const at::Tensor& shift, double scale_constant);
|
||||||
|
|
||||||
|
// fused_norm_scale_shift
|
||||||
|
at::Tensor fused_norm_scale_shift_cpu(
|
||||||
|
const at::Tensor& input,
|
||||||
|
const std::optional<at::Tensor>& weight,
|
||||||
|
const std::optional<at::Tensor>& bias,
|
||||||
|
const at::Tensor& scale,
|
||||||
|
const at::Tensor& shift,
|
||||||
|
const std::string& norm_type,
|
||||||
|
double eps);
|
||||||
|
|
||||||
|
// fused_scale_residual_norm_scale_shift
|
||||||
|
std::tuple<at::Tensor, at::Tensor> fused_scale_residual_norm_scale_shift_cpu(
|
||||||
|
const at::Tensor& residual,
|
||||||
|
const at::Tensor& input,
|
||||||
|
const std::optional<at::Tensor>& gate,
|
||||||
|
const std::optional<at::Tensor>& weight,
|
||||||
|
const std::optional<at::Tensor>& bias,
|
||||||
|
const at::Tensor& scale,
|
||||||
|
const at::Tensor& shift,
|
||||||
|
const std::string& norm_type,
|
||||||
|
double eps);
|
||||||
|
|
||||||
// qwen3_next_rmsnorm_gated
|
// qwen3_next_rmsnorm_gated
|
||||||
at::Tensor fused_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps);
|
at::Tensor fused_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps);
|
||||||
|
|
||||||
@@ -651,6 +677,34 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
|||||||
"head_dim, int num_head) -> "
|
"head_dim, int num_head) -> "
|
||||||
"(Tensor, Tensor, Tensor)");
|
"(Tensor, Tensor, Tensor)");
|
||||||
m.impl("fused_qk_gemma_rmsnorm_with_gate_cpu", torch::kCPU, &fused_qk_gemma_rmsnorm_with_gate_cpu);
|
m.impl("fused_qk_gemma_rmsnorm_with_gate_cpu", torch::kCPU, &fused_qk_gemma_rmsnorm_with_gate_cpu);
|
||||||
|
m.def("fused_scale_shift_cpu(Tensor input, Tensor scale, Tensor shift, float scale_constant) -> Tensor");
|
||||||
|
m.impl("fused_scale_shift_cpu", torch::kCPU, &fused_scale_shift_cpu);
|
||||||
|
m.def(
|
||||||
|
"fused_norm_scale_shift_cpu("
|
||||||
|
"Tensor input, "
|
||||||
|
"Tensor? weight, "
|
||||||
|
"Tensor? bias, "
|
||||||
|
"Tensor scale, "
|
||||||
|
"Tensor shift, "
|
||||||
|
"str norm_type, "
|
||||||
|
"float eps"
|
||||||
|
") -> Tensor");
|
||||||
|
|
||||||
|
m.impl("fused_norm_scale_shift_cpu", torch::kCPU, &fused_norm_scale_shift_cpu);
|
||||||
|
m.def(
|
||||||
|
"fused_scale_residual_norm_scale_shift_cpu("
|
||||||
|
"Tensor residual, "
|
||||||
|
"Tensor input, "
|
||||||
|
"Tensor? gate, "
|
||||||
|
"Tensor? weight, "
|
||||||
|
"Tensor? bias, "
|
||||||
|
"Tensor scale, "
|
||||||
|
"Tensor shift, "
|
||||||
|
"str norm_type, "
|
||||||
|
"float eps"
|
||||||
|
") -> (Tensor, Tensor)");
|
||||||
|
|
||||||
|
m.impl("fused_scale_residual_norm_scale_shift_cpu", torch::kCPU, &fused_scale_residual_norm_scale_shift_cpu);
|
||||||
|
|
||||||
// speculative decoding
|
// speculative decoding
|
||||||
m.def(
|
m.def(
|
||||||
|
|||||||
@@ -526,6 +526,66 @@ def fuse_scale_shift_kernel(
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def expand_scale_shift_cpu_param(
|
||||||
|
tensor: torch.Tensor,
|
||||||
|
x: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
B, L, C = x.shape
|
||||||
|
|
||||||
|
if tensor.numel() == 1:
|
||||||
|
return tensor.reshape(1, 1, 1).expand(B, L, C)
|
||||||
|
|
||||||
|
if tensor.dim() == 1:
|
||||||
|
if tensor.shape[0] != C:
|
||||||
|
raise ValueError(f"1D modulation tensor must have shape [{C}]")
|
||||||
|
tensor = tensor.reshape(1, 1, C)
|
||||||
|
|
||||||
|
elif tensor.dim() == 2:
|
||||||
|
tensor = tensor[:, None, :]
|
||||||
|
|
||||||
|
elif tensor.dim() == 3:
|
||||||
|
pass
|
||||||
|
|
||||||
|
elif tensor.dim() == 4:
|
||||||
|
# [B, F, 1, C] -> [B, L, C]
|
||||||
|
if tensor.shape[2] != 1:
|
||||||
|
raise ValueError("4D modulation tensor must have shape [B, F, 1, C]")
|
||||||
|
num_frames = tensor.shape[1]
|
||||||
|
if L % num_frames != 0:
|
||||||
|
raise ValueError("sequence length must be divisible by num_frames")
|
||||||
|
frame_seqlen = L // num_frames
|
||||||
|
tensor = tensor.expand(
|
||||||
|
tensor.shape[0], num_frames, frame_seqlen, tensor.shape[-1]
|
||||||
|
).reshape(tensor.shape[0], L, tensor.shape[-1])
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise ValueError("modulation tensor must be scalar or 1D/2D/3D/4D")
|
||||||
|
return tensor.expand(B, L, C)
|
||||||
|
|
||||||
|
|
||||||
|
def _fuse_scale_shift_kernel_cpu(
|
||||||
|
x: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
shift: torch.Tensor,
|
||||||
|
scale_constant: float = 1.0,
|
||||||
|
block_l: int = 128,
|
||||||
|
block_c: int = 128,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
import sgl_kernel # noqa: F401
|
||||||
|
|
||||||
|
del block_l, block_c
|
||||||
|
|
||||||
|
scale = expand_scale_shift_cpu_param(scale, x)
|
||||||
|
shift = expand_scale_shift_cpu_param(shift, x)
|
||||||
|
|
||||||
|
return torch.ops.sgl_kernel.fused_scale_shift_cpu(
|
||||||
|
x,
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
scale_constant,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def fuse_layernorm_scale_shift_gate_select01_kernel(
|
def fuse_layernorm_scale_shift_gate_select01_kernel(
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
weight: torch.Tensor | None,
|
weight: torch.Tensor | None,
|
||||||
@@ -733,5 +793,5 @@ fuse_scale_shift_kernel = select_impl(
|
|||||||
npu=lazy_fallback("npu", "fuse_scale_shift_native"),
|
npu=lazy_fallback("npu", "fuse_scale_shift_native"),
|
||||||
mps=lazy_fallback("torch", "fuse_scale_shift_kernel_native"),
|
mps=lazy_fallback("torch", "fuse_scale_shift_kernel_native"),
|
||||||
musa=lazy_fallback("torch", "fuse_scale_shift_kernel_native"),
|
musa=lazy_fallback("torch", "fuse_scale_shift_kernel_native"),
|
||||||
cpu=lazy_fallback("torch", "fuse_scale_shift_kernel_native"),
|
cpu=_fuse_scale_shift_kernel_cpu,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -46,8 +46,9 @@ class CustomOp(nn.Module):
|
|||||||
return self.forward_cuda(*args, **kwargs)
|
return self.forward_cuda(*args, **kwargs)
|
||||||
|
|
||||||
def forward_cpu(self, *args, **kwargs) -> Any:
|
def forward_cpu(self, *args, **kwargs) -> Any:
|
||||||
# By default, we assume that CPU ops are compatible with CUDA ops.
|
# By default, we assume that CPU ops are compatible with the
|
||||||
return self.forward_cuda(*args, **kwargs)
|
# PyTorch-native implementation.
|
||||||
|
return self.forward_native(*args, **kwargs)
|
||||||
|
|
||||||
def forward_tpu(self, *args, **kwargs) -> Any:
|
def forward_tpu(self, *args, **kwargs) -> Any:
|
||||||
# By default, we assume that TPU ops are compatible with the
|
# By default, we assume that TPU ops are compatible with the
|
||||||
@@ -79,6 +80,8 @@ class CustomOp(nn.Module):
|
|||||||
return self.forward_xpu
|
return self.forward_xpu
|
||||||
elif current_platform.is_musa():
|
elif current_platform.is_musa():
|
||||||
return self.forward_musa
|
return self.forward_musa
|
||||||
|
elif current_platform.is_cpu():
|
||||||
|
return self.forward_cpu
|
||||||
else:
|
else:
|
||||||
return self.forward_native
|
return self.forward_native
|
||||||
|
|
||||||
|
|||||||
@@ -17,6 +17,9 @@ from sglang.kernels.ops.diffusion import (
|
|||||||
fused_inplace_qknorm_rope,
|
fused_inplace_qknorm_rope,
|
||||||
triton_one_pass_rms_norm,
|
triton_one_pass_rms_norm,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.diffusion.modulate.scale_shift_triton import (
|
||||||
|
expand_scale_shift_cpu_param,
|
||||||
|
)
|
||||||
from sglang.kernels.ops.layernorm.norm import (
|
from sglang.kernels.ops.layernorm.norm import (
|
||||||
can_use_fused_inplace_qknorm,
|
can_use_fused_inplace_qknorm,
|
||||||
fused_inplace_qknorm,
|
fused_inplace_qknorm,
|
||||||
@@ -746,6 +749,42 @@ class _ScaleResidualNormScaleShift(CustomOp):
|
|||||||
modulated = normalized * (1 + scale) + shift
|
modulated = normalized * (1 + scale) + shift
|
||||||
return modulated, residual_output
|
return modulated, residual_output
|
||||||
|
|
||||||
|
def forward_cpu(
|
||||||
|
self,
|
||||||
|
residual: torch.Tensor,
|
||||||
|
x: torch.Tensor,
|
||||||
|
gate: torch.Tensor | int,
|
||||||
|
shift: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
weight = getattr(self.norm, "weight", None)
|
||||||
|
bias = getattr(self.norm, "bias", None)
|
||||||
|
|
||||||
|
if isinstance(gate, torch.Tensor):
|
||||||
|
gate_tensor = gate
|
||||||
|
elif gate == 1:
|
||||||
|
gate_tensor = None
|
||||||
|
else:
|
||||||
|
return self.forward_native(residual, x, gate, shift, scale)
|
||||||
|
|
||||||
|
scale = expand_scale_shift_cpu_param(scale, x)
|
||||||
|
shift = expand_scale_shift_cpu_param(shift, x)
|
||||||
|
|
||||||
|
if gate_tensor is not None:
|
||||||
|
gate_tensor = expand_scale_shift_cpu_param(gate_tensor, x)
|
||||||
|
|
||||||
|
return torch.ops.sgl_kernel.fused_scale_residual_norm_scale_shift_cpu(
|
||||||
|
residual,
|
||||||
|
x,
|
||||||
|
gate_tensor,
|
||||||
|
_ensure_contiguous(weight),
|
||||||
|
_ensure_contiguous(bias),
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
self.norm_type,
|
||||||
|
self.eps,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ScaleResidualLayerNormScaleShift(_ScaleResidualNormScaleShift):
|
class ScaleResidualLayerNormScaleShift(_ScaleResidualNormScaleShift):
|
||||||
norm_type = "layer"
|
norm_type = "layer"
|
||||||
@@ -867,6 +906,28 @@ class _NormScaleShift(CustomOp):
|
|||||||
|
|
||||||
return (normalized * (1 + scale) + shift).to(x.dtype)
|
return (normalized * (1 + scale) + shift).to(x.dtype)
|
||||||
|
|
||||||
|
def forward_cpu(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
shift: torch.Tensor,
|
||||||
|
scale: torch.Tensor,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
weight = getattr(self.norm, "weight", None)
|
||||||
|
bias = getattr(self.norm, "bias", None)
|
||||||
|
|
||||||
|
scale = expand_scale_shift_cpu_param(scale, x)
|
||||||
|
shift = expand_scale_shift_cpu_param(shift, x)
|
||||||
|
|
||||||
|
return torch.ops.sgl_kernel.fused_norm_scale_shift_cpu(
|
||||||
|
x,
|
||||||
|
_ensure_contiguous(weight),
|
||||||
|
_ensure_contiguous(bias),
|
||||||
|
scale,
|
||||||
|
shift,
|
||||||
|
self.norm_type,
|
||||||
|
self.eps,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class LayerNormScaleShift(_NormScaleShift):
|
class LayerNormScaleShift(_NormScaleShift):
|
||||||
norm_type = "layer"
|
norm_type = "layer"
|
||||||
|
|||||||
@@ -0,0 +1,184 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import sgl_kernel # noqa: F401
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.cpu_test_utils import precision
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="stage-a-test-cpu-intel")
|
||||||
|
register_cpu_ci(est_time=10, suite="base-b-test-cpu-arm64")
|
||||||
|
|
||||||
|
torch.manual_seed(1234)
|
||||||
|
|
||||||
|
eps = 1e-6
|
||||||
|
|
||||||
|
DTYPE_PAIRS = [
|
||||||
|
(torch.bfloat16, torch.bfloat16),
|
||||||
|
(torch.bfloat16, torch.float32),
|
||||||
|
(torch.float16, torch.float16),
|
||||||
|
(torch.float16, torch.float32),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestDiffusionNorm:
|
||||||
|
def rmsnorm_ref(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
weight: torch.Tensor | None,
|
||||||
|
eps: float,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
x_fp32 = x.float()
|
||||||
|
variance = x_fp32.square().mean(dim=-1, keepdim=True)
|
||||||
|
out = x_fp32 * torch.rsqrt(variance + eps)
|
||||||
|
|
||||||
|
if weight is not None:
|
||||||
|
out = out * weight.float()
|
||||||
|
|
||||||
|
return out
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("input_dtype,param_dtype", DTYPE_PAIRS)
|
||||||
|
@pytest.mark.parametrize("broadcast_c", [False, True])
|
||||||
|
def test_fused_scale_shift(
|
||||||
|
self,
|
||||||
|
input_dtype,
|
||||||
|
param_dtype,
|
||||||
|
broadcast_c,
|
||||||
|
):
|
||||||
|
B, S, D = 2, 4, 67
|
||||||
|
x = torch.randn(B, S, D, dtype=input_dtype)
|
||||||
|
|
||||||
|
if broadcast_c:
|
||||||
|
# hidden dimension broadcast -> stride_c == 0
|
||||||
|
scale = torch.randn(B, 1, 1, dtype=param_dtype)
|
||||||
|
shift = torch.randn(B, S, 1, dtype=param_dtype)
|
||||||
|
else:
|
||||||
|
# normal vector load -> stride_c == 1
|
||||||
|
scale = torch.randn(B, 1, D, dtype=param_dtype)
|
||||||
|
shift = torch.randn(B, S, D, dtype=param_dtype)
|
||||||
|
|
||||||
|
scale_expanded = scale.expand_as(x)
|
||||||
|
shift_expanded = shift.expand_as(x)
|
||||||
|
|
||||||
|
if broadcast_c:
|
||||||
|
assert scale_expanded.stride(2) == 0
|
||||||
|
assert shift_expanded.stride(2) == 0
|
||||||
|
else:
|
||||||
|
assert scale_expanded.stride(2) == 1
|
||||||
|
assert shift_expanded.stride(2) == 1
|
||||||
|
|
||||||
|
out = torch.ops.sgl_kernel.fused_scale_shift_cpu(
|
||||||
|
x,
|
||||||
|
scale_expanded,
|
||||||
|
shift_expanded,
|
||||||
|
1.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
ref = (x.float() * (1.0 + scale.float()) + shift.float()).to(input_dtype)
|
||||||
|
|
||||||
|
torch.testing.assert_close(
|
||||||
|
out,
|
||||||
|
ref,
|
||||||
|
atol=precision[input_dtype],
|
||||||
|
rtol=precision[input_dtype],
|
||||||
|
)
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("input_dtype", [torch.bfloat16, torch.float16])
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"gate_type,norm_dtype,param_type,norm_type",
|
||||||
|
[
|
||||||
|
("input", None, "input", "rms"),
|
||||||
|
("fp32", torch.float32, "input", "layer"),
|
||||||
|
(None, None, "fp32", "layer"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_fused_scale_residual_norm_scale_shift(
|
||||||
|
self,
|
||||||
|
input_dtype,
|
||||||
|
gate_type,
|
||||||
|
norm_dtype,
|
||||||
|
param_type,
|
||||||
|
norm_type,
|
||||||
|
):
|
||||||
|
B, S, D = 2, 4, 67
|
||||||
|
|
||||||
|
x = torch.randn(B, S, D, dtype=input_dtype)
|
||||||
|
residual = torch.randn(B, S, D, dtype=input_dtype)
|
||||||
|
|
||||||
|
gate_dtype = (
|
||||||
|
input_dtype
|
||||||
|
if gate_type == "input"
|
||||||
|
else torch.float32
|
||||||
|
if gate_type == "fp32"
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
param_dtype = input_dtype if param_type == "input" else torch.float32
|
||||||
|
|
||||||
|
gate = torch.randn(D, dtype=gate_dtype) if gate_dtype is not None else None
|
||||||
|
weight = torch.randn(D, dtype=norm_dtype) if norm_dtype is not None else None
|
||||||
|
bias = (
|
||||||
|
torch.randn(D, dtype=norm_dtype)
|
||||||
|
if norm_dtype is not None and norm_type == "layer"
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
|
scale = torch.randn(B, 1, D, dtype=param_dtype)
|
||||||
|
shift = torch.randn(B, S, D, dtype=param_dtype)
|
||||||
|
|
||||||
|
scale_expanded = scale.expand_as(x)
|
||||||
|
shift_expanded = shift.expand_as(x)
|
||||||
|
gate_expanded = gate.view(1, 1, D).expand_as(x) if gate is not None else None
|
||||||
|
|
||||||
|
out, residual_out = (
|
||||||
|
torch.ops.sgl_kernel.fused_scale_residual_norm_scale_shift_cpu(
|
||||||
|
residual,
|
||||||
|
x,
|
||||||
|
gate_expanded,
|
||||||
|
weight,
|
||||||
|
bias,
|
||||||
|
scale_expanded,
|
||||||
|
shift_expanded,
|
||||||
|
norm_type,
|
||||||
|
eps=eps,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if gate is None:
|
||||||
|
residual_fp32 = residual.float() + x.float()
|
||||||
|
else:
|
||||||
|
residual_fp32 = residual.float() + x.float() * gate.float()
|
||||||
|
|
||||||
|
ref_residual = residual_fp32.to(input_dtype)
|
||||||
|
norm_input = ref_residual.float()
|
||||||
|
|
||||||
|
if norm_type == "rms":
|
||||||
|
normalized = self.rmsnorm_ref(norm_input, weight, eps)
|
||||||
|
else:
|
||||||
|
normalized = torch.nn.functional.layer_norm(
|
||||||
|
norm_input,
|
||||||
|
(D,),
|
||||||
|
weight.float() if weight is not None else None,
|
||||||
|
bias.float() if bias is not None else None,
|
||||||
|
eps,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Match CUDA activation boundary after norm.
|
||||||
|
normalized = normalized.to(input_dtype).float()
|
||||||
|
|
||||||
|
ref_out = (normalized * (1.0 + scale.float()) + shift.float()).to(input_dtype)
|
||||||
|
|
||||||
|
torch.testing.assert_close(
|
||||||
|
residual_out,
|
||||||
|
ref_residual,
|
||||||
|
atol=precision[input_dtype],
|
||||||
|
rtol=precision[input_dtype],
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.testing.assert_close(
|
||||||
|
out, ref_out, atol=precision[input_dtype], rtol=precision[input_dtype]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__]))
|
||||||
Reference in New Issue
Block a user