[CPU] improve silu performance by replacing fp32 div with rcp14 (#31304)
This commit is contained in:
+17
-107
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user