Optimize MiniMax-M2.7 on CPU (#31956)

This commit is contained in:
Xinguo Zhu
2026-08-13 15:04:16 +08:00
committed by GitHub
parent 889c2f31aa
commit 3f6ef01322
8 changed files with 687 additions and 72 deletions
+257
View File
@@ -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(
+99 -50
View File
@@ -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, "
+18 -11
View File
@@ -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)
+20 -4
View File
@@ -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):