[CPU] Add Qwen3.5 model optimization for CPU (#19484)

Co-authored-by: Zheng, Beilei <beilei.zheng@intel.com>
Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
jianan-gu
2026-04-26 10:12:36 -07:00
committed by GitHub
co-authored by Zheng, Beilei Ma Mingfei Xinyuan Tong
parent 7d49564431
commit 10fd0faccd
20 changed files with 768 additions and 209 deletions
+50
View File
@@ -97,6 +97,43 @@ namespace {
TORCH_CHECK(false, "Unsupported floating data type."); \
}
// Helper MICRO for CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT:
// TYPE1: the primary dtype (input, output, weight);
// TYPE2: defined as PARAM_T input
#define CPU_DISPATCH_TYPE1_WITH_PARAM_REDUCED(TYPE1, PARAM_T, ...) \
switch (TYPE1) { \
case at::ScalarType::BFloat16: { \
using scalar_t = at::BFloat16; \
using param_t = PARAM_T; \
return __VA_ARGS__(); \
} \
case at::ScalarType::Half: { \
using scalar_t = at::Half; \
using param_t = PARAM_T; \
return __VA_ARGS__(); \
} \
default: \
TORCH_CHECK(false, "Unsupported floating data type."); \
}
// Helper MICRO for CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT:
// TYPE1: the dtype both for scalar_t and param_t
#define CPU_DISPATCH_TYPE1_WITH_SAME_PARAM_REDUCED(TYPE1, ...) \
switch (TYPE1) { \
case at::ScalarType::BFloat16: { \
using scalar_t = at::BFloat16; \
using param_t = at::BFloat16; \
return __VA_ARGS__(); \
} \
case at::ScalarType::Half: { \
using scalar_t = at::Half; \
using param_t = at::Half; \
return __VA_ARGS__(); \
} \
default: \
TORCH_CHECK(false, "Unsupported reduced floating data type."); \
}
// dispatch with mixed dtypes (TYPE1, TYPE2):
// TYPE1: the primary dtype (input, output, weight);
// TYPE2: the secondary dtype (bias, etc.).
@@ -113,6 +150,19 @@ namespace {
} \
}()
// dispatch with mixed dtypes (reduced one, no float for TYPE1) (TYPE1, TYPE2):
// TYPE1: the primary dtype (input, output, weight);
// TYPE2: the secondary dtype (bias, etc.).
#define CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(TYPE1, TYPE2, ...) \
[&] { \
if (TYPE2 == at::kFloat) { \
CPU_DISPATCH_TYPE1_WITH_PARAM_REDUCED(TYPE1, float, __VA_ARGS__) \
} else { \
TORCH_CHECK(TYPE1 == TYPE2); \
CPU_DISPATCH_TYPE1_WITH_SAME_PARAM_REDUCED(TYPE1, __VA_ARGS__) \
} \
}()
#define UNUSED(x) (void)(x)
#define CHECK_CPU(x) TORCH_CHECK(x.device().type() == at::kCPU, #x " must be a CPU tensor")
+93 -37
View File
@@ -814,12 +814,12 @@ inline at::vec::Vectorized<float> softplus(const at::vec::Vectorized<float>& x,
return Vec::blendv(Vec::blendv(log1pex, expx, mask_lo), x, mask_hi);
}
template <typename scalar_t>
template <typename scalar_t, typename param_t>
void fused_sigmoid_gating_delta_rule_update_kernel_impl(
const scalar_t* __restrict__ q_ptr,
const scalar_t* __restrict__ k_ptr,
const scalar_t* __restrict__ v_ptr,
const float* __restrict__ A_log_ptr,
const param_t* __restrict__ A_log_ptr,
const scalar_t* __restrict__ a_ptr,
const scalar_t* __restrict__ dt_bias_ptr,
const scalar_t* __restrict__ b_ptr,
@@ -903,7 +903,7 @@ void fused_sigmoid_gating_delta_rule_update_kernel_impl(
for (int64_t i = begin; i < end; ++i) {
int64_t cache_index = indices_ptr[bi];
int64_t state_offset = (cache_index * v_num_heads + ni) * head_dim * v_head_dim;
float g_val = -std::exp(A_log_ptr[ni]) *
float g_val = -std::exp(float(A_log_ptr[ni])) *
softplus(float(a_ptr[bi * v_num_heads + ni]) + float(dt_bias_ptr[ni]), softplus_threshold);
float g_val_exp = std::exp(g_val);
fVec g_val_exp_vec = fVec(g_val_exp);
@@ -1021,6 +1021,55 @@ void fused_gdn_gating_kernel_impl(
});
}
template <typename scalar_t>
void fused_gdn_gating_kernel_impl(
scalar_t* __restrict__ A_log,
const scalar_t* __restrict__ a,
const scalar_t* __restrict__ b,
const scalar_t* __restrict__ dt_bias,
float* __restrict__ out,
scalar_t* __restrict__ beta,
int64_t batch,
int64_t num_heads) {
using bVec = at::vec::Vectorized<scalar_t>;
using fVec = at::vec::Vectorized<float>;
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);
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());
g0.store(out + i * num_heads + j);
g1.store(out + i * num_heads + j + fvec_size);
bVec beta_vec = at::vec::convert_from_float<scalar_t>(beta0, beta1);
beta_vec.store(beta + i * num_heads + j);
}
for (; j < num_heads; ++j) {
out[i * num_heads + j] = -std::exp(float(A_log[j])) * softplus(float(a[i * num_heads + j]) + float(dt_bias[j]));
beta[i * num_heads + j] = 1 / (1 + std::exp(-b[i * num_heads + j]));
}
}
});
}
} // anonymous namespace
template <bool is_last_dim_contiguous>
@@ -1242,7 +1291,6 @@ at::Tensor fused_sigmoid_gating_delta_rule_update_cpu(
int64_t v_head_dim = v.size(3);
CHECK_INPUT_SHAPE_DTYPE<true>(k, {seq_len, batch_size, num_heads, head_dim}, q.scalar_type());
CHECK_INPUT_SHAPE_DTYPE<true>(v, {seq_len, batch_size, v_num_heads, v_head_dim}, q.scalar_type());
CHECK_INPUT_SHAPE_DTYPE<true>(A_log, {v_num_heads}, at::kFloat);
CHECK_INPUT_SHAPE_DTYPE<true>(a, {batch_size, v_num_heads}, q.scalar_type());
CHECK_INPUT_SHAPE_DTYPE<true>(dt_bias, {v_num_heads}, q.scalar_type());
CHECK_INPUT_SHAPE_DTYPE<true>(b, {batch_size, v_num_heads}, q.scalar_type());
@@ -1252,6 +1300,12 @@ at::Tensor fused_sigmoid_gating_delta_rule_update_cpu(
initial_state_source, {initial_state_source.size(0), v_num_heads, head_dim, v_head_dim}, at::kFloat);
CHECK(initial_state_source.size(0) >= batch_size);
CHECK_EQ(v_num_heads % num_heads, 0);
TORCH_CHECK(
A_log.sizes() == at::IntArrayRef({v_num_heads}),
"Input tensor shape mismatch: expected ",
at::IntArrayRef({v_num_heads}),
", got ",
A_log.sizes());
int64_t q_strideB = q.stride(1);
int64_t q_strideS = q.stride(0);
@@ -1264,37 +1318,39 @@ at::Tensor fused_sigmoid_gating_delta_rule_update_cpu(
int64_t v_strideH = v.stride(2);
at::Tensor core_attn_out = at::empty({batch_size, seq_len, v_num_heads, v_head_dim}, q.options());
at::Tensor qk_scale_buf = at::empty({2 * batch_size, seq_len, num_heads}, at::kFloat);
AT_DISPATCH_REDUCED_FLOATING_TYPES(q.scalar_type(), "fused_sigmoid_gating_delta_rule_update_kernel_impl", [&] {
fused_sigmoid_gating_delta_rule_update_kernel_impl<scalar_t>(
q.data_ptr<scalar_t>(),
k.data_ptr<scalar_t>(),
v.data_ptr<scalar_t>(),
A_log.data_ptr<float>(),
a.data_ptr<scalar_t>(),
dt_bias.data_ptr<scalar_t>(),
b.data_ptr<scalar_t>(),
initial_state_indices.data_ptr<int32_t>(),
initial_state_source.data_ptr<float>(),
core_attn_out.data_ptr<scalar_t>(),
qk_scale_buf.data_ptr<float>(),
seq_len,
batch_size,
num_heads,
head_dim,
v_num_heads,
v_head_dim,
q_strideB,
q_strideS,
q_strideH,
k_strideB,
k_strideS,
k_strideH,
v_strideB,
v_strideS,
v_strideH,
use_qk_l2norm_in_kernel,
softplus_threshold);
});
CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(
q.scalar_type(), A_log.scalar_type(), "fused_sigmoid_gating_delta_rule_update_kernel_impl", [&] {
fused_sigmoid_gating_delta_rule_update_kernel_impl<scalar_t, param_t>(
q.data_ptr<scalar_t>(),
k.data_ptr<scalar_t>(),
v.data_ptr<scalar_t>(),
A_log.data_ptr<param_t>(),
a.data_ptr<scalar_t>(),
dt_bias.data_ptr<scalar_t>(),
b.data_ptr<scalar_t>(),
initial_state_indices.data_ptr<int32_t>(),
initial_state_source.data_ptr<float>(),
core_attn_out.data_ptr<scalar_t>(),
qk_scale_buf.data_ptr<float>(),
seq_len,
batch_size,
num_heads,
head_dim,
v_num_heads,
v_head_dim,
q_strideB,
q_strideS,
q_strideH,
k_strideB,
k_strideS,
k_strideH,
v_strideB,
v_strideS,
v_strideH,
use_qk_l2norm_in_kernel,
softplus_threshold);
});
return core_attn_out;
}
@@ -1318,9 +1374,9 @@ fused_gdn_gating_cpu(const at::Tensor& A_log, const at::Tensor& a, const at::Ten
CHECK_EQ(b.size(1), num_heads);
at::Tensor out = at::empty({1, batch, num_heads}, a.options().dtype(at::kFloat));
at::Tensor beta = at::empty({1, batch, num_heads}, b.options());
AT_DISPATCH_REDUCED_FLOATING_TYPES(a.scalar_type(), "fused_gdn_gating_kernel", [&] {
CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(a.scalar_type(), A_log.scalar_type(), "fused_gdn_gating_kernel", [&] {
fused_gdn_gating_kernel_impl<scalar_t>(
A_log.data_ptr<float>(),
A_log.data_ptr<param_t>(),
a.data_ptr<scalar_t>(),
b.data_ptr<scalar_t>(),
dt_bias.data_ptr<scalar_t>(),
+86
View File
@@ -61,6 +61,41 @@ void fused_qkvzba_split_reshape_cat_impl(
}
});
}
template <typename scalar_t>
void fused_qkvzba_split_reshape_cat_contiguous_impl(
const scalar_t* __restrict__ mixed_qkvz,
const scalar_t* __restrict__ mixed_ba,
scalar_t* __restrict__ mixed_qkv,
scalar_t* __restrict__ z,
scalar_t* __restrict__ b,
scalar_t* __restrict__ a,
int64_t batch,
int64_t k_tp,
int64_t v_tp,
int64_t num_heads_v,
int64_t qkv_dim,
int64_t qkv_strideB,
int64_t qkvz_strideB,
int64_t ba_strideB) {
at::parallel_for(0, batch, 0, [&](int64_t begin, int64_t end) {
for (int64_t bi = begin; bi < end; ++bi) {
scalar_t* __restrict__ qkv_out_ptr = mixed_qkv + bi * qkv_strideB;
const scalar_t* __restrict__ qkv_in_ptr = mixed_qkvz + bi * qkvz_strideB;
scalar_t* __restrict__ z_out_ptr = z + bi * v_tp;
const scalar_t* __restrict__ z_in_ptr = qkv_in_ptr + qkv_dim;
copy_stub(qkv_out_ptr, qkv_in_ptr, qkv_dim);
copy_stub(z_out_ptr, z_in_ptr, v_tp);
scalar_t* __restrict__ b_out_ptr = b + bi * num_heads_v;
const scalar_t* __restrict__ b_in_ptr = mixed_ba + bi * ba_strideB;
scalar_t* __restrict__ a_out_ptr = a + bi * num_heads_v;
const scalar_t* __restrict__ a_in_ptr = b_in_ptr + num_heads_v;
copy_stub(b_out_ptr, b_in_ptr, num_heads_v);
copy_stub(a_out_ptr, a_in_ptr, num_heads_v);
}
});
}
} // anonymous namespace
// mixed_qkvz: [batch, num_heads_qk * head_qk * 2 + num_heads_v * head_v * 2]
@@ -83,6 +118,7 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> fused_qkvzba_split_re
CHECK_EQ(mixed_qkvz.size(1), expected_dim);
CHECK_EQ(mixed_ba.size(0), batch);
CHECK_EQ(mixed_ba.size(1), ba_dim);
TORCH_CHECK(mixed_ba.scalar_type() == mixed_qkvz.scalar_type(), "mixed_ba and mixed_qkvz must share same dtype");
CHECK_EQ(num_heads_v % num_heads_qk, 0);
at::Tensor mixed_qkv = at::empty({batch, qkv_dim}, mixed_qkvz.options());
at::Tensor z = at::empty({batch, num_heads_v, head_v}, mixed_qkvz.options());
@@ -112,3 +148,53 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> fused_qkvzba_split_re
});
return std::make_tuple(mixed_qkv, z, b, a);
}
// mixed_qkvz: [batch, num_heads_qk * head_qk * 2 + num_heads_v * head_v * 2]
// mixed_ba: [batch, num_heads_v * 2]
std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> fused_qkvzba_split_reshape_cat_contiguous_cpu(
const at::Tensor& mixed_qkvz,
const at::Tensor& mixed_ba,
int64_t num_heads_qk,
int64_t num_heads_v,
int64_t head_qk,
int64_t head_v) {
CHECK_DIM(2, mixed_qkvz);
CHECK_DIM(2, mixed_ba);
CHECK_INPUT(mixed_qkvz);
CHECK_INPUT(mixed_ba);
int64_t batch = mixed_qkvz.size(0);
int64_t k_tp = num_heads_qk * head_qk;
int64_t v_tp = num_heads_v * head_v;
int64_t qkv_dim = k_tp * 2 + v_tp;
int64_t ba_dim = num_heads_v * 2;
int64_t expected_dim = qkv_dim + v_tp;
CHECK_EQ(mixed_qkvz.size(1), expected_dim);
CHECK_EQ(mixed_ba.size(0), batch);
CHECK_EQ(mixed_ba.size(1), ba_dim);
TORCH_CHECK(mixed_ba.scalar_type() == mixed_qkvz.scalar_type(), "mixed_ba and mixed_qkvz must share same dtype");
at::Tensor mixed_qkv = at::empty({batch, qkv_dim}, mixed_qkvz.options());
at::Tensor z = at::empty({batch, num_heads_v, head_v}, mixed_qkvz.options());
at::Tensor b = at::empty({batch, num_heads_v}, mixed_ba.options());
at::Tensor a = at::empty({batch, num_heads_v}, mixed_ba.options());
int64_t qkvz_strideB = mixed_qkvz.size(1);
int64_t qkv_strideB = mixed_qkv.size(1);
int64_t ba_strideB = mixed_ba.size(1);
AT_DISPATCH_REDUCED_FLOATING_TYPES(mixed_qkvz.scalar_type(), "fused_qkvzba_split_reshape_cat_contiguous_impl", [&] {
fused_qkvzba_split_reshape_cat_contiguous_impl<scalar_t>(
mixed_qkvz.data_ptr<scalar_t>(),
mixed_ba.data_ptr<scalar_t>(),
mixed_qkv.data_ptr<scalar_t>(),
z.data_ptr<scalar_t>(),
b.data_ptr<scalar_t>(),
a.data_ptr<scalar_t>(),
batch,
k_tp,
v_tp,
num_heads_v,
qkv_dim,
qkv_strideB,
qkvz_strideB,
ba_strideB);
});
return std::make_tuple(mixed_qkv, z, b, a);
}
+7 -5
View File
@@ -495,13 +495,15 @@ void fused_experts_kernel_impl(
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;
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
int64_t m_size = offsets[mb + 1] - offsets[mb];
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, input + index * K, K);
if (nb_offset == 0) {
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, input + index * K, K);
}
}
if (use_brgemm) {
+7 -5
View File
@@ -65,13 +65,15 @@ void fused_experts_fp8_kernel_impl(
int32_t pre_expert_id = mb == 0 ? -1 : expert_ids[mb - 1];
bool do_unpack = (mb == mb0) || (expert_id != pre_expert_id);
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
int64_t m_size = offsets[mb + 1] - offsets[mb];
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, input + index * K, K);
if (nb_offset == 0) {
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, input + index * K, K);
}
}
const int64_t offset = offsets[mb];
+8 -6
View File
@@ -550,14 +550,16 @@ void fused_experts_int8_kernel_impl(
const float* __restrict__ Bs0 = w1s + expert_id * 2 * N + nb_upper * BLOCK_N;
const float* __restrict__ Bs1 = w1s + expert_id * 2 * N + nb_lower * BLOCK_N;
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
int64_t m_size = offsets[mb + 1] - offsets[mb];
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, Aq_tmp + index * K, K);
As[m] = As_tmp[index];
if (nb_offset == 0) {
// 1.a load A
const int32_t* A_ids = sorted_ids + mb * BLOCK_M;
for (int64_t m = 0; m < m_size; ++m) {
int32_t index = A_ids[m] / topk;
copy_stub(A + m * K, Aq_tmp + index * K, K);
As[m] = As_tmp[index];
}
}
if (use_brgemm) {
@@ -374,6 +374,15 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> fused_qkvzba_split_re
int64_t head_qk,
int64_t head_v);
// fused_qkvzba_split_reshape_cat_cpu_contiguous
std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> fused_qkvzba_split_reshape_cat_contiguous_cpu(
const at::Tensor& mixed_qkvz,
const at::Tensor& mixed_ba,
int64_t num_heads_qk,
int64_t num_heads_v,
int64_t head_qk,
int64_t head_v);
// image preprocessor
std::tuple<at::Tensor, at::Tensor> image_preprocess_cpu(
at::TensorList images,
@@ -621,6 +630,12 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
"fused_qkvzba_split_reshape_cat_cpu(Tensor mixed_qkvz, Tensor mixed_ba, int num_heads_qk, int num_heads_v, int "
"head_qk, int head_v) -> (Tensor, Tensor, Tensor, Tensor)");
m.impl("fused_qkvzba_split_reshape_cat_cpu", torch::kCPU, &fused_qkvzba_split_reshape_cat_cpu);
// fused_qkvzba_split_reshape_cat_contiguous_cpu
m.def(
"fused_qkvzba_split_reshape_cat_contiguous_cpu(Tensor mixed_qkvz, Tensor mixed_ba, int num_heads_qk, int "
"num_heads_v, int "
"head_qk, int head_v) -> (Tensor, Tensor, Tensor, Tensor)");
m.impl("fused_qkvzba_split_reshape_cat_contiguous_cpu", torch::kCPU, &fused_qkvzba_split_reshape_cat_contiguous_cpu);
// image preprocessor
m.def(