[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
+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>(),