Optimize MiniMax-M2.7 on CPU (#31956)
This commit is contained in:
@@ -467,6 +467,167 @@ void fused_qk_norm4d_kernel_impl(
|
||||
});
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
float sum_squares(const scalar_t* __restrict__ input, int64_t size) {
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
constexpr int kVecSize = bVec::size();
|
||||
|
||||
fVec sum_fvec{0.f};
|
||||
float sum_val{0.f};
|
||||
int64_t i = 0;
|
||||
for (; i <= size - kVecSize; i += kVecSize) {
|
||||
auto [input_fvec0, input_fvec1] = load_float_vec2(input + i);
|
||||
sum_fvec += input_fvec0 * input_fvec0;
|
||||
sum_fvec += input_fvec1 * input_fvec1;
|
||||
}
|
||||
for (; i < size; ++i) {
|
||||
const float input_val = static_cast<float>(input[i]);
|
||||
sum_val += input_val * input_val;
|
||||
}
|
||||
return sum_val + vec_reduce_sum(sum_fvec);
|
||||
}
|
||||
|
||||
template <typename scalar_t>
|
||||
void fused_qk_norm_sumsq_kernel_impl(
|
||||
float* __restrict__ sum_sq,
|
||||
const scalar_t* __restrict__ q,
|
||||
const scalar_t* __restrict__ k,
|
||||
const NormParams& q_params,
|
||||
const NormParams& k_params) {
|
||||
at::parallel_for(0, q_params.B, 0, [&](int64_t begin, int64_t end) {
|
||||
for (int64_t b = begin; b < end; ++b) {
|
||||
sum_sq[b * 2] = sum_squares(q + q_params.input_offset(b, 0, 0), q_params.D);
|
||||
sum_sq[b * 2 + 1] = sum_squares(k + k_params.input_offset(b, 0, 0), k_params.D);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
template <NormMode M, typename scalar_t>
|
||||
void apply_norm_from_stats(
|
||||
scalar_t* __restrict__ output,
|
||||
const scalar_t* __restrict__ input,
|
||||
const NormParams& params,
|
||||
int64_t size,
|
||||
float sum,
|
||||
float sum_sq,
|
||||
int64_t tp_world_size) {
|
||||
using bVec = at::vec::Vectorized<scalar_t>;
|
||||
using fVec = at::vec::Vectorized<float>;
|
||||
constexpr int kVecSize = bVec::size();
|
||||
|
||||
const bool use_bias = params.bias != nullptr;
|
||||
float mean = 0.f;
|
||||
float variance = 0.f;
|
||||
const float global_size = static_cast<float>(size * tp_world_size);
|
||||
|
||||
if constexpr (NormTraits<M>::has_mean) {
|
||||
mean = sum / global_size;
|
||||
variance = (sum_sq / global_size) - (mean * mean);
|
||||
} else {
|
||||
variance = sum_sq / global_size;
|
||||
}
|
||||
const float scale = 1.f / std::sqrt(variance + params.eps);
|
||||
const fVec scale_fvec{scale};
|
||||
const fVec mean_fvec{mean};
|
||||
const fVec shift_fvec{params.shift};
|
||||
|
||||
int64_t i = 0;
|
||||
for (; i <= size - kVecSize; i += kVecSize) {
|
||||
auto [x_fvec0, x_fvec1] = load_float_vec2(input + i);
|
||||
if constexpr (NormTraits<M>::has_mean) {
|
||||
x_fvec0 = (x_fvec0 - mean_fvec) * scale_fvec;
|
||||
x_fvec1 = (x_fvec1 - mean_fvec) * scale_fvec;
|
||||
} else {
|
||||
x_fvec0 = x_fvec0 * scale_fvec;
|
||||
x_fvec1 = x_fvec1 * scale_fvec;
|
||||
}
|
||||
if constexpr (NormTraits<M>::has_weight) {
|
||||
auto [w_fvec0, w_fvec1] = load_float_vec2(static_cast<const scalar_t*>(params.weight) + i);
|
||||
if constexpr (NormTraits<M>::has_shift) {
|
||||
w_fvec0 = NormTraits<M>::apply_shift(w_fvec0, shift_fvec);
|
||||
w_fvec1 = NormTraits<M>::apply_shift(w_fvec1, shift_fvec);
|
||||
}
|
||||
x_fvec0 = NormTraits<M>::apply_weight(x_fvec0, w_fvec0);
|
||||
x_fvec1 = NormTraits<M>::apply_weight(x_fvec1, w_fvec1);
|
||||
}
|
||||
if constexpr (NormTraits<M>::has_bias) {
|
||||
if (use_bias) {
|
||||
auto [b_fvec0, b_fvec1] = load_float_vec2(static_cast<const scalar_t*>(params.bias) + i);
|
||||
x_fvec0 = NormTraits<M>::apply_bias(x_fvec0, b_fvec0);
|
||||
x_fvec1 = NormTraits<M>::apply_bias(x_fvec1, b_fvec1);
|
||||
}
|
||||
}
|
||||
|
||||
convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1).store(output + i);
|
||||
}
|
||||
|
||||
for (; i < size; ++i) {
|
||||
float x_val = static_cast<float>(input[i]);
|
||||
if constexpr (NormTraits<M>::has_mean) {
|
||||
x_val = (x_val - mean) * scale;
|
||||
} else {
|
||||
x_val = x_val * scale;
|
||||
}
|
||||
if constexpr (NormTraits<M>::has_weight) {
|
||||
float w_val = static_cast<float>(static_cast<const scalar_t*>(params.weight)[i]);
|
||||
if constexpr (NormTraits<M>::has_shift) {
|
||||
w_val = NormTraits<M>::apply_shift(w_val, params.shift);
|
||||
}
|
||||
x_val = NormTraits<M>::apply_weight(x_val, w_val);
|
||||
}
|
||||
if constexpr (NormTraits<M>::has_bias) {
|
||||
if (use_bias) {
|
||||
const float b_val = static_cast<float>(static_cast<const scalar_t*>(params.bias)[i]);
|
||||
x_val = NormTraits<M>::apply_bias(x_val, b_val);
|
||||
}
|
||||
}
|
||||
output[i] = static_cast<scalar_t>(x_val);
|
||||
}
|
||||
}
|
||||
|
||||
template <NormMode M, typename scalar_t>
|
||||
void fused_qk_norm_apply_from_stats_kernel_impl(
|
||||
scalar_t* __restrict__ q_out,
|
||||
scalar_t* __restrict__ k_out,
|
||||
const scalar_t* __restrict__ q,
|
||||
const scalar_t* __restrict__ k,
|
||||
const float* __restrict__ sum,
|
||||
const float* __restrict__ sum_sq,
|
||||
const NormParams& q_params,
|
||||
const NormParams& k_params,
|
||||
int64_t tp_world_size) {
|
||||
if constexpr (NormTraits<M>::has_mean) {
|
||||
TORCH_INTERNAL_ASSERT(sum != nullptr);
|
||||
}
|
||||
at::parallel_for(0, q_params.B, 0, [&](int64_t begin, int64_t end) {
|
||||
for (int64_t b = begin; b < end; ++b) {
|
||||
float q_sum = 0.f;
|
||||
float k_sum = 0.f;
|
||||
if constexpr (NormTraits<M>::has_mean) {
|
||||
q_sum = sum[b * 2];
|
||||
k_sum = sum[b * 2 + 1];
|
||||
}
|
||||
apply_norm_from_stats<M>(
|
||||
q_out + q_params.output_offset(b, 0, 0),
|
||||
q + q_params.input_offset(b, 0, 0),
|
||||
q_params,
|
||||
q_params.D,
|
||||
q_sum,
|
||||
sum_sq[b * 2],
|
||||
tp_world_size);
|
||||
apply_norm_from_stats<M>(
|
||||
k_out + k_params.output_offset(b, 0, 0),
|
||||
k + k_params.input_offset(b, 0, 0),
|
||||
k_params,
|
||||
k_params.D,
|
||||
k_sum,
|
||||
sum_sq[b * 2 + 1],
|
||||
tp_world_size);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#undef LAUNCH_PARALLEL_LOOP
|
||||
#undef LAUNCH_PARALLEL_LOOP_HD
|
||||
} // anonymous namespace
|
||||
@@ -694,6 +855,102 @@ at::Tensor fused_add_layernorm_cpu(
|
||||
return output;
|
||||
}
|
||||
|
||||
// q: {batch_size, q_hidden_size} 2D
|
||||
// k: {batch_size, k_hidden_size} 2D
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_rmsnorm_cpu(
|
||||
const at::Tensor& q, const at::Tensor& k, const at::Tensor& q_weight, const at::Tensor& k_weight, double eps) {
|
||||
const auto st = q.scalar_type();
|
||||
CHECK_INPUT_ND<2>(q);
|
||||
CHECK_INPUT_ND<2>(k);
|
||||
|
||||
CHECK_EQ(k.size(0), q.size(0));
|
||||
CHECK_EQ(k.scalar_type(), st);
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(q_weight, {q.size(1)}, st);
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(k_weight, {k.size(1)}, st);
|
||||
|
||||
NormParams q_params{q, static_cast<float>(eps)};
|
||||
q_params.weight = q_weight.data_ptr();
|
||||
|
||||
NormParams k_params{k, static_cast<float>(eps)};
|
||||
k_params.weight = k_weight.data_ptr();
|
||||
|
||||
at::Tensor q_out = at::empty_like(q);
|
||||
at::Tensor k_out = at::empty_like(k);
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "fused_qk_rmsnorm_kernel", [&] {
|
||||
fused_qk_norm4d_kernel_impl<NormMode::RMSNorm, scalar_t, false>(
|
||||
q_out.data_ptr<scalar_t>(),
|
||||
k_out.data_ptr<scalar_t>(),
|
||||
nullptr,
|
||||
q.data_ptr<scalar_t>(),
|
||||
k.data_ptr<scalar_t>(),
|
||||
q_params,
|
||||
k_params);
|
||||
});
|
||||
return std::make_tuple(q_out, k_out);
|
||||
}
|
||||
|
||||
// q: {batch_size, local_q_hidden_size} 2D
|
||||
// k: {batch_size, local_k_hidden_size} 2D
|
||||
// output: local Q/K squared sums, {batch_size, 2} FP32
|
||||
at::Tensor fused_qk_rmsnorm_sumsq_cpu(const at::Tensor& q, const at::Tensor& k) {
|
||||
const auto st = q.scalar_type();
|
||||
CHECK_INPUT_ND<2>(q);
|
||||
CHECK_INPUT_ND<2>(k);
|
||||
CHECK_EQ(k.size(0), q.size(0));
|
||||
CHECK_EQ(k.scalar_type(), st);
|
||||
|
||||
NormParams q_params{q, 0.f};
|
||||
NormParams k_params{k, 0.f};
|
||||
at::Tensor sum_sq = at::empty({q.size(0), 2}, q.options().dtype(at::kFloat));
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "fused_qk_rmsnorm_sumsq_kernel", [&] {
|
||||
fused_qk_norm_sumsq_kernel_impl<scalar_t>(
|
||||
sum_sq.data_ptr<float>(), q.data_ptr<scalar_t>(), k.data_ptr<scalar_t>(), q_params, k_params);
|
||||
});
|
||||
return sum_sq;
|
||||
}
|
||||
|
||||
// q: {batch_size, local_q_hidden_size} 2D
|
||||
// k: {batch_size, local_k_hidden_size} 2D
|
||||
// sum_sq: globally reduced Q/K squared sums, {batch_size, 2} FP32
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_rmsnorm_apply_from_stats_cpu(
|
||||
const at::Tensor& q,
|
||||
const at::Tensor& k,
|
||||
const at::Tensor& q_weight,
|
||||
const at::Tensor& k_weight,
|
||||
const at::Tensor& sum_sq,
|
||||
int64_t tp_world_size,
|
||||
double eps) {
|
||||
const auto st = q.scalar_type();
|
||||
CHECK_INPUT_ND<2>(q);
|
||||
CHECK_INPUT_ND<2>(k);
|
||||
CHECK_EQ(k.size(0), q.size(0));
|
||||
CHECK_EQ(k.scalar_type(), st);
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(q_weight, {q.size(1)}, st);
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(k_weight, {k.size(1)}, st);
|
||||
CHECK_INPUT_SHAPE_DTYPE<true>(sum_sq, {q.size(0), 2}, at::kFloat);
|
||||
TORCH_CHECK(tp_world_size > 0, "tp_world_size must be positive, got ", tp_world_size);
|
||||
|
||||
NormParams q_params{q, static_cast<float>(eps)};
|
||||
q_params.weight = q_weight.data_ptr();
|
||||
NormParams k_params{k, static_cast<float>(eps)};
|
||||
k_params.weight = k_weight.data_ptr();
|
||||
at::Tensor q_out = at::empty_like(q);
|
||||
at::Tensor k_out = at::empty_like(k);
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "fused_qk_rmsnorm_apply_kernel", [&] {
|
||||
fused_qk_norm_apply_from_stats_kernel_impl<NormMode::RMSNorm, scalar_t>(
|
||||
q_out.data_ptr<scalar_t>(),
|
||||
k_out.data_ptr<scalar_t>(),
|
||||
q.data_ptr<scalar_t>(),
|
||||
k.data_ptr<scalar_t>(),
|
||||
nullptr,
|
||||
sum_sq.data_ptr<float>(),
|
||||
q_params,
|
||||
k_params,
|
||||
tp_world_size);
|
||||
});
|
||||
return std::make_tuple(q_out, k_out);
|
||||
}
|
||||
|
||||
// q : {batch_size, num_head * head_dim} 2D
|
||||
// k : {batch_size, num_head_kv * head_dim} 2D
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
||||
|
||||
@@ -12,24 +12,39 @@ inline void softmax(float* __restrict__ out, const scalar_t* __restrict__ input)
|
||||
|
||||
// step 1: get max
|
||||
fVec max_fvec = fVec(-std::numeric_limits<float>::infinity());
|
||||
if constexpr (SIZE < kVecSize) {
|
||||
// SIZE = 1, 2, 4, 8, 16; only the top half is used
|
||||
bVec x_bvec = bVec::loadu(input, SIZE);
|
||||
fVec x_fvec0, x_fvec1;
|
||||
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
|
||||
x_fvec0 = fVec::set(max_fvec, x_fvec0, SIZE);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec0);
|
||||
x_fvec0.store(out, SIZE);
|
||||
if constexpr (std::is_same_v<scalar_t, float>) {
|
||||
if constexpr (SIZE < kVecSize) {
|
||||
fVec x_fvec = fVec::loadu(input, SIZE);
|
||||
x_fvec = fVec::set(max_fvec, x_fvec, SIZE);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec);
|
||||
x_fvec.store(out, SIZE);
|
||||
} else {
|
||||
for (int d = 0; d < SIZE; d += kVecSize) {
|
||||
fVec x_fvec = fVec::loadu(input + d);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec);
|
||||
x_fvec.store(out + d);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
for (int d = 0; d < SIZE; d += kVecSize) {
|
||||
bVec x_bvec = bVec::loadu(input + d);
|
||||
if constexpr (SIZE < kVecSize) {
|
||||
// SIZE = 1, 2, 4, 8, 16; only the top half is used
|
||||
bVec x_bvec = bVec::loadu(input, SIZE);
|
||||
fVec x_fvec0, x_fvec1;
|
||||
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
|
||||
|
||||
x_fvec0 = fVec::set(max_fvec, x_fvec0, SIZE);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec0);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec1);
|
||||
x_fvec0.store(out + d);
|
||||
x_fvec1.store(out + d + fVec::size());
|
||||
x_fvec0.store(out, SIZE);
|
||||
} else {
|
||||
for (int d = 0; d < SIZE; d += kVecSize) {
|
||||
bVec x_bvec = bVec::loadu(input + d);
|
||||
fVec x_fvec0, x_fvec1;
|
||||
std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec);
|
||||
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec0);
|
||||
max_fvec = at::vec::maximum(max_fvec, x_fvec1);
|
||||
x_fvec0.store(out + d);
|
||||
x_fvec1.store(out + d + fVec::size());
|
||||
}
|
||||
}
|
||||
}
|
||||
float max_val = vec_reduce_max(max_fvec);
|
||||
@@ -174,41 +189,44 @@ void topk_sigmoid_kernel_impl(
|
||||
float* __restrict__ topk_weights,
|
||||
int32_t* __restrict__ topk_ids,
|
||||
const scalar_t* __restrict__ gating_output,
|
||||
const float* __restrict__ correction_bias,
|
||||
int64_t num_tokens,
|
||||
int64_t topk,
|
||||
bool renormalize) {
|
||||
using Vec = at::vec::Vectorized<float>;
|
||||
const int64_t num_experts_per_group = NUM_EXPERTS;
|
||||
using elem_t = std::pair<float, int32_t>;
|
||||
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
|
||||
alignas(64) float scores[NUM_EXPERTS];
|
||||
using elem_t = std::pair<float, int32_t>;
|
||||
std::vector<elem_t> queue(num_experts_per_group);
|
||||
alignas(64) elem_t queue[NUM_EXPERTS];
|
||||
|
||||
for (int64_t i = begin; i < end; ++i) {
|
||||
at::vec::convert<scalar_t, float>(gating_output + i * NUM_EXPERTS, scores, NUM_EXPERTS);
|
||||
const scalar_t* token_logits = gating_output + i * NUM_EXPERTS;
|
||||
|
||||
float gmax = at::vec::reduce_all<float>(
|
||||
[](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, scores, num_experts_per_group);
|
||||
|
||||
// find position of first max,
|
||||
// note that we may have multiple max values.
|
||||
int first_max_idx = -1;
|
||||
for (int64_t e = 0; e < num_experts_per_group; ++e) {
|
||||
if (scores[e] == gmax) {
|
||||
first_max_idx = e;
|
||||
break;
|
||||
if (correction_bias == nullptr) {
|
||||
at::vec::convert<scalar_t, float>(token_logits, scores, NUM_EXPERTS);
|
||||
for (int32_t expert = 0; expert < NUM_EXPERTS; ++expert) {
|
||||
queue[expert] = {scores[expert], expert};
|
||||
}
|
||||
} else {
|
||||
sigmoid<scalar_t, NUM_EXPERTS>(scores, token_logits);
|
||||
for (int32_t expert = 0; expert < NUM_EXPERTS; ++expert) {
|
||||
queue[expert] = {scores[expert] + correction_bias[expert], expert};
|
||||
}
|
||||
}
|
||||
|
||||
// scalar sigmoid
|
||||
topk_weights[i] = 1.0 / (1.0 + exp(0.0 - gmax));
|
||||
topk_ids[i] = first_max_idx;
|
||||
std::partial_sort(queue, queue + topk, queue + NUM_EXPERTS, [](const elem_t& x, const elem_t& y) -> bool {
|
||||
return x.first > y.first;
|
||||
});
|
||||
|
||||
float sum = 0.f;
|
||||
for (int64_t j = 0; j < topk; ++j) {
|
||||
int32_t expert_idx = queue[j].second;
|
||||
float weight = correction_bias == nullptr ? 1.f / (1.f + std::exp(-scores[expert_idx])) : scores[expert_idx];
|
||||
topk_weights[i * topk + j] = weight;
|
||||
topk_ids[i * topk + j] = expert_idx;
|
||||
sum += weight;
|
||||
}
|
||||
|
||||
if (renormalize) {
|
||||
float sum = 0.f;
|
||||
for (int64_t j = 0; j < topk; ++j) {
|
||||
sum += topk_weights[i * topk + j];
|
||||
}
|
||||
float scale = 1.f / sum;
|
||||
for (int64_t j = 0; j < topk; ++j) {
|
||||
topk_weights[i * topk + j] *= scale;
|
||||
@@ -223,10 +241,12 @@ void topk_softmax_kernel_impl(
|
||||
float* __restrict__ topk_weights,
|
||||
int32_t* __restrict__ topk_ids,
|
||||
const scalar_t* __restrict__ gating_output,
|
||||
const float* __restrict__ correction_bias,
|
||||
int64_t num_tokens,
|
||||
int64_t topk,
|
||||
bool renormalize) {
|
||||
const int64_t num_experts_per_group = NUM_EXPERTS;
|
||||
const bool use_correction_bias = correction_bias != nullptr;
|
||||
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
|
||||
alignas(64) float scores[NUM_EXPERTS];
|
||||
using elem_t = std::pair<float, int32_t>;
|
||||
@@ -236,7 +256,8 @@ void topk_softmax_kernel_impl(
|
||||
softmax<scalar_t, NUM_EXPERTS>(scores, gating_output + i * NUM_EXPERTS);
|
||||
|
||||
for (int64_t e = 0; e < num_experts_per_group; ++e) {
|
||||
queue[e] = {scores[e], e};
|
||||
const float score = use_correction_bias ? scores[e] + correction_bias[e] : scores[e];
|
||||
queue[e] = {score, e};
|
||||
}
|
||||
|
||||
std::partial_sort(queue.begin(), queue.begin() + topk, queue.end(), [](const elem_t& x, const elem_t& y) -> bool {
|
||||
@@ -244,8 +265,9 @@ void topk_softmax_kernel_impl(
|
||||
});
|
||||
|
||||
for (int64_t j = 0; j < topk; ++j) {
|
||||
topk_weights[i * topk + j] = queue[j].first;
|
||||
topk_ids[i * topk + j] = queue[j].second;
|
||||
int32_t expert_idx = queue[j].second;
|
||||
topk_weights[i * topk + j] = scores[expert_idx];
|
||||
topk_ids[i * topk + j] = expert_idx;
|
||||
}
|
||||
|
||||
if (renormalize) {
|
||||
@@ -420,6 +442,7 @@ void biased_grouped_topk_kernel_impl(
|
||||
topk_weights.data_ptr<float>(), \
|
||||
topk_ids.data_ptr<int32_t>(), \
|
||||
gating_output.data_ptr<scalar_t>(), \
|
||||
correction_bias_ptr, \
|
||||
num_tokens, \
|
||||
topk, \
|
||||
renormalize);
|
||||
@@ -429,6 +452,7 @@ void biased_grouped_topk_kernel_impl(
|
||||
topk_weights.data_ptr<float>(), \
|
||||
topk_ids.data_ptr<int32_t>(), \
|
||||
gating_output.data_ptr<scalar_t>(), \
|
||||
correction_bias_ptr, \
|
||||
num_tokens, \
|
||||
topk, \
|
||||
renormalize);
|
||||
@@ -447,21 +471,30 @@ void biased_grouped_topk_kernel_impl(
|
||||
|
||||
} // anonymous namespace
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor>
|
||||
topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) {
|
||||
std::tuple<at::Tensor, at::Tensor> topk_sigmoid_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
at::Tensor& gating_output,
|
||||
int64_t topk,
|
||||
bool renormalize,
|
||||
const std::optional<at::Tensor>& correction_bias) {
|
||||
CHECK_INPUT(gating_output);
|
||||
|
||||
const auto st = hidden_states.scalar_type();
|
||||
CHECK_EQ(gating_output.scalar_type(), st);
|
||||
const auto st = gating_output.scalar_type();
|
||||
|
||||
int64_t num_tokens = hidden_states.size(0);
|
||||
int64_t num_experts = gating_output.size(1);
|
||||
TORCH_CHECK(gating_output.size(0) == num_tokens, "Number of tokens mismatch");
|
||||
TORCH_CHECK(topk == 1, "topk_sigmoid only supports topk=1 case");
|
||||
TORCH_CHECK(topk > 0 && topk <= num_experts, "topk must satisfy 0 < topk <= num_experts");
|
||||
const float* correction_bias_ptr = nullptr;
|
||||
if (correction_bias.has_value()) {
|
||||
const auto& correction_bias_tensor = correction_bias.value();
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(correction_bias_tensor, {num_experts}, at::kFloat);
|
||||
correction_bias_ptr = correction_bias_tensor.data_ptr<float>();
|
||||
}
|
||||
at::Tensor topk_weights = at::empty({num_tokens, topk}, hidden_states.options().dtype(at::kFloat));
|
||||
at::Tensor topk_ids = at::empty({num_tokens, topk}, hidden_states.options().dtype(at::kInt));
|
||||
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "topk_sigmoid_kernel", [&] {
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES_AND(at::kFloat, st, "topk_sigmoid_kernel", [&] {
|
||||
switch (num_experts) {
|
||||
case 1:
|
||||
LAUNCH_TOPK_SIGMOID_KERNEL(1);
|
||||
@@ -493,6 +526,12 @@ topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t t
|
||||
case 256:
|
||||
LAUNCH_TOPK_SIGMOID_KERNEL(256);
|
||||
break;
|
||||
case 384:
|
||||
LAUNCH_TOPK_SIGMOID_KERNEL(384);
|
||||
break;
|
||||
case 512:
|
||||
LAUNCH_TOPK_SIGMOID_KERNEL(512);
|
||||
break;
|
||||
default:
|
||||
TORCH_CHECK(false, "Unexpected num_experts: ", num_experts);
|
||||
}
|
||||
@@ -500,21 +539,31 @@ topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t t
|
||||
return std::make_tuple(topk_weights, topk_ids);
|
||||
}
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor>
|
||||
topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) {
|
||||
std::tuple<at::Tensor, at::Tensor> topk_softmax_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
at::Tensor& gating_output,
|
||||
int64_t topk,
|
||||
bool renormalize,
|
||||
const std::optional<at::Tensor>& correction_bias) {
|
||||
CHECK_INPUT(gating_output);
|
||||
|
||||
const auto st = hidden_states.scalar_type();
|
||||
CHECK_EQ(gating_output.scalar_type(), st);
|
||||
const auto st = gating_output.scalar_type();
|
||||
|
||||
int64_t num_tokens = hidden_states.size(0);
|
||||
int64_t num_experts = gating_output.size(1);
|
||||
TORCH_CHECK(gating_output.size(0) == num_tokens, "Number of tokens mismatch");
|
||||
TORCH_CHECK(topk > 0 && topk <= num_experts, "topk must satisfy 0 < topk <= num_experts");
|
||||
const float* correction_bias_ptr = nullptr;
|
||||
if (correction_bias.has_value()) {
|
||||
const auto& correction_bias_tensor = correction_bias.value();
|
||||
CHECK_INPUT_SHAPE_DTYPE<false>(correction_bias_tensor, {num_experts}, at::kFloat);
|
||||
correction_bias_ptr = correction_bias_tensor.data_ptr<float>();
|
||||
}
|
||||
|
||||
at::Tensor topk_weights = at::empty({num_tokens, topk}, hidden_states.options().dtype(at::kFloat));
|
||||
at::Tensor topk_ids = at::empty({num_tokens, topk}, hidden_states.options().dtype(at::kInt));
|
||||
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(st, "topk_softmax_cpu", [&] {
|
||||
AT_DISPATCH_REDUCED_FLOATING_TYPES_AND(at::kFloat, st, "topk_softmax_cpu", [&] {
|
||||
switch (num_experts) {
|
||||
case 1:
|
||||
LAUNCH_TOPK_SOFTMAX_KERNEL(1);
|
||||
|
||||
@@ -58,6 +58,19 @@ at::Tensor fused_add_layernorm_cpu(
|
||||
const std::optional<at::Tensor>& bias,
|
||||
double eps);
|
||||
|
||||
// fused_qk_rmsnorm
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_rmsnorm_cpu(
|
||||
const at::Tensor& q, const at::Tensor& k, const at::Tensor& q_weight, const at::Tensor& k_weight, double eps);
|
||||
at::Tensor fused_qk_rmsnorm_sumsq_cpu(const at::Tensor& q, const at::Tensor& k);
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_rmsnorm_apply_from_stats_cpu(
|
||||
const at::Tensor& q,
|
||||
const at::Tensor& k,
|
||||
const at::Tensor& q_weight,
|
||||
const at::Tensor& k_weight,
|
||||
const at::Tensor& sum_sq,
|
||||
int64_t tp_world_size,
|
||||
double eps);
|
||||
|
||||
// fused_qk_gemma_rmsnorm
|
||||
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
||||
const at::Tensor& q,
|
||||
@@ -157,10 +170,18 @@ void rotate_input_ids_cpu(
|
||||
const std::optional<at::Tensor>& select_index_opt);
|
||||
|
||||
// topk
|
||||
std::tuple<at::Tensor, at::Tensor>
|
||||
topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize);
|
||||
std::tuple<at::Tensor, at::Tensor>
|
||||
topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize);
|
||||
std::tuple<at::Tensor, at::Tensor> topk_sigmoid_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
at::Tensor& gating_output,
|
||||
int64_t topk,
|
||||
bool renormalize,
|
||||
const std::optional<at::Tensor>& correction_bias);
|
||||
std::tuple<at::Tensor, at::Tensor> topk_softmax_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
at::Tensor& gating_output,
|
||||
int64_t topk,
|
||||
bool renormalize,
|
||||
const std::optional<at::Tensor>& correction_bias);
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor> grouped_topk_cpu(
|
||||
at::Tensor& hidden_states,
|
||||
@@ -580,6 +601,16 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
"fused_add_layernorm_cpu(Tensor input, Tensor residual, Tensor weight, Tensor? bias, float eps) -> "
|
||||
"Tensor");
|
||||
m.impl("fused_add_layernorm_cpu", torch::kCPU, &fused_add_layernorm_cpu);
|
||||
m.def(
|
||||
"fused_qk_rmsnorm_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, float eps) -> "
|
||||
"(Tensor, Tensor)");
|
||||
m.impl("fused_qk_rmsnorm_cpu", torch::kCPU, &fused_qk_rmsnorm_cpu);
|
||||
m.def("fused_qk_rmsnorm_sumsq_cpu(Tensor q, Tensor k) -> Tensor");
|
||||
m.impl("fused_qk_rmsnorm_sumsq_cpu", torch::kCPU, &fused_qk_rmsnorm_sumsq_cpu);
|
||||
m.def(
|
||||
"fused_qk_rmsnorm_apply_from_stats_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, Tensor sum_sq, int "
|
||||
"tp_world_size, float eps) -> (Tensor, Tensor)");
|
||||
m.impl("fused_qk_rmsnorm_apply_from_stats_cpu", torch::kCPU, &fused_qk_rmsnorm_apply_from_stats_cpu);
|
||||
m.def(
|
||||
"fused_qk_gemma_rmsnorm_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, float eps, int head_dim) -> "
|
||||
"(Tensor, Tensor)");
|
||||
@@ -649,9 +680,13 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) {
|
||||
m.impl("reconstruct_indices_from_tree_mask_cpu", torch::kCPU, &reconstruct_indices_from_tree_mask_cpu);
|
||||
|
||||
// topk
|
||||
m.def("topk_sigmoid_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize) -> (Tensor, Tensor)");
|
||||
m.def(
|
||||
"topk_sigmoid_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize, "
|
||||
"Tensor? correction_bias=None) -> (Tensor, Tensor)");
|
||||
m.impl("topk_sigmoid_cpu", torch::kCPU, &topk_sigmoid_cpu);
|
||||
m.def("topk_softmax_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize) -> (Tensor, Tensor)");
|
||||
m.def(
|
||||
"topk_softmax_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize, "
|
||||
"Tensor? correction_bias=None) -> (Tensor, Tensor)");
|
||||
m.impl("topk_softmax_cpu", torch::kCPU, &topk_softmax_cpu);
|
||||
m.def(
|
||||
"grouped_topk_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize, int num_expert_group, "
|
||||
|
||||
@@ -781,11 +781,24 @@ def fused_topk_cpu(
|
||||
if num_token_non_padded is not None:
|
||||
raise ValueError("num_token_non_padded is not supported for CPU fused topk")
|
||||
|
||||
# TODO: add c++ kernel for cpu
|
||||
# The topk_softmax_cpu kernel only handles vanilla softmax scoring with no
|
||||
# correction bias. Fall back to the torch-native impl for the rest
|
||||
# (e.g. MiniMax sets both correction_bias and scoring_func).
|
||||
if correction_bias is not None or scoring_func != "softmax":
|
||||
if scoring_func == "softmax":
|
||||
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_softmax_cpu(
|
||||
hidden_states=hidden_states,
|
||||
gating_output=gating_output,
|
||||
topk=topk,
|
||||
renormalize=renormalize,
|
||||
correction_bias=correction_bias,
|
||||
)
|
||||
elif scoring_func == "sigmoid":
|
||||
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_sigmoid_cpu(
|
||||
hidden_states=hidden_states,
|
||||
gating_output=gating_output,
|
||||
topk=topk,
|
||||
renormalize=renormalize,
|
||||
correction_bias=correction_bias,
|
||||
)
|
||||
else:
|
||||
# Fall back to the torch-native impl for the rest
|
||||
return fused_topk_torch_native(
|
||||
hidden_states,
|
||||
gating_output,
|
||||
@@ -795,12 +808,6 @@ def fused_topk_cpu(
|
||||
scoring_func=scoring_func,
|
||||
)
|
||||
|
||||
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_softmax_cpu(
|
||||
hidden_states=hidden_states,
|
||||
gating_output=gating_output,
|
||||
topk=topk,
|
||||
renormalize=renormalize,
|
||||
)
|
||||
return topk_weights, topk_ids
|
||||
|
||||
|
||||
|
||||
@@ -201,6 +201,18 @@ def register_fake_ops(tp_size: int):
|
||||
def _(input, *args, **kwargs):
|
||||
return torch.empty_like(input)
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_rmsnorm_cpu")
|
||||
def _(q, k, *args, **kwargs):
|
||||
return torch.empty_like(q), torch.empty_like(k)
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_rmsnorm_sumsq_cpu")
|
||||
def _(q, k):
|
||||
return torch.empty((q.shape[0], 2), dtype=torch.float32, device=q.device)
|
||||
|
||||
@register_cpu_compile_fake("fused_qk_rmsnorm_apply_from_stats_cpu")
|
||||
def _(q, k, *args, **kwargs):
|
||||
return torch.empty_like(q), torch.empty_like(k)
|
||||
|
||||
@register_cpu_compile_fake("shm_allgather")
|
||||
def _(data, dim):
|
||||
return torch.cat([data] * tp_size, dim=dim)
|
||||
@@ -385,7 +397,7 @@ def register_fake_ops(tp_size: int):
|
||||
return topk_weights, topk_ids
|
||||
|
||||
@register_cpu_compile_fake("topk_sigmoid_cpu")
|
||||
def _(hidden_states, gating_output, topk, renormalize):
|
||||
def _(hidden_states, gating_output, topk, renormalize, correction_bias=None):
|
||||
num_tokens = hidden_states.shape[0]
|
||||
shape = (num_tokens, topk)
|
||||
return (
|
||||
@@ -399,6 +411,7 @@ def register_fake_ops(tp_size: int):
|
||||
gating_output,
|
||||
topk,
|
||||
renormalize,
|
||||
correction_bias=None,
|
||||
):
|
||||
num_tokens = hidden_states.shape[0]
|
||||
shape = (num_tokens, topk)
|
||||
|
||||
@@ -482,10 +482,26 @@ class MiniMaxM2QKRMSNorm:
|
||||
return q, k
|
||||
|
||||
def _forward_cpu(self, q: torch.Tensor, k: torch.Tensor):
|
||||
# TODO: add c++ kernel for cpu
|
||||
q = self._q_norm(q.contiguous())
|
||||
k = self._k_norm(k.contiguous())
|
||||
return q, k
|
||||
if self._world_size > 1:
|
||||
sum_sq = torch.ops.sgl_kernel.fused_qk_rmsnorm_sumsq_cpu(q, k)
|
||||
sum_sq = attn_tp_all_reduce(sum_sq)
|
||||
return torch.ops.sgl_kernel.fused_qk_rmsnorm_apply_from_stats_cpu(
|
||||
q,
|
||||
k,
|
||||
self._q_norm.weight,
|
||||
self._k_norm.weight,
|
||||
sum_sq,
|
||||
self._world_size,
|
||||
self._eps,
|
||||
)
|
||||
|
||||
return torch.ops.sgl_kernel.fused_qk_rmsnorm_cpu(
|
||||
q,
|
||||
k,
|
||||
self._q_norm.weight,
|
||||
self._k_norm.weight,
|
||||
self._eps,
|
||||
)
|
||||
|
||||
|
||||
class MiniMaxM2MoE(nn.Module):
|
||||
|
||||
Reference in New Issue
Block a user