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