[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:
co-authored by
Zheng, Beilei
Ma Mingfei
Xinyuan Tong
parent
7d49564431
commit
10fd0faccd
@@ -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")
|
||||
|
||||
@@ -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>(),
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user