[CPU] Add GPT-OSS model optimization for CPU (#16775)
Co-authored-by: mingfeima <mingfei.ma@intel.com> Co-authored-by: jianan-gu <jianan.gu@intel.com>
This commit is contained in:
co-authored by
mingfeima
jianan-gu
parent
5601b7139d
commit
3ecf2c76ad
+217
-37
@@ -158,6 +158,56 @@ inline void silu_and_mul(
|
||||
}
|
||||
}
|
||||
|
||||
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(
|
||||
@@ -450,6 +500,8 @@ void fused_experts_kernel_impl(
|
||||
const scalar_t* __restrict__ input,
|
||||
const scalar_t* __restrict__ packed_w1,
|
||||
const scalar_t* __restrict__ packed_w2,
|
||||
const float* __restrict__ w1_bias,
|
||||
const float* __restrict__ w2_bias,
|
||||
const float* __restrict__ topk_weights,
|
||||
const int32_t* __restrict__ sorted_ids,
|
||||
const int32_t* __restrict__ expert_ids,
|
||||
@@ -459,7 +511,11 @@ void fused_experts_kernel_impl(
|
||||
int64_t K,
|
||||
int64_t E,
|
||||
int64_t topk,
|
||||
int64_t num_tokens_post_pad) {
|
||||
int64_t num_tokens_post_pad,
|
||||
float alpha,
|
||||
float limit,
|
||||
CPUActMethod act_func,
|
||||
bool with_bias) {
|
||||
// handle 2 tiles per block
|
||||
constexpr int64_t BLOCK_M = block_size_m();
|
||||
constexpr int64_t BLOCK_N = block_size_n();
|
||||
@@ -494,6 +550,8 @@ void fused_experts_kernel_impl(
|
||||
int32_t expert_id = expert_ids[mb];
|
||||
const scalar_t* __restrict__ B0 = packed_w1 + expert_id * stride_e + nb_upper * BLOCK_N * stride_n;
|
||||
const scalar_t* __restrict__ B1 = packed_w1 + expert_id * stride_e + nb_lower * BLOCK_N * stride_n;
|
||||
const float* __restrict__ B0_bias = w1_bias + expert_id * 2 * N + nb_upper * BLOCK_N;
|
||||
const float* __restrict__ B1_bias = w1_bias + expert_id * 2 * N + nb_lower * BLOCK_N;
|
||||
|
||||
int64_t m_size = offsets[mb + 1] - offsets[mb];
|
||||
|
||||
@@ -533,23 +591,58 @@ void fused_experts_kernel_impl(
|
||||
/* B */ B1,
|
||||
/* C */ C1);
|
||||
|
||||
// 1.d silu and mul
|
||||
const int64_t offset = offsets[mb];
|
||||
silu_and_mul<scalar_t, BLOCK_N>(ic1 + offset * N + nb * BLOCK_N, C0, C1, m_size, N);
|
||||
} else {
|
||||
// fused 1.bcd: silu_and_mul(A @ B0, A @ B1)
|
||||
const int64_t offset = offsets[mb];
|
||||
tinygemm_kernel(
|
||||
/* A */ A,
|
||||
/* B0 */ B0,
|
||||
/* B1 */ B1,
|
||||
/* C */ ic1 + offset * N + nb * BLOCK_N,
|
||||
/* M */ m_size,
|
||||
/* N */ n_size,
|
||||
/* K */ K,
|
||||
/* lda */ K,
|
||||
/* ldb */ n_size,
|
||||
/* ldc */ N);
|
||||
if (act_func == CPUActMethod::swiglu) {
|
||||
tinygemm_kernel(
|
||||
/* A */ A,
|
||||
/* B */ B0,
|
||||
/* C */ C0,
|
||||
/* M */ m_size,
|
||||
/* N */ n_size,
|
||||
/* K */ K,
|
||||
/* lda */ K,
|
||||
/* ldb */ n_size,
|
||||
/* ldc */ BLOCK_N);
|
||||
tinygemm_kernel(
|
||||
/* A */ A,
|
||||
/* B */ B1,
|
||||
/* C */ C1,
|
||||
/* M */ m_size,
|
||||
/* N */ n_size,
|
||||
/* K */ K,
|
||||
/* lda */ K,
|
||||
/* ldb */ n_size,
|
||||
/* ldc */ BLOCK_N);
|
||||
} else {
|
||||
// fused 1.bcd: silu_and_mul(A @ B0, A @ B1)
|
||||
tinygemm_kernel(
|
||||
/* A */ A,
|
||||
/* B0 */ B0,
|
||||
/* B1 */ B1,
|
||||
/* C */ ic1 + offset * N + nb * BLOCK_N,
|
||||
/* M */ m_size,
|
||||
/* N */ n_size,
|
||||
/* K */ K,
|
||||
/* lda */ K,
|
||||
/* ldb */ n_size,
|
||||
/* ldc */ N);
|
||||
}
|
||||
}
|
||||
if (with_bias) {
|
||||
for (int64_t m = 0; m < m_size; ++m) {
|
||||
add_bias_stub(C0 + m * BLOCK_N, B0_bias, n_size);
|
||||
add_bias_stub(C1 + m * BLOCK_N, B1_bias, n_size);
|
||||
}
|
||||
}
|
||||
// 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);
|
||||
} 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);
|
||||
}
|
||||
});
|
||||
|
||||
@@ -586,6 +679,7 @@ void fused_experts_kernel_impl(
|
||||
// B shape [IC, n_size] in vnni format
|
||||
int32_t expert_id = expert_ids[mb];
|
||||
const scalar_t* __restrict__ B = packed_w2 + expert_id * stride_e2 + nb * BLOCK_N * stride_oc;
|
||||
const float* __restrict__ B_bias = w2_bias + expert_id * OC + nb * BLOCK_N;
|
||||
|
||||
// 2.a gemm: C = A @ B
|
||||
if (use_brgemm) {
|
||||
@@ -613,6 +707,11 @@ void fused_experts_kernel_impl(
|
||||
/* ldc */ BLOCK_N);
|
||||
}
|
||||
|
||||
if (with_bias) {
|
||||
for (int64_t m = 0; m < m_size; ++m) {
|
||||
add_bias_stub(C + m * BLOCK_N, B_bias, n_size);
|
||||
}
|
||||
}
|
||||
// 2.b copy from C to ic2 in original order
|
||||
// and also mul topk_weights in float32
|
||||
for (int64_t m = 0; m < m_size; ++m) {
|
||||
@@ -804,21 +903,38 @@ void shared_expert_kernel_impl(
|
||||
} // anonymous namespace
|
||||
|
||||
// common checks
|
||||
template <CPUQuantMethod quant>
|
||||
static inline void check_moe_scales(
|
||||
bool use_int8_w8a8,
|
||||
bool use_fp8_w8a16,
|
||||
const std::optional<at::Tensor>& w1_scale,
|
||||
const std::optional<at::Tensor>& w2_scale,
|
||||
const std::optional<std::vector<int64_t>> block_size) {
|
||||
if (use_int8_w8a8) {
|
||||
if constexpr (quant == CPUQuantMethod::INT8_W8A8) {
|
||||
TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for int8 w8a8.");
|
||||
TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for int8 w8a8.");
|
||||
}
|
||||
if (use_fp8_w8a16) {
|
||||
} else if constexpr (quant == CPUQuantMethod::FP8_W8A16) {
|
||||
TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for fp8 w8a16.");
|
||||
TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for fp8 w8a16.");
|
||||
TORCH_CHECK(block_size.has_value(), "missing block_size for fp8 w8a16.");
|
||||
TORCH_CHECK(block_size.value().size() == 2, "expect block_size.size() to be 2.");
|
||||
} else if constexpr (quant == CPUQuantMethod::MXFP4) {
|
||||
TORCH_CHECK(w1_scale.has_value(), "missing w1_scale for mxfp4.");
|
||||
TORCH_CHECK(w2_scale.has_value(), "missing w2_scale for mxfp4.");
|
||||
TORCH_CHECK(w1_scale.value().scalar_type() == at::kByte, "expect w1_scale to be uint8.");
|
||||
TORCH_CHECK(w2_scale.value().scalar_type() == at::kByte, "expect w2_scale to be uint8.");
|
||||
}
|
||||
}
|
||||
|
||||
static inline void check_moe_scales(
|
||||
int64_t moe_comp_method,
|
||||
const std::optional<at::Tensor>& w1_scale,
|
||||
const std::optional<at::Tensor>& w2_scale,
|
||||
const std::optional<std::vector<int64_t>> block_size) {
|
||||
if (moe_comp_method == CPUQuantMethod::INT8_W8A8) {
|
||||
check_moe_scales<CPUQuantMethod::INT8_W8A8>(w1_scale, w2_scale, block_size);
|
||||
} else if (moe_comp_method == CPUQuantMethod::FP8_W8A16) {
|
||||
check_moe_scales<CPUQuantMethod::FP8_W8A16>(w1_scale, w2_scale, block_size);
|
||||
} else if (moe_comp_method == CPUQuantMethod::MXFP4) {
|
||||
check_moe_scales<CPUQuantMethod::MXFP4>(w1_scale, w2_scale, block_size);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -834,8 +950,8 @@ static inline void check_moe_scales(
|
||||
TORCH_CHECK(w2s.size(DIM1) == div_up(N, block_size_K))
|
||||
|
||||
// hidden_states: [M, K]
|
||||
// w1: [E, 2N, K]
|
||||
// w2: [E, K, N]
|
||||
// w1: [E, 2N, K] or [E, 2N, K / 2] for uint8
|
||||
// w2: [E, K, N] or [E, K, N / 2] for uint8
|
||||
// topk_weights: [M, topk]
|
||||
// topk_ids: [M, topk] (int32_t)
|
||||
//
|
||||
@@ -853,6 +969,10 @@ at::Tensor fused_experts_cpu(
|
||||
const std::optional<at::Tensor>& w1_zero,
|
||||
const std::optional<at::Tensor>& w2_zero,
|
||||
const std::optional<std::vector<int64_t>> block_size,
|
||||
const std::optional<at::Tensor>& w1_bias,
|
||||
const std::optional<at::Tensor>& w2_bias,
|
||||
const std::optional<double>& alpha,
|
||||
const std::optional<double>& limit,
|
||||
bool is_vnni) {
|
||||
auto packed_w1 = is_vnni ? w1 : convert_weight_packed(w1);
|
||||
auto packed_w2 = is_vnni ? w2 : convert_weight_packed(w2);
|
||||
@@ -895,8 +1015,8 @@ at::Tensor fused_experts_cpu(
|
||||
int64_t topk = topk_weights_.size(1);
|
||||
|
||||
// we use int32_t compensation for int8 w8a8
|
||||
int64_t packed_K = get_row_size(K, moe_comp_method == CPUQuantMethod::INT8_W8A8);
|
||||
int64_t packed_N = get_row_size(N, moe_comp_method == CPUQuantMethod::INT8_W8A8);
|
||||
int64_t packed_K = get_row_size(static_cast<CPUQuantMethod>(moe_comp_method), K);
|
||||
int64_t packed_N = get_row_size(static_cast<CPUQuantMethod>(moe_comp_method), N);
|
||||
|
||||
// check weight shapes
|
||||
CHECK_EQ(w2.size(0), E);
|
||||
@@ -906,12 +1026,7 @@ at::Tensor fused_experts_cpu(
|
||||
CHECK_EQ(packed_w2.size(2), packed_N / (moe_comp_method == CPUQuantMethod::INT4_W4A8 ? 2 : 1));
|
||||
}
|
||||
// check scales
|
||||
check_moe_scales(
|
||||
moe_comp_method == CPUQuantMethod::INT8_W8A8,
|
||||
moe_comp_method == CPUQuantMethod::FP8_W8A16,
|
||||
w1_scale,
|
||||
w2_scale,
|
||||
block_size);
|
||||
check_moe_scales(moe_comp_method, w1_scale, w2_scale, block_size);
|
||||
|
||||
at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states);
|
||||
|
||||
@@ -963,7 +1078,7 @@ at::Tensor fused_experts_cpu(
|
||||
// 5. Aq_tmp : [M, K] or [M * topk, N]
|
||||
// 6. As_tmp : [M * topk]
|
||||
//
|
||||
// for fp8 w8a16:
|
||||
// for fp8 w8a16 and mxfp4:
|
||||
// 7. intermediate_cache0 : [M * topk, 2N]
|
||||
// 8. B_tmp : [T, MAX_CACHE_BLOCK_SIZE, BLOCK_N, std::max(K, N)]
|
||||
//
|
||||
@@ -976,7 +1091,7 @@ at::Tensor fused_experts_cpu(
|
||||
if (moe_comp_method == CPUQuantMethod::INT8_W8A8) {
|
||||
buffer_size_nbytes += std::max(M * K, M * topk * N) + M * topk * sizeof(float);
|
||||
}
|
||||
if (moe_comp_method == CPUQuantMethod::FP8_W8A16) {
|
||||
if (moe_comp_method == CPUQuantMethod::FP8_W8A16 || moe_comp_method == CPUQuantMethod::MXFP4) {
|
||||
buffer_size_nbytes += M * topk * 2 * N * 2 + num_threads * MAX_CACHE_BLOCK_SIZE * BLOCK_N * std::max(K, N) * 2;
|
||||
}
|
||||
if (moe_comp_method == CPUQuantMethod::INT4_W4A8) {
|
||||
@@ -1029,9 +1144,11 @@ at::Tensor fused_experts_cpu(
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
scalar_t* __restrict__ intermediate_cache0 = (scalar_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N));
|
||||
scalar_t* __restrict__ B_tmp = (scalar_t*)((void*)(intermediate_cache0 + M * topk * 2 * N));
|
||||
bool with_bias = w1_bias.has_value();
|
||||
auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul;
|
||||
|
||||
CHECK_MOE_SCALES_FP8(1, 2);
|
||||
fused_experts_fp8_kernel_impl(
|
||||
fused_experts_fp_kernel_impl<scalar_t, at::Float8_e4m3fn, float, false>(
|
||||
out_hidden_states.data_ptr<scalar_t>(),
|
||||
intermediate_cache0,
|
||||
intermediate_cache1,
|
||||
@@ -1042,6 +1159,8 @@ at::Tensor fused_experts_cpu(
|
||||
hidden_states.data_ptr<scalar_t>(),
|
||||
packed_w1.data_ptr<at::Float8_e4m3fn>(),
|
||||
packed_w2.data_ptr<at::Float8_e4m3fn>(),
|
||||
with_bias ? w1_bias.value().data_ptr<float>() : nullptr,
|
||||
with_bias ? w2_bias.value().data_ptr<float>() : nullptr,
|
||||
w1s.data_ptr<float>(),
|
||||
w2s.data_ptr<float>(),
|
||||
block_size_N,
|
||||
@@ -1055,7 +1174,56 @@ at::Tensor fused_experts_cpu(
|
||||
K,
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad);
|
||||
num_tokens_post_pad,
|
||||
alpha.has_value() ? float(alpha.value()) : 0,
|
||||
limit.has_value() ? float(limit.value()) : 0,
|
||||
act_func,
|
||||
with_bias);
|
||||
} else if (moe_comp_method == CPUQuantMethod::MXFP4) {
|
||||
scalar_t* __restrict__ A_tmp = (scalar_t*)((void*)(intermediate_cache2 + M * topk * K));
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
scalar_t* __restrict__ intermediate_cache0 = (scalar_t*)((void*)(C_tmp + num_threads * 2 * BLOCK_M * BLOCK_N));
|
||||
scalar_t* __restrict__ B_tmp = (scalar_t*)((void*)(intermediate_cache0 + M * topk * 2 * N));
|
||||
bool with_bias = w1_bias.has_value();
|
||||
auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul;
|
||||
|
||||
// mxfp4 supports only group size of 32 (2^5)
|
||||
constexpr int64_t group_size = 32;
|
||||
auto w1s = w1_scale.value();
|
||||
auto w2s = w2_scale.value();
|
||||
TORCH_CHECK(w1s.numel(), E * 2 * N * K >> 5);
|
||||
TORCH_CHECK(w2s.numel(), E * K * N >> 5);
|
||||
fused_experts_fp_kernel_impl<scalar_t, uint8_t, uint8_t, true>(
|
||||
out_hidden_states.data_ptr<scalar_t>(),
|
||||
intermediate_cache0,
|
||||
intermediate_cache1,
|
||||
intermediate_cache2,
|
||||
A_tmp,
|
||||
B_tmp,
|
||||
C_tmp,
|
||||
hidden_states.data_ptr<scalar_t>(),
|
||||
packed_w1.data_ptr<uint8_t>(),
|
||||
packed_w2.data_ptr<uint8_t>(),
|
||||
with_bias ? w1_bias.value().data_ptr<float>() : nullptr,
|
||||
with_bias ? w2_bias.value().data_ptr<float>() : nullptr,
|
||||
w1s.data_ptr<uint8_t>(),
|
||||
w2s.data_ptr<uint8_t>(),
|
||||
/*block_size_N*/ 1,
|
||||
/*block_size_K*/ group_size,
|
||||
topk_weights_.data_ptr<float>(),
|
||||
sorted_ids,
|
||||
expert_ids,
|
||||
offsets,
|
||||
M,
|
||||
N,
|
||||
K,
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad,
|
||||
alpha.has_value() ? float(alpha.value()) : 0,
|
||||
limit.has_value() ? float(limit.value()) : 0,
|
||||
act_func,
|
||||
with_bias);
|
||||
} else if (moe_comp_method == CPUQuantMethod::INT4_W4A8) {
|
||||
uint8_t* __restrict__ A_tmp = (uint8_t*)((void*)(intermediate_cache2 + M * topk * K));
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
@@ -1101,6 +1269,8 @@ at::Tensor fused_experts_cpu(
|
||||
} else {
|
||||
scalar_t* __restrict__ A_tmp = intermediate_cache2 + M * topk * K;
|
||||
float* __restrict__ C_tmp = (float*)((void*)(A_tmp + num_threads * BLOCK_M * K));
|
||||
bool with_bias = w1_bias.has_value();
|
||||
auto act_func = alpha.has_value() && limit.has_value() ? CPUActMethod::swiglu : CPUActMethod::silu_and_mul;
|
||||
|
||||
fused_experts_kernel_impl<scalar_t>(
|
||||
out_hidden_states.data_ptr<scalar_t>(),
|
||||
@@ -1111,6 +1281,8 @@ at::Tensor fused_experts_cpu(
|
||||
hidden_states.data_ptr<scalar_t>(),
|
||||
packed_w1.data_ptr<scalar_t>(),
|
||||
packed_w2.data_ptr<scalar_t>(),
|
||||
with_bias ? w1_bias.value().data_ptr<float>() : nullptr,
|
||||
with_bias ? w2_bias.value().data_ptr<float>() : nullptr,
|
||||
topk_weights_.data_ptr<float>(),
|
||||
sorted_ids,
|
||||
expert_ids,
|
||||
@@ -1120,7 +1292,11 @@ at::Tensor fused_experts_cpu(
|
||||
K,
|
||||
E,
|
||||
topk,
|
||||
num_tokens_post_pad);
|
||||
num_tokens_post_pad,
|
||||
alpha.has_value() ? float(alpha.value()) : 0,
|
||||
limit.has_value() ? float(limit.value()) : 0,
|
||||
act_func,
|
||||
with_bias);
|
||||
}
|
||||
});
|
||||
return out_hidden_states;
|
||||
@@ -1183,7 +1359,11 @@ at::Tensor shared_expert_cpu(
|
||||
CHECK_EQ(packed_w2.size(1), packed_N);
|
||||
|
||||
// check scales
|
||||
check_moe_scales(use_int8_w8a8, use_fp8_w8a16, w1_scale, w2_scale, block_size);
|
||||
if (use_int8_w8a8) {
|
||||
check_moe_scales<CPUQuantMethod::INT8_W8A8>(w1_scale, w2_scale, block_size);
|
||||
} else if (use_fp8_w8a16) {
|
||||
check_moe_scales<CPUQuantMethod::FP8_W8A16>(w1_scale, w2_scale, block_size);
|
||||
}
|
||||
|
||||
at::Tensor out_hidden_states = inplace ? hidden_states : at::empty_like(hidden_states);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user