[CPU] improve silu performance by replacing fp32 div with rcp14 (#31304)

This commit is contained in:
Ma Mingfei
2026-07-17 12:02:37 +08:00
committed by GitHub
parent 27a52d2530
commit 486a56be56
15 changed files with 509 additions and 557 deletions
+6 -13
View File
@@ -25,13 +25,8 @@ void act_and_mul_kernel_impl(
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= dim - kVecSize; d += kVecSize) {
bVec x_bvec = bVec::loadu(input_ptr + d);
fVec x_fvec0, x_fvec1;
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
bVec y_bvec = bVec::loadu(input_other_ptr + d);
fVec y_fvec0, y_fvec1;
std::tie(y_fvec0, y_fvec1) = at::vec::convert_to_float(y_bvec);
auto [x_fvec0, x_fvec1] = load_float_vec2(input_ptr + d);
auto [y_fvec0, y_fvec1] = load_float_vec2(input_other_ptr + d);
x_fvec0 = vf(x_fvec0);
x_fvec1 = vf(x_fvec1);
@@ -39,8 +34,7 @@ void act_and_mul_kernel_impl(
x_fvec0 = x_fvec0 * y_fvec0;
x_fvec1 = x_fvec1 * y_fvec1;
x_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
x_bvec.store(output_ptr + d);
convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1).store(output_ptr + d);
}
#pragma GCC unroll 4
for (; d < dim; ++d) {
@@ -69,7 +63,6 @@ void fused_sigmoid_mul_kernel_impl(
using fVec = at::vec::Vectorized<float>;
constexpr int64_t kVecSize = bVec::size();
const fVec one = fVec(1.f);
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
for (int64_t i = begin; i < end; ++i) {
const scalar_t* __restrict__ i_ptr = input + i * dim;
@@ -86,8 +79,8 @@ void fused_sigmoid_mul_kernel_impl(
for (; d <= head_dim - kVecSize; d += kVecSize) {
auto [x_fvec0, x_fvec1] = load_float_vec2(attn_ptr + d);
auto [g_fvec0, g_fvec1] = load_float_vec2(gate_ptr + d);
x_fvec0 = x_fvec0 / (one + g_fvec0.neg().exp_u20());
x_fvec1 = x_fvec1 / (one + g_fvec1.neg().exp_u20());
x_fvec0 = x_fvec0 * fast_sigmoid(g_fvec0);
x_fvec1 = x_fvec1 * fast_sigmoid(g_fvec1);
convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1).store(out_ptr + d);
}
#pragma GCC unroll 4
@@ -121,7 +114,7 @@ at::Tensor silu_and_mul_cpu(at::Tensor& input) {
num_tokens,
d,
[](float x) { return x / (1.f + std::exp(-x)); },
[](Vec x) { return x / (Vec(1.f) + x.neg().exp_u20()); });
[](Vec x) { return fast_silu(x); });
});
return out;
}
+4 -5
View File
@@ -59,13 +59,12 @@ inline void copy_add_stub(
constexpr int kVecSize = bVec::size();
for (int64_t d = 0; d < N; d += kVecSize) {
fVec bias0, bias1;
bVec bias_vec = bVec::loadu(bias + d);
std::tie(bias0, bias1) = at::vec::convert_to_float(bias_vec);
auto [bias0, bias1] = load_float_vec2(bias + d);
for (int64_t m = 0; m < M; ++m) {
fVec data0 = fVec::loadu(Ctmp + m * N + d) + bias0;
fVec data1 = fVec::loadu(Ctmp + m * N + d + fVec::size()) + bias1;
auto [data0, data1] = load_float_vec2(Ctmp + m * N + d);
data0 = data0 + bias0;
data1 = data1 + bias1;
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
out_vec.store(C + m * ldc + d);
}
+4 -4
View File
@@ -150,8 +150,9 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc,
int64_t d = 0;
#pragma GCC unroll 4
for (; d <= size - kVecSize; d += kVecSize) {
fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec;
fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec;
auto [a_fvec0, a_fvec1] = load_float_vec2(acc + d);
a_fvec0 = a_fvec0 * s_fvec;
a_fvec1 = a_fvec1 * s_fvec;
bVec out_bvec = convert_from_float_ext<scalar_t>(a_fvec0, a_fvec1);
out_bvec.store(out + d);
}
@@ -186,8 +187,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu
constexpr int col = i % COLS;
// for COLS = 2, 4 use 512bit store
if constexpr (col % 2 == 0) {
fVec a_fvec0 = fVec::loadu(input + col * 16);
fVec a_fvec1 = fVec::loadu(input + col * 16 + 16);
auto [a_fvec0, a_fvec1] = load_float_vec2(input + col * 16);
bVec out_bvec = convert_from_float_ext<scalar_t>(a_fvec0, a_fvec1);
out_bvec.store(out + col * 16);
}
+4 -4
View File
@@ -29,8 +29,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu
constexpr int col = i % COLS;
// for COLS = 2, 4 use 512bit store
if constexpr (col % 2 == 0) {
fVec a_fvec0 = fVec::loadu(input + col * 16);
fVec a_fvec1 = fVec::loadu(input + col * 16 + 16);
auto [a_fvec0, a_fvec1] = load_float_vec2(input + col * 16);
bVec out_bvec = convert_from_float_ext<scalar_t>(a_fvec0, a_fvec1);
out_bvec.store(out + col * 16);
}
@@ -47,8 +46,9 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ acc,
int d = 0;
#pragma GCC unroll 4
for (; d <= size - kVecSize; d += kVecSize) {
fVec a_fvec0 = fVec::loadu(acc + d) * s_fvec;
fVec a_fvec1 = fVec::loadu(acc + d + fVec::size()) * s_fvec;
auto [a_fvec0, a_fvec1] = load_float_vec2(acc + d);
a_fvec0 = a_fvec0 * s_fvec;
a_fvec1 = a_fvec1 * s_fvec;
bVec out_bvec = convert_from_float_ext<scalar_t>(a_fvec0, a_fvec1);
out_bvec.store(out + d);
}
+8 -17
View File
@@ -111,8 +111,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d);
fVec data1 = fVec::loadu(input + d + fVec::size());
auto [data0, data1] = load_float_vec2(input + d);
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
out_vec.store(out + d);
}
@@ -130,9 +129,7 @@ inline void copy_stub(float* __restrict__ out, const scalar_t* __restrict__ inpu
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0, data1;
bVec b_vec = bVec::loadu(input + d);
std::tie(data0, data1) = at::vec::convert_to_float(b_vec);
auto [data0, data1] = load_float_vec2(input + d);
data0.store(out + d);
data1.store(out + d + fVec::size());
}
@@ -151,9 +148,9 @@ inline void copy_add_stub(
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d) + fVec::loadu(bias + d);
fVec data1 = fVec::loadu(input + d + fVec::size()) + fVec::loadu(bias + d + fVec::size());
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
auto [data0, data1] = load_float_vec2(input + d);
auto [bias0, bias1] = load_float_vec2(bias + d);
bVec out_vec = convert_from_float_ext<scalar_t>(data0 + bias0, data1 + bias1);
out_vec.store(out + d);
}
for (; d < size; ++d) {
@@ -171,7 +168,6 @@ inline void scalar_sigmoid_and_mul(
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
// scalar sigmoid
const fVec one = fVec(1.f);
fVec X;
if constexpr (has_bias) {
assert(bias != nullptr);
@@ -179,18 +175,13 @@ inline void scalar_sigmoid_and_mul(
} else {
X = fVec(input[0]);
}
X = one / (one + X.neg().exp_u20());
X = fast_sigmoid(X);
// vec mul
constexpr int kVecSize = bVec::size();
for (int d = 0; d < SIZE; d += kVecSize) {
bVec m_bvec = bVec::loadu(mul + d);
fVec m_fvec0, m_fvec1;
std::tie(m_fvec0, m_fvec1) = at::vec::convert_to_float(m_bvec);
m_fvec0 = m_fvec0 * X;
m_fvec1 = m_fvec1 * X;
bVec out_vec = convert_from_float_ext<scalar_t>(m_fvec0, m_fvec1);
auto [m_fvec0, m_fvec1] = load_float_vec2(mul + d);
bVec out_vec = convert_from_float_ext<scalar_t>(m_fvec0 * X, m_fvec1 * X);
out_vec.store(out + d);
}
}
+6 -8
View File
@@ -13,8 +13,7 @@ inline void copy_stub(scalar_t* __restrict__ out, const float* __restrict__ inpu
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d);
fVec data1 = fVec::loadu(input + d + fVec::size());
auto [data0, data1] = load_float_vec2(input + d);
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
out_vec.store(out + d);
}
@@ -33,9 +32,9 @@ inline void copy_add_stub(
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d) + fVec::loadu(bias + d);
fVec data1 = fVec::loadu(input + d + fVec::size()) + fVec::loadu(bias + d + fVec::size());
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
auto [data0, data1] = load_float_vec2(input + d);
auto [bias0, bias1] = load_float_vec2(bias + d);
bVec out_vec = convert_from_float_ext<scalar_t>(data0 + bias0, data1 + bias1);
out_vec.store(out + d);
}
for (; d < size; ++d) {
@@ -52,9 +51,8 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__
int d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d) * vscale;
fVec data1 = fVec::loadu(input + d + fVec::size()) * vscale;
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
auto [data0, data1] = load_float_vec2(input + d);
bVec out_vec = convert_from_float_ext<scalar_t>(data0 * vscale, data1 * vscale);
out_vec.store(out + d);
}
for (; d < size; ++d) {
+2 -3
View File
@@ -177,14 +177,13 @@ struct tinygemm_kernel<at::BFloat16, K, BLOCK_N, has_bias, has_silu> {
using fVec = at::vec::Vectorized<float>;
using bVec = at::vec::Vectorized<at::BFloat16>;
const fVec one = fVec(1.f);
auto storec = [&](auto i, int64_t m) {
constexpr int col = i;
fVec x0 = fVec(vc[col * 2 + 0]);
fVec x1 = fVec(vc[col * 2 + 1]);
if constexpr (has_silu) {
x0 = x0 / (one + x0.neg().exp_u20());
x1 = x1 / (one + x1.neg().exp_u20());
x0 = fast_silu(x0);
x1 = fast_silu(x1);
}
bVec out_vec = convert_from_float_ext<at::BFloat16>(x0, x1);
out_vec.store(C + m * lda + col * 32);
+17 -38
View File
@@ -1146,14 +1146,10 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl(
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= head_dim - VecSize; d += VecSize) {
bVec q_bvec = bVec::loadu(q_ptr + q_offset + d);
fVec q_fvec0, q_fvec1;
std::tie(q_fvec0, q_fvec1) = at::vec::convert_to_float(q_bvec);
auto [q_fvec0, q_fvec1] = load_float_vec2(q_ptr + q_offset + d);
sum_q_fvec += q_fvec0 * q_fvec0;
sum_q_fvec += q_fvec1 * q_fvec1;
bVec k_bvec = bVec::loadu(k_ptr + k_offset + d);
fVec k_fvec0, k_fvec1;
std::tie(k_fvec0, k_fvec1) = at::vec::convert_to_float(k_bvec);
auto [k_fvec0, k_fvec1] = load_float_vec2(k_ptr + k_offset + d);
sum_k_fvec += k_fvec0 * k_fvec0;
sum_k_fvec += k_fvec1 * k_fvec1;
}
@@ -1200,14 +1196,11 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl(
fVec kv_mem_vec1 = fVec(float(0));
for (int di = 0; di < head_dim; ++di) {
fVec k_val_vec = fVec(k_ptr[k_offset + di] * k_scale);
fVec state_vec0 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi);
fVec state_vec1 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi + fVecSize);
auto [state_vec0, state_vec1] = load_float_vec2(state_ptr + state_offset + di * v_head_dim + dvi);
kv_mem_vec0 = kv_mem_vec0 + state_vec0 * g_val_exp_vec * k_val_vec;
kv_mem_vec1 = kv_mem_vec1 + state_vec1 * g_val_exp_vec * k_val_vec;
}
bVec v_bvec = bVec::loadu(v_ptr + v_offset + dvi);
fVec v_vec0, v_vec1;
std::tie(v_vec0, v_vec1) = at::vec::convert_to_float(v_bvec);
auto [v_vec0, v_vec1] = load_float_vec2(v_ptr + v_offset + dvi);
fVec dt_vec0 = (v_vec0 - kv_mem_vec0) * beta_vec;
fVec dt_vec1 = (v_vec1 - kv_mem_vec1) * beta_vec;
fVec o_vec0 = fVec(float(0));
@@ -1215,8 +1208,7 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl(
for (int di = 0; di < head_dim; ++di) {
fVec q_vec = fVec(q_ptr[q_offset + di] * q_scale);
fVec k_vec = fVec(k_ptr[k_offset + di] * k_scale);
fVec state_vec0 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi);
fVec state_vec1 = fVec::loadu(state_ptr + state_offset + di * v_head_dim + dvi + fVecSize);
auto [state_vec0, state_vec1] = load_float_vec2(state_ptr + state_offset + di * v_head_dim + dvi);
state_vec0 = state_vec0 * g_val_exp_vec + k_vec * dt_vec0;
state_vec1 = state_vec1 * g_val_exp_vec + k_vec * dt_vec1;
o_vec0 = o_vec0 + state_vec0 * q_vec * scale_vec;
@@ -1265,25 +1257,19 @@ void fused_gdn_gating_kernel_impl(
constexpr int vec_size = bVec::size();
constexpr int fvec_size = fVec::size();
const fVec neg_one(-1.0f);
const fVec one(1.0f);
at::parallel_for(0, batch, 0, [&](int64_t begin, int64_t end) {
for (int64_t i = begin; i < end; ++i) {
int64_t j = 0;
for (; j < num_heads - (num_heads % vec_size); j += vec_size) {
fVec A_log_vec0 = fVec::loadu(A_log + j);
fVec A_log_vec1 = fVec::loadu(A_log + j + fvec_size);
bVec dt_bias_vec = bVec::loadu(dt_bias + j);
bVec a_bvec = bVec::loadu(a + i * num_heads + j);
bVec b_bvec = bVec::loadu(b + i * num_heads + j);
fVec a0, a1, dt_bias_vec0, dt_bias_vec1, b0, b1;
std::tie(a0, a1) = at::vec::convert_to_float(a_bvec);
std::tie(b0, b1) = at::vec::convert_to_float(b_bvec);
std::tie(dt_bias_vec0, dt_bias_vec1) = at::vec::convert_to_float(dt_bias_vec);
auto [A_log_vec0, A_log_vec1] = load_float_vec2(A_log + j);
auto [dt_bias_vec0, dt_bias_vec1] = load_float_vec2(dt_bias + j);
auto [a0, a1] = load_float_vec2(a + i * num_heads + j);
auto [b0, b1] = load_float_vec2(b + i * num_heads + j);
fVec g0 = neg_one * A_log_vec0.exp_u20() * softplus(a0 + dt_bias_vec0);
fVec g1 = neg_one * A_log_vec1.exp_u20() * softplus(a1 + dt_bias_vec1);
fVec beta0 = one / (one + (neg_one * b0).exp_u20());
fVec beta1 = one / (one + (neg_one * b1).exp_u20());
fVec beta0 = fast_sigmoid(b0);
fVec beta1 = fast_sigmoid(b1);
g0.store(out + i * num_heads + j);
g1.store(out + i * num_heads + j + fvec_size);
@@ -1313,26 +1299,19 @@ void fused_gdn_gating_kernel_impl(
constexpr int vec_size = bVec::size();
constexpr int fvec_size = fVec::size();
const fVec neg_one(-1.0f);
const fVec one(1.0f);
at::parallel_for(0, batch, 0, [&](int64_t begin, int64_t end) {
for (int64_t i = begin; i < end; ++i) {
int64_t j = 0;
for (; j < num_heads - (num_heads % vec_size); j += vec_size) {
bVec A_log_bvec = bVec::loadu(A_log + j);
fVec A_log_vec0, A_log_vec1;
std::tie(A_log_vec0, A_log_vec1) = at::vec::convert_to_float(A_log_bvec);
bVec dt_bias_vec = bVec::loadu(dt_bias + j);
bVec a_bvec = bVec::loadu(a + i * num_heads + j);
bVec b_bvec = bVec::loadu(b + i * num_heads + j);
fVec a0, a1, dt_bias_vec0, dt_bias_vec1, b0, b1;
std::tie(a0, a1) = at::vec::convert_to_float(a_bvec);
std::tie(b0, b1) = at::vec::convert_to_float(b_bvec);
std::tie(dt_bias_vec0, dt_bias_vec1) = at::vec::convert_to_float(dt_bias_vec);
auto [A_log_vec0, A_log_vec1] = load_float_vec2(A_log + j);
auto [dt_bias_vec0, dt_bias_vec1] = load_float_vec2(dt_bias + j);
auto [a0, a1] = load_float_vec2(a + i * num_heads + j);
auto [b0, b1] = load_float_vec2(b + i * num_heads + j);
fVec g0 = neg_one * A_log_vec0.exp_u20() * softplus(a0 + dt_bias_vec0);
fVec g1 = neg_one * A_log_vec1.exp_u20() * softplus(a1 + dt_bias_vec1);
fVec beta0 = one / (one + (neg_one * b0).exp_u20());
fVec beta1 = one / (one + (neg_one * b1).exp_u20());
fVec beta0 = fast_sigmoid(b0);
fVec beta1 = fast_sigmoid(b1);
g0.store(out + i * num_heads + j);
g1.store(out + i * num_heads + j + fvec_size);
+17 -107
View File
@@ -118,96 +118,6 @@ int moe_align_block_size(
return num_tokens_post_pad;
}
// silu : shape leading dimension
// input0 [m_size, BLOCK_N] BLOCK_N
// input1 [m_size, BLOCK_N] BLOCK_N
// output [M * topk, N] N
template <typename scalar_t, int BLOCK_N>
inline void silu_and_mul(
scalar_t* __restrict__ output,
const float* __restrict__ input0, // x: x0, x1
const float* __restrict__ input1, // y: y0, y1
int64_t m_size,
int64_t N) {
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
// no remainder
for (int64_t m = 0; m < m_size; ++m) {
scalar_t* __restrict__ out = output + m * N;
const float* __restrict__ x = input0 + m * BLOCK_N;
const float* __restrict__ y = input1 + m * BLOCK_N;
for (int64_t d = 0; d < BLOCK_N; d += bVec::size()) {
fVec x0 = fVec::loadu(x + d);
fVec x1 = fVec::loadu(x + d + fVec::size());
fVec y0 = fVec::loadu(y + d);
fVec y1 = fVec::loadu(y + d + fVec::size());
// silu
x0 = x0 / (one + x0.neg().exp_u20());
x1 = x1 / (one + x1.neg().exp_u20());
// mul
x0 = x0 * y0;
x1 = x1 * y1;
// convert
bVec out_vec = convert_from_float_ext<scalar_t>(x0, x1);
out_vec.store(out + d);
}
}
}
template <typename scalar_t, int BLOCK_N>
inline void clamp_sigmoid_and_mul(
scalar_t* __restrict__ output,
const float* __restrict__ input0,
int64_t m_size,
int64_t N,
const float alpha,
const float limit,
int64_t offset) {
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
const fVec zero = fVec(0.f);
const fVec limit_v = fVec(limit);
const fVec nlimit_v = fVec(-limit);
const fVec alpha_v = fVec(alpha);
// no remainder
for (int64_t m = 0; m < m_size; ++m) {
scalar_t* __restrict__ out = output + m * N;
const float* __restrict__ cur_ptr = input0 + m * BLOCK_N;
for (int64_t d = 0; d < BLOCK_N; d += bVec::size()) {
float tmp_glu0[fVec::size()]; // 16
float tmp_linear0[fVec::size()]; // 16
// interleaved: x[2i] = glu, x[2i+1] = linear
for (int j = 0; j < fVec::size(); ++j) {
// x0 [0,2,..30]
tmp_glu0[j] = cur_ptr[d + j * 2];
// y0 [1,3,...31]
tmp_linear0[j] = cur_ptr[d + j * 2 + 1];
}
fVec x0 = fVec::loadu(tmp_glu0);
fVec y0 = fVec::loadu(tmp_linear0);
// clamp
x0 = at::vec::minimum(x0, limit_v);
y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0));
// x * sigmoid(x * alpha)
x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20());
// (y + 1) * x
y0 = y0 + one;
x0 = x0 * y0;
// // convert
convert_from_float_and_store<scalar_t>(out + d / 2 + offset, x0);
}
}
}
template <typename scalar_t, int BLOCK_M, int BLOCK_N>
struct tinygemm_kernel_nn2 {
static inline void apply(
@@ -284,23 +194,17 @@ struct tinygemm_kernel_nn2<at::BFloat16, BLOCK_M, BLOCK_N> {
Unroll<ROWS * COLS>{}(compute, k);
}
using Vec = at::vec::Vectorized<float>;
const Vec one = Vec(1.f);
auto storec = [&](auto i) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
// for COLS = 2, 4 use 512bit store
if constexpr (col % 2 == 0) {
Vec x0 = vc0[row * COLS + col + 0];
Vec x1 = vc0[row * COLS + col + 1];
Vec y0 = vc1[row * COLS + col + 0];
Vec y1 = vc1[row * COLS + col + 1];
// silu
x0 = x0 / (one + x0.neg().exp_u20());
x1 = x1 / (one + x1.neg().exp_u20());
// mul
x0 = x0 * y0;
x1 = x1 * y1;
__m512 x0 = vc0[row * COLS + col + 0];
__m512 x1 = vc0[row * COLS + col + 1];
__m512 y0 = vc1[row * COLS + col + 0];
__m512 y1 = vc1[row * COLS + col + 1];
x0 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x0), y0);
x1 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x1), y1);
_mm512_storeu_si512(
reinterpret_cast<__m512i*>((C + row * ldc + col * 16)),
@@ -638,11 +542,15 @@ void fused_experts_kernel_impl(
// 1.d silu and mul
const int64_t offset = offsets[mb];
if (act_func == CPUActMethod::silu_and_mul && use_brgemm) {
silu_and_mul<scalar_t, BLOCK_N>(ic1 + offset * N + nb * BLOCK_N, C0, C1, m_size, N);
for (int64_t m = 0; m < m_size; ++m) {
silu_and_mul_stub(ic1 + (offset + m) * N + nb * BLOCK_N, C0 + m * BLOCK_N, C1 + m * BLOCK_N, BLOCK_N);
}
} else if (act_func == CPUActMethod::swiglu) {
clamp_sigmoid_and_mul<scalar_t, BLOCK_N>(ic1 + offset * N, C0, m_size, N, alpha, limit, 0 + nb * BLOCK_N / 2);
clamp_sigmoid_and_mul<scalar_t, BLOCK_N>(
ic1 + offset * N, C1, m_size, N, alpha, limit, N / 2 + nb * BLOCK_N / 2);
for (int64_t m = 0; m < m_size; ++m) {
scalar_t* __restrict__ ic1_row = ic1 + (offset + m) * N;
clamp_sigmoid_and_mul_stub(ic1_row + nb * BLOCK_N / 2, C0 + m * BLOCK_N, BLOCK_N / 2, alpha, limit);
clamp_sigmoid_and_mul_stub(ic1_row + N / 2 + nb * BLOCK_N / 2, C1 + m * BLOCK_N, BLOCK_N / 2, alpha, limit);
}
}
});
@@ -811,7 +719,9 @@ void shared_expert_kernel_impl(
/* C */ C1);
// 1.d silu and mul
silu_and_mul<scalar_t, BLOCK_N>(ic1 + mb * BLOCK_M * N + nb * BLOCK_N, C0, C1, m_size, N);
for (int64_t m = 0; m < m_size; ++m) {
silu_and_mul_stub(ic1 + (mb * BLOCK_M + m) * N + nb * BLOCK_N, C0 + m * BLOCK_N, C1 + m * BLOCK_N, BLOCK_N);
}
} else {
// fused 1.bcd: silu_and_mul(A @ B0, A @ B1)
tinygemm_kernel(
+41 -82
View File
@@ -59,9 +59,7 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const input_t* __restrict_
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
auto [x0, x1] = load_float_vec2(input + d);
x0 = x0 * weight_vec;
x1 = x1 * weight_vec;
bVec out_vec = convert_from_float_ext<scalar_t>(x0, x1);
bVec out_vec = convert_from_float_ext<scalar_t>(x0 * weight_vec, x1 * weight_vec);
out_vec.store(out + d);
}
for (; d < size; ++d) {
@@ -86,10 +84,7 @@ inline void sum_stub(scalar_t* __restrict__ out, const scalar_t* __restrict__ in
fVec sum_fvec0 = fVec(0.f);
fVec sum_fvec1 = fVec(0.f);
for (int t = 0; t < topk; ++t) {
bVec x_bvec = bVec::loadu(input + t * K + d);
fVec x_fvec0, x_fvec1;
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
auto [x_fvec0, x_fvec1] = load_float_vec2(input + t * K + d);
sum_fvec0 += x_fvec0;
sum_fvec1 += x_fvec1;
}
@@ -132,11 +127,7 @@ inline void add_mul_stub(
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
auto [x0, x1] = load_float_vec2(input + d);
bVec y_bvec = bVec::loadu(input2 + d);
fVec y0, y1;
std::tie(y0, y1) = at::vec::convert_to_float(y_bvec);
auto [y0, y1] = load_float_vec2(input2 + d);
x0 = x0 + y0 * s_vec;
x1 = x1 + y1 * s_vec;
bVec out_vec = convert_from_float_ext<scalar_t>(x0, x1);
@@ -147,31 +138,52 @@ inline void add_mul_stub(
}
}
template <typename scalar_t>
template <typename scalar_t, typename input_t>
inline void silu_and_mul_stub(
scalar_t* __restrict__ out, const scalar_t* __restrict__ input, const scalar_t* __restrict__ input2, int64_t size) {
scalar_t* __restrict__ out, const input_t* __restrict__ input, const input_t* __restrict__ input2, int64_t size) {
static_assert(
std::is_same_v<input_t, float> || std::is_same_v<input_t, scalar_t>,
"silu_and_mul_stub only supports input_t == float or input_t == scalar_t");
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
// no remainder
#pragma GCC unroll 4
for (int64_t d = 0; d < size; d += bVec::size()) {
bVec x = bVec::loadu(input + d);
fVec x0, x1;
std::tie(x0, x1) = at::vec::convert_to_float(x);
bVec y = bVec::loadu(input2 + d);
fVec y0, y1;
std::tie(y0, y1) = at::vec::convert_to_float(y);
x0 = x0 / (one + x0.neg().exp_u20());
x1 = x1 / (one + x1.neg().exp_u20());
x0 = x0 * y0;
x1 = x1 * y1;
auto [x0, x1] = load_float_vec2(input + d);
auto [y0, y1] = load_float_vec2(input2 + d);
x0 = fast_silu(x0) * y0;
x1 = fast_silu(x1) * y1;
bVec out_vec = convert_from_float_ext<scalar_t>(x0, x1);
out_vec.store(out + d);
}
}
template <typename scalar_t, typename input_t>
inline void clamp_sigmoid_and_mul_stub(
scalar_t* __restrict__ out, const input_t* __restrict__ input, int64_t size, const float alpha, const float limit) {
static_assert(
std::is_same_v<input_t, float> || std::is_same_v<input_t, scalar_t>,
"clamp_sigmoid_and_mul_stub only supports input_t == float or input_t == scalar_t");
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
const fVec limit_v = fVec(limit);
const fVec nlimit_v = fVec(-limit);
const fVec alpha_v = fVec(alpha);
#pragma GCC unroll 4
for (int64_t d = 0; d < 2 * size; d += bVec::size()) {
auto [x0_, y0_] = load_float_vec2(input + d);
auto [x0, y0] = at::vec::deinterleave2<float>(x0_, y0_);
x0 = at::vec::minimum(x0, limit_v);
y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0));
x0 = fast_sigmoid_glu(x0, alpha_v) * (y0 + one);
store_from_float_ext(out + d / 2, x0);
}
}
template <typename scalar_t>
inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__ input, float weight, int64_t size) {
using bVec = at::vec::Vectorized<scalar_t>;
@@ -181,9 +193,8 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const float* __restrict__
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
fVec data0 = fVec::loadu(input + d) * weight_vec;
fVec data1 = fVec::loadu(input + d + fVec::size()) * weight_vec;
bVec out_vec = convert_from_float_ext<scalar_t>(data0, data1);
auto [x0, x1] = load_float_vec2(input + d);
bVec out_vec = convert_from_float_ext<scalar_t>(x0 * weight_vec, x1 * weight_vec);
out_vec.store(out + d);
}
for (; d < size; ++d) {
@@ -217,63 +228,11 @@ inline void copy_mul_stub(scalar_t* __restrict__ out, const scalar_t* __restrict
int64_t d;
#pragma GCC unroll 4
for (d = 0; d <= size - kVecSize; d += kVecSize) {
bVec x = bVec::loadu(input + d);
fVec x0, x1;
std::tie(x0, x1) = at::vec::convert_to_float(x);
x0 = x0 * weight_vec;
x1 = x1 * weight_vec;
bVec out_vec = convert_from_float_ext<scalar_t>(x0, x1);
auto [x0, x1] = load_float_vec2(input + d);
bVec out_vec = convert_from_float_ext<scalar_t>(x0 * weight_vec, x1 * weight_vec);
out_vec.store(out + d);
}
for (; d < size; ++d) {
out[d] = static_cast<scalar_t>(input[d] * weight);
}
}
template <typename scalar_t>
inline void clamp_sigmoid_and_mul_stub(
scalar_t* __restrict__ out,
const scalar_t* __restrict__ input,
int64_t size,
const float alpha,
const float limit) {
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
const fVec zero = fVec(0.f);
const fVec limit_v = fVec(limit);
const fVec nlimit_v = fVec(-limit);
const fVec alpha_v = fVec(alpha);
// no remainder
#pragma GCC unroll 4
for (int64_t d = 0; d < size; d += bVec::size()) {
bVec x = bVec::loadu(input + d);
fVec x0_, y0_;
std::tie(x0_, y0_) = at::vec::convert_to_float(x);
float tmp_buffer[fVec::size() * 2]; // 32
float tmp_glu[fVec::size()]; // 16
float tmp_linear[fVec::size()]; // 16
x0_.store(tmp_buffer);
y0_.store(tmp_buffer + fVec::size());
// interleaved: x[2i] = glu, x[2i+1] = linear
for (int j = 0; j < fVec::size(); ++j) {
// x0 [0,2,..30]
tmp_glu[j] = tmp_buffer[j * 2];
// y0 [1,3,...31]
tmp_linear[j] = tmp_buffer[j * 2 + 1];
}
fVec x0 = fVec::loadu(tmp_glu);
fVec y0 = fVec::loadu(tmp_linear);
// clamp
x0 = at::vec::minimum(x0, limit_v);
y0 = at::vec::minimum(limit_v, at::vec::maximum(nlimit_v, y0));
// x * sigmoid(x * alpha)
x0 = x0 / (one + (x0 * alpha_v).neg().exp_u20());
// (y + 1) * x
y0 = y0 + one;
x0 = x0 * y0;
convert_from_float_and_store<scalar_t>(out + d / 2, x0);
}
}
+2 -2
View File
@@ -126,8 +126,8 @@ void fused_experts_fp_kernel_impl(
} else if (act_func == CPUActMethod::swiglu) {
at::parallel_for(0, M * topk, 0, [&](int64_t begin, int64_t end) {
for (int64_t m = begin; m < end; ++m) {
clamp_sigmoid_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, N, alpha, limit);
clamp_sigmoid_and_mul_stub(ic1 + m * N + N / 2, ic0 + m * 2 * N + N, N, alpha, limit);
clamp_sigmoid_and_mul_stub(ic1 + m * N, ic0 + m * 2 * N, N / 2, alpha, limit);
clamp_sigmoid_and_mul_stub(ic1 + m * N + N / 2, ic0 + m * 2 * N + N, N / 2, alpha, limit);
}
});
}
+11 -19
View File
@@ -48,16 +48,14 @@ inline void silu_and_mul(
vc1[col] = _mm512_mul_ps(_mm512_mul_ps(vc1[col], vas), vbs1[col]);
};
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
const fVec one = fVec(1.f);
auto silu_and_mul = [&](auto col) {
fVec x = fVec(vc0[col]);
fVec y = fVec(vc1[col]);
x = x / (one + x.neg().exp_u20());
vc0[col] = x * y;
__m512 x = vc0[col];
__m512 y = vc1[col];
vc0[col] = _mm512_mul_ps(_mm512_rcp14_silu_ps(x), y);
};
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
auto storec = [&](auto col, int64_t m) {
if constexpr (col % 2 == 0) {
fVec x0 = fVec(vc0[col + 0]);
@@ -224,23 +222,17 @@ struct tinygemm_kernel_vnni<at::BFloat16, BLOCK_M, BLOCK_N> {
};
Unroll<ROWS * COLS>{}(scalec);
using Vec = at::vec::Vectorized<float>;
const Vec one = Vec(1.f);
auto storec = [&](auto i) {
constexpr int row = i / COLS;
constexpr int col = i % COLS;
// for COLS = 2, 4 use 512bit store
if constexpr (col % 2 == 0) {
Vec x0 = _mm512_castsi512_ps(vc0[row * COLS + col + 0]);
Vec x1 = _mm512_castsi512_ps(vc0[row * COLS + col + 1]);
Vec y0 = _mm512_castsi512_ps(vc1[row * COLS + col + 0]);
Vec y1 = _mm512_castsi512_ps(vc1[row * COLS + col + 1]);
// silu
x0 = x0 / (one + x0.neg().exp_u20());
x1 = x1 / (one + x1.neg().exp_u20());
// mul
x0 = x0 * y0;
x1 = x1 * y1;
__m512 x0 = _mm512_castsi512_ps(vc0[row * COLS + col + 0]);
__m512 x1 = _mm512_castsi512_ps(vc0[row * COLS + col + 1]);
__m512 y0 = _mm512_castsi512_ps(vc1[row * COLS + col + 0]);
__m512 y1 = _mm512_castsi512_ps(vc1[row * COLS + col + 1]);
x0 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x0), y0);
x1 = _mm512_mul_ps(_mm512_rcp14_silu_ps(x1), y1);
_mm512_storeu_si512(
reinterpret_cast<__m512i*>((C + row * ldc + col * 16)),
+3 -9
View File
@@ -240,9 +240,7 @@ inline float reduce(const scalar_t* __restrict__ x, int64_t size) {
// no remainder
#pragma GCC unroll 4
for (int64_t d = 0; d < size; d += bVec::size()) {
bVec x_bvec = bVec::loadu(x + d);
fVec x_fvec0, x_fvec1;
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
auto [x_fvec0, x_fvec1] = load_float_vec2(x + d);
sum_fvec += x_fvec0 * x_fvec0;
sum_fvec += x_fvec1 * x_fvec1;
}
@@ -259,12 +257,8 @@ inline void map2(scalar_t* y, const scalar_t* x, const scalar_t* __restrict__ w,
// no remainder
#pragma GCC unroll 4
for (int64_t d = 0; d < size; d += bVec::size()) {
bVec x_bvec = bVec::loadu(x + d);
fVec x_fvec0, x_fvec1;
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
bVec w_bvec = bVec::loadu(w + d);
fVec w_fvec0, w_fvec1;
std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec);
auto [x_fvec0, x_fvec1] = load_float_vec2(x + d);
auto [w_fvec0, w_fvec1] = load_float_vec2(w + d);
x_fvec0 = x_fvec0 * scale_fvec * w_fvec0;
x_fvec1 = x_fvec1 * scale_fvec * w_fvec1;
bVec out_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
+57 -6
View File
@@ -17,11 +17,11 @@ inline Vectorized<scalar_t> convert_from_float_ext(const Vectorized<float>& a, c
}
template <typename scalar_t>
inline void convert_from_float_and_store(scalar_t* out, const Vectorized<float>& a) {
float out_buffer[at::vec::Vectorized<float>::size()];
inline void store_from_float_ext(scalar_t* out, const Vectorized<float>& a) {
float out_buffer[Vectorized<float>::size()];
a.store(out_buffer);
for (int i = 0; i < 16; i++) {
out[i] = (scalar_t)out_buffer[i];
for (int i = 0; i < Vectorized<float>::size(); ++i) {
out[i] = static_cast<scalar_t>(out_buffer[i]);
}
}
@@ -55,8 +55,14 @@ convert_from_float_ext<at::BFloat16>(const Vectorized<float>& a, const Vectorize
}
template <>
inline void convert_from_float_and_store<at::BFloat16>(at::BFloat16* out, const Vectorized<float>& a) {
_mm256_storeu_si256((__m256i*)out, (__m256i)(_mm512_cvtneps_pbh(__m512(a))));
inline void store_from_float_ext<at::BFloat16>(at::BFloat16* out, const Vectorized<float>& a) {
_mm256_storeu_si256(reinterpret_cast<__m256i*>(out), (__m256i)(_mm512_cvtneps_pbh(__m512(a))));
}
template <>
inline void store_from_float_ext<at::Half>(at::Half* out, const Vectorized<float>& a) {
_mm256_storeu_si256(
reinterpret_cast<__m256i*>(out), _mm512_cvtps_ph(__m512(a), _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC));
}
#define CVT_BF16_TO_FP32(a) _mm512_castsi512_ps(_mm512_slli_epi32(_mm512_cvtepu16_epi32(a), 16))
@@ -557,6 +563,51 @@ inline __attribute__((always_inline)) __m512 _mm512_fexp_u20_ps(const __m512 val
// final interpretation to float
return _mm512_castsi512_ps(casted_integer);
}
// sigmoid(x) = 1 / (1 + exp(-x)); avoid vdivps via rcp14
inline __attribute__((always_inline)) __m512 _mm512_rcp14_sigmoid_ps(__m512 x) {
__m512 minus_x = _mm512_xor_ps(_mm512_set1_ps(-0.f), x);
__m512 denom = _mm512_add_ps(_mm512_exp_u20_ps(minus_x), _mm512_set1_ps(1.f));
return _mm512_rcp14_ps(denom);
}
// SiLU(x) = x * sigmoid(x)
inline __attribute__((always_inline)) __m512 _mm512_rcp14_silu_ps(__m512 x) {
return _mm512_mul_ps(x, _mm512_rcp14_sigmoid_ps(x));
}
// x * sigmoid(x * alpha) for clamped SwiGLU
inline __attribute__((always_inline)) __m512 _mm512_rcp14_sigmoid_glu_ps(__m512 x, __m512 alpha) {
__m512 xa = _mm512_mul_ps(x, alpha);
return _mm512_mul_ps(x, _mm512_rcp14_sigmoid_ps(xa));
}
#endif
inline at::vec::Vectorized<float> fast_sigmoid(const at::vec::Vectorized<float>& x) {
#if defined(CPU_CAPABILITY_AVX512)
return at::vec::Vectorized<float>(_mm512_rcp14_sigmoid_ps(x));
#else
const auto one = at::vec::Vectorized<float>(1.f);
return one / (one + x.neg().exp_u20());
#endif
}
inline at::vec::Vectorized<float> fast_silu(const at::vec::Vectorized<float>& x) {
#if defined(CPU_CAPABILITY_AVX512)
return at::vec::Vectorized<float>(_mm512_rcp14_silu_ps(x));
#else
return x * fast_sigmoid(x);
#endif
}
inline at::vec::Vectorized<float>
fast_sigmoid_glu(const at::vec::Vectorized<float>& x, const at::vec::Vectorized<float>& alpha) {
#if defined(CPU_CAPABILITY_AVX512)
return at::vec::Vectorized<float>(_mm512_rcp14_sigmoid_glu_ps(x, alpha));
#else
return x * fast_sigmoid(x * alpha);
#endif
}
} // anonymous namespace