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
|
||||||
#undef LAUNCH_PARALLEL_LOOP_HD
|
#undef LAUNCH_PARALLEL_LOOP_HD
|
||||||
} // anonymous namespace
|
} // anonymous namespace
|
||||||
@@ -694,6 +855,102 @@ at::Tensor fused_add_layernorm_cpu(
|
|||||||
return output;
|
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
|
// q : {batch_size, num_head * head_dim} 2D
|
||||||
// k : {batch_size, num_head_kv * head_dim} 2D
|
// k : {batch_size, num_head_kv * head_dim} 2D
|
||||||
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
||||||
|
|||||||
@@ -12,6 +12,20 @@ inline void softmax(float* __restrict__ out, const scalar_t* __restrict__ input)
|
|||||||
|
|
||||||
// step 1: get max
|
// step 1: get max
|
||||||
fVec max_fvec = fVec(-std::numeric_limits<float>::infinity());
|
fVec max_fvec = fVec(-std::numeric_limits<float>::infinity());
|
||||||
|
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 {
|
||||||
if constexpr (SIZE < kVecSize) {
|
if constexpr (SIZE < kVecSize) {
|
||||||
// SIZE = 1, 2, 4, 8, 16; only the top half is used
|
// SIZE = 1, 2, 4, 8, 16; only the top half is used
|
||||||
bVec x_bvec = bVec::loadu(input, SIZE);
|
bVec x_bvec = bVec::loadu(input, SIZE);
|
||||||
@@ -32,6 +46,7 @@ inline void softmax(float* __restrict__ out, const scalar_t* __restrict__ input)
|
|||||||
x_fvec1.store(out + d + fVec::size());
|
x_fvec1.store(out + d + fVec::size());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
}
|
||||||
float max_val = vec_reduce_max(max_fvec);
|
float max_val = vec_reduce_max(max_fvec);
|
||||||
max_fvec = fVec(max_val);
|
max_fvec = fVec(max_val);
|
||||||
|
|
||||||
@@ -174,41 +189,44 @@ void topk_sigmoid_kernel_impl(
|
|||||||
float* __restrict__ topk_weights,
|
float* __restrict__ topk_weights,
|
||||||
int32_t* __restrict__ topk_ids,
|
int32_t* __restrict__ topk_ids,
|
||||||
const scalar_t* __restrict__ gating_output,
|
const scalar_t* __restrict__ gating_output,
|
||||||
|
const float* __restrict__ correction_bias,
|
||||||
int64_t num_tokens,
|
int64_t num_tokens,
|
||||||
int64_t topk,
|
int64_t topk,
|
||||||
bool renormalize) {
|
bool renormalize) {
|
||||||
using Vec = at::vec::Vectorized<float>;
|
using elem_t = std::pair<float, int32_t>;
|
||||||
const int64_t num_experts_per_group = NUM_EXPERTS;
|
|
||||||
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
|
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
|
||||||
alignas(64) float scores[NUM_EXPERTS];
|
alignas(64) float scores[NUM_EXPERTS];
|
||||||
using elem_t = std::pair<float, int32_t>;
|
alignas(64) elem_t queue[NUM_EXPERTS];
|
||||||
std::vector<elem_t> queue(num_experts_per_group);
|
|
||||||
|
|
||||||
for (int64_t i = begin; i < end; ++i) {
|
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>(
|
if (correction_bias == nullptr) {
|
||||||
[](Vec& x, Vec& y) { return at::vec::maximum(x, y); }, scores, num_experts_per_group);
|
at::vec::convert<scalar_t, float>(token_logits, scores, NUM_EXPERTS);
|
||||||
|
for (int32_t expert = 0; expert < NUM_EXPERTS; ++expert) {
|
||||||
// find position of first max,
|
queue[expert] = {scores[expert], expert};
|
||||||
// note that we may have multiple max values.
|
}
|
||||||
int first_max_idx = -1;
|
} else {
|
||||||
for (int64_t e = 0; e < num_experts_per_group; ++e) {
|
sigmoid<scalar_t, NUM_EXPERTS>(scores, token_logits);
|
||||||
if (scores[e] == gmax) {
|
for (int32_t expert = 0; expert < NUM_EXPERTS; ++expert) {
|
||||||
first_max_idx = e;
|
queue[expert] = {scores[expert] + correction_bias[expert], expert};
|
||||||
break;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// scalar sigmoid
|
std::partial_sort(queue, queue + topk, queue + NUM_EXPERTS, [](const elem_t& x, const elem_t& y) -> bool {
|
||||||
topk_weights[i] = 1.0 / (1.0 + exp(0.0 - gmax));
|
return x.first > y.first;
|
||||||
topk_ids[i] = first_max_idx;
|
});
|
||||||
|
|
||||||
if (renormalize) {
|
|
||||||
float sum = 0.f;
|
float sum = 0.f;
|
||||||
for (int64_t j = 0; j < topk; ++j) {
|
for (int64_t j = 0; j < topk; ++j) {
|
||||||
sum += topk_weights[i * 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 scale = 1.f / sum;
|
float scale = 1.f / sum;
|
||||||
for (int64_t j = 0; j < topk; ++j) {
|
for (int64_t j = 0; j < topk; ++j) {
|
||||||
topk_weights[i * topk + j] *= scale;
|
topk_weights[i * topk + j] *= scale;
|
||||||
@@ -223,10 +241,12 @@ void topk_softmax_kernel_impl(
|
|||||||
float* __restrict__ topk_weights,
|
float* __restrict__ topk_weights,
|
||||||
int32_t* __restrict__ topk_ids,
|
int32_t* __restrict__ topk_ids,
|
||||||
const scalar_t* __restrict__ gating_output,
|
const scalar_t* __restrict__ gating_output,
|
||||||
|
const float* __restrict__ correction_bias,
|
||||||
int64_t num_tokens,
|
int64_t num_tokens,
|
||||||
int64_t topk,
|
int64_t topk,
|
||||||
bool renormalize) {
|
bool renormalize) {
|
||||||
const int64_t num_experts_per_group = NUM_EXPERTS;
|
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) {
|
at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) {
|
||||||
alignas(64) float scores[NUM_EXPERTS];
|
alignas(64) float scores[NUM_EXPERTS];
|
||||||
using elem_t = std::pair<float, int32_t>;
|
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);
|
softmax<scalar_t, NUM_EXPERTS>(scores, gating_output + i * NUM_EXPERTS);
|
||||||
|
|
||||||
for (int64_t e = 0; e < num_experts_per_group; ++e) {
|
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 {
|
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) {
|
for (int64_t j = 0; j < topk; ++j) {
|
||||||
topk_weights[i * topk + j] = queue[j].first;
|
int32_t expert_idx = queue[j].second;
|
||||||
topk_ids[i * topk + j] = queue[j].second;
|
topk_weights[i * topk + j] = scores[expert_idx];
|
||||||
|
topk_ids[i * topk + j] = expert_idx;
|
||||||
}
|
}
|
||||||
|
|
||||||
if (renormalize) {
|
if (renormalize) {
|
||||||
@@ -420,6 +442,7 @@ void biased_grouped_topk_kernel_impl(
|
|||||||
topk_weights.data_ptr<float>(), \
|
topk_weights.data_ptr<float>(), \
|
||||||
topk_ids.data_ptr<int32_t>(), \
|
topk_ids.data_ptr<int32_t>(), \
|
||||||
gating_output.data_ptr<scalar_t>(), \
|
gating_output.data_ptr<scalar_t>(), \
|
||||||
|
correction_bias_ptr, \
|
||||||
num_tokens, \
|
num_tokens, \
|
||||||
topk, \
|
topk, \
|
||||||
renormalize);
|
renormalize);
|
||||||
@@ -429,6 +452,7 @@ void biased_grouped_topk_kernel_impl(
|
|||||||
topk_weights.data_ptr<float>(), \
|
topk_weights.data_ptr<float>(), \
|
||||||
topk_ids.data_ptr<int32_t>(), \
|
topk_ids.data_ptr<int32_t>(), \
|
||||||
gating_output.data_ptr<scalar_t>(), \
|
gating_output.data_ptr<scalar_t>(), \
|
||||||
|
correction_bias_ptr, \
|
||||||
num_tokens, \
|
num_tokens, \
|
||||||
topk, \
|
topk, \
|
||||||
renormalize);
|
renormalize);
|
||||||
@@ -447,21 +471,30 @@ void biased_grouped_topk_kernel_impl(
|
|||||||
|
|
||||||
} // anonymous namespace
|
} // anonymous namespace
|
||||||
|
|
||||||
std::tuple<at::Tensor, at::Tensor>
|
std::tuple<at::Tensor, at::Tensor> topk_sigmoid_cpu(
|
||||||
topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) {
|
at::Tensor& hidden_states,
|
||||||
|
at::Tensor& gating_output,
|
||||||
|
int64_t topk,
|
||||||
|
bool renormalize,
|
||||||
|
const std::optional<at::Tensor>& correction_bias) {
|
||||||
CHECK_INPUT(gating_output);
|
CHECK_INPUT(gating_output);
|
||||||
|
|
||||||
const auto st = hidden_states.scalar_type();
|
const auto st = gating_output.scalar_type();
|
||||||
CHECK_EQ(gating_output.scalar_type(), st);
|
|
||||||
|
|
||||||
int64_t num_tokens = hidden_states.size(0);
|
int64_t num_tokens = hidden_states.size(0);
|
||||||
int64_t num_experts = gating_output.size(1);
|
int64_t num_experts = gating_output.size(1);
|
||||||
TORCH_CHECK(gating_output.size(0) == num_tokens, "Number of tokens mismatch");
|
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_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::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) {
|
switch (num_experts) {
|
||||||
case 1:
|
case 1:
|
||||||
LAUNCH_TOPK_SIGMOID_KERNEL(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:
|
case 256:
|
||||||
LAUNCH_TOPK_SIGMOID_KERNEL(256);
|
LAUNCH_TOPK_SIGMOID_KERNEL(256);
|
||||||
break;
|
break;
|
||||||
|
case 384:
|
||||||
|
LAUNCH_TOPK_SIGMOID_KERNEL(384);
|
||||||
|
break;
|
||||||
|
case 512:
|
||||||
|
LAUNCH_TOPK_SIGMOID_KERNEL(512);
|
||||||
|
break;
|
||||||
default:
|
default:
|
||||||
TORCH_CHECK(false, "Unexpected num_experts: ", num_experts);
|
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);
|
return std::make_tuple(topk_weights, topk_ids);
|
||||||
}
|
}
|
||||||
|
|
||||||
std::tuple<at::Tensor, at::Tensor>
|
std::tuple<at::Tensor, at::Tensor> topk_softmax_cpu(
|
||||||
topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) {
|
at::Tensor& hidden_states,
|
||||||
|
at::Tensor& gating_output,
|
||||||
|
int64_t topk,
|
||||||
|
bool renormalize,
|
||||||
|
const std::optional<at::Tensor>& correction_bias) {
|
||||||
CHECK_INPUT(gating_output);
|
CHECK_INPUT(gating_output);
|
||||||
|
|
||||||
const auto st = hidden_states.scalar_type();
|
const auto st = gating_output.scalar_type();
|
||||||
CHECK_EQ(gating_output.scalar_type(), st);
|
|
||||||
|
|
||||||
int64_t num_tokens = hidden_states.size(0);
|
int64_t num_tokens = hidden_states.size(0);
|
||||||
int64_t num_experts = gating_output.size(1);
|
int64_t num_experts = gating_output.size(1);
|
||||||
TORCH_CHECK(gating_output.size(0) == num_tokens, "Number of tokens mismatch");
|
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_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::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) {
|
switch (num_experts) {
|
||||||
case 1:
|
case 1:
|
||||||
LAUNCH_TOPK_SOFTMAX_KERNEL(1);
|
LAUNCH_TOPK_SOFTMAX_KERNEL(1);
|
||||||
|
|||||||
@@ -58,6 +58,19 @@ at::Tensor fused_add_layernorm_cpu(
|
|||||||
const std::optional<at::Tensor>& bias,
|
const std::optional<at::Tensor>& bias,
|
||||||
double eps);
|
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
|
// fused_qk_gemma_rmsnorm
|
||||||
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
std::tuple<at::Tensor, at::Tensor> fused_qk_gemma_rmsnorm_cpu(
|
||||||
const at::Tensor& q,
|
const at::Tensor& q,
|
||||||
@@ -157,10 +170,18 @@ void rotate_input_ids_cpu(
|
|||||||
const std::optional<at::Tensor>& select_index_opt);
|
const std::optional<at::Tensor>& select_index_opt);
|
||||||
|
|
||||||
// topk
|
// topk
|
||||||
std::tuple<at::Tensor, at::Tensor>
|
std::tuple<at::Tensor, at::Tensor> topk_sigmoid_cpu(
|
||||||
topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize);
|
at::Tensor& hidden_states,
|
||||||
std::tuple<at::Tensor, at::Tensor>
|
at::Tensor& gating_output,
|
||||||
topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize);
|
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(
|
std::tuple<at::Tensor, at::Tensor> grouped_topk_cpu(
|
||||||
at::Tensor& hidden_states,
|
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) -> "
|
"fused_add_layernorm_cpu(Tensor input, Tensor residual, Tensor weight, Tensor? bias, float eps) -> "
|
||||||
"Tensor");
|
"Tensor");
|
||||||
m.impl("fused_add_layernorm_cpu", torch::kCPU, &fused_add_layernorm_cpu);
|
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(
|
m.def(
|
||||||
"fused_qk_gemma_rmsnorm_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, float eps, int head_dim) -> "
|
"fused_qk_gemma_rmsnorm_cpu(Tensor q, Tensor k, Tensor q_weight, Tensor k_weight, float eps, int head_dim) -> "
|
||||||
"(Tensor, Tensor)");
|
"(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);
|
m.impl("reconstruct_indices_from_tree_mask_cpu", torch::kCPU, &reconstruct_indices_from_tree_mask_cpu);
|
||||||
|
|
||||||
// topk
|
// 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.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.impl("topk_softmax_cpu", torch::kCPU, &topk_softmax_cpu);
|
||||||
m.def(
|
m.def(
|
||||||
"grouped_topk_cpu(Tensor hidden_states, Tensor gating_output, int topk, bool renormalize, int num_expert_group, "
|
"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:
|
if num_token_non_padded is not None:
|
||||||
raise ValueError("num_token_non_padded is not supported for CPU fused topk")
|
raise ValueError("num_token_non_padded is not supported for CPU fused topk")
|
||||||
|
|
||||||
# TODO: add c++ kernel for cpu
|
if scoring_func == "softmax":
|
||||||
# The topk_softmax_cpu kernel only handles vanilla softmax scoring with no
|
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_softmax_cpu(
|
||||||
# correction bias. Fall back to the torch-native impl for the rest
|
hidden_states=hidden_states,
|
||||||
# (e.g. MiniMax sets both correction_bias and scoring_func).
|
gating_output=gating_output,
|
||||||
if correction_bias is not None or scoring_func != "softmax":
|
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(
|
return fused_topk_torch_native(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
gating_output,
|
gating_output,
|
||||||
@@ -795,12 +808,6 @@ def fused_topk_cpu(
|
|||||||
scoring_func=scoring_func,
|
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
|
return topk_weights, topk_ids
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -201,6 +201,18 @@ def register_fake_ops(tp_size: int):
|
|||||||
def _(input, *args, **kwargs):
|
def _(input, *args, **kwargs):
|
||||||
return torch.empty_like(input)
|
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")
|
@register_cpu_compile_fake("shm_allgather")
|
||||||
def _(data, dim):
|
def _(data, dim):
|
||||||
return torch.cat([data] * tp_size, dim=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
|
return topk_weights, topk_ids
|
||||||
|
|
||||||
@register_cpu_compile_fake("topk_sigmoid_cpu")
|
@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]
|
num_tokens = hidden_states.shape[0]
|
||||||
shape = (num_tokens, topk)
|
shape = (num_tokens, topk)
|
||||||
return (
|
return (
|
||||||
@@ -399,6 +411,7 @@ def register_fake_ops(tp_size: int):
|
|||||||
gating_output,
|
gating_output,
|
||||||
topk,
|
topk,
|
||||||
renormalize,
|
renormalize,
|
||||||
|
correction_bias=None,
|
||||||
):
|
):
|
||||||
num_tokens = hidden_states.shape[0]
|
num_tokens = hidden_states.shape[0]
|
||||||
shape = (num_tokens, topk)
|
shape = (num_tokens, topk)
|
||||||
|
|||||||
@@ -482,10 +482,26 @@ class MiniMaxM2QKRMSNorm:
|
|||||||
return q, k
|
return q, k
|
||||||
|
|
||||||
def _forward_cpu(self, q: torch.Tensor, k: torch.Tensor):
|
def _forward_cpu(self, q: torch.Tensor, k: torch.Tensor):
|
||||||
# TODO: add c++ kernel for cpu
|
if self._world_size > 1:
|
||||||
q = self._q_norm(q.contiguous())
|
sum_sq = torch.ops.sgl_kernel.fused_qk_rmsnorm_sumsq_cpu(q, k)
|
||||||
k = self._k_norm(k.contiguous())
|
sum_sq = attn_tp_all_reduce(sum_sq)
|
||||||
return q, k
|
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):
|
class MiniMaxM2MoE(nn.Module):
|
||||||
|
|||||||
@@ -234,6 +234,87 @@ class TestFusedRMSNormGated:
|
|||||||
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
|
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
|
||||||
|
|
||||||
|
|
||||||
|
class TestFusedQKRMSNorm:
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS)
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"batch_size,q_size,k_size,v_size",
|
||||||
|
[(1, 256, 64, 64), (17, 512, 128, 128)],
|
||||||
|
)
|
||||||
|
def test_fused_qk_rmsnorm(
|
||||||
|
self, batch_size: int, q_size: int, k_size: int, v_size: int, dtype
|
||||||
|
):
|
||||||
|
"""Q and K split views must be normalized over their distinct full widths."""
|
||||||
|
qkv = torch.randn([batch_size, q_size + k_size + v_size], dtype=dtype)
|
||||||
|
q, k, _ = qkv.split([q_size, k_size, v_size], dim=-1)
|
||||||
|
q_weight = torch.randn(q_size, dtype=dtype)
|
||||||
|
k_weight = torch.randn(k_size, dtype=dtype)
|
||||||
|
|
||||||
|
q_out, k_out = torch.ops.sgl_kernel.fused_qk_rmsnorm_cpu(
|
||||||
|
q, k, q_weight, k_weight, eps
|
||||||
|
)
|
||||||
|
ref_q_out = TestNorm()._forward_native(q, q_weight, eps)
|
||||||
|
ref_k_out = TestNorm()._forward_native(k, k_weight, eps)
|
||||||
|
|
||||||
|
atol = rtol = precision[dtype]
|
||||||
|
torch.testing.assert_close(q_out, ref_q_out, atol=atol, rtol=rtol)
|
||||||
|
torch.testing.assert_close(k_out, ref_k_out, atol=atol, rtol=rtol)
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("dtype", DTYPES, ids=DTYPE_IDS)
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"batch_size,q_size,k_size,tp_world_size",
|
||||||
|
[(1, 256, 64, 2), (17, 512, 128, 4)],
|
||||||
|
)
|
||||||
|
def test_fused_qk_rmsnorm_tp(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
q_size: int,
|
||||||
|
k_size: int,
|
||||||
|
tp_world_size: int,
|
||||||
|
dtype,
|
||||||
|
):
|
||||||
|
q = torch.randn([batch_size, q_size], dtype=dtype)
|
||||||
|
k = torch.randn([batch_size, k_size], dtype=dtype)
|
||||||
|
q_weight = torch.randn(q_size, dtype=dtype)
|
||||||
|
k_weight = torch.randn(k_size, dtype=dtype)
|
||||||
|
|
||||||
|
q_shards = q.chunk(tp_world_size, dim=-1)
|
||||||
|
k_shards = k.chunk(tp_world_size, dim=-1)
|
||||||
|
q_weight_shards = q_weight.chunk(tp_world_size)
|
||||||
|
k_weight_shards = k_weight.chunk(tp_world_size)
|
||||||
|
local_sum_sq = [
|
||||||
|
torch.ops.sgl_kernel.fused_qk_rmsnorm_sumsq_cpu(q_shard, k_shard)
|
||||||
|
for q_shard, k_shard in zip(q_shards, k_shards)
|
||||||
|
]
|
||||||
|
global_sum_sq = torch.stack(local_sum_sq).sum(dim=0)
|
||||||
|
|
||||||
|
shard_outputs = [
|
||||||
|
torch.ops.sgl_kernel.fused_qk_rmsnorm_apply_from_stats_cpu(
|
||||||
|
q_shard,
|
||||||
|
k_shard,
|
||||||
|
q_weight_shard,
|
||||||
|
k_weight_shard,
|
||||||
|
global_sum_sq,
|
||||||
|
tp_world_size,
|
||||||
|
eps,
|
||||||
|
)
|
||||||
|
for q_shard, k_shard, q_weight_shard, k_weight_shard in zip(
|
||||||
|
q_shards, k_shards, q_weight_shards, k_weight_shards
|
||||||
|
)
|
||||||
|
]
|
||||||
|
q_out = torch.cat([output[0] for output in shard_outputs], dim=-1)
|
||||||
|
k_out = torch.cat([output[1] for output in shard_outputs], dim=-1)
|
||||||
|
|
||||||
|
ref_q_out = TestNorm()._forward_native(q, q_weight, eps)
|
||||||
|
ref_k_out = TestNorm()._forward_native(k, k_weight, eps)
|
||||||
|
atol = rtol = precision[dtype]
|
||||||
|
torch.testing.assert_close(q_out, ref_q_out, atol=atol, rtol=rtol)
|
||||||
|
torch.testing.assert_close(k_out, ref_k_out, atol=atol, rtol=rtol)
|
||||||
|
|
||||||
|
assert global_sum_sq.shape == (batch_size, 2)
|
||||||
|
assert global_sum_sq.dtype == torch.float32
|
||||||
|
|
||||||
|
|
||||||
class TestLayerNorm:
|
class TestLayerNorm:
|
||||||
|
|
||||||
def _forward_native(
|
def _forward_native(
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import itertools
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -207,6 +208,89 @@ class TestTopK(CustomTestCase):
|
|||||||
self._run_single_test(123, 256, 4, renormalize, torch.bfloat16)
|
self._run_single_test(123, 256, 4, renormalize, torch.bfloat16)
|
||||||
self._run_single_test(123, 160, 6, renormalize, torch.bfloat16)
|
self._run_single_test(123, 160, 6, renormalize, torch.bfloat16)
|
||||||
|
|
||||||
|
def test_topk_softmax_mixed_input_dtypes(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
hidden_states = torch.randn((17, 16), dtype=torch.bfloat16)
|
||||||
|
gating_output = torch.randn((17, 128), dtype=torch.float32)
|
||||||
|
correction_bias = torch.randn(128, dtype=torch.float32)
|
||||||
|
|
||||||
|
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_softmax_cpu(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
gating_output=gating_output,
|
||||||
|
topk=8,
|
||||||
|
renormalize=True,
|
||||||
|
correction_bias=correction_bias,
|
||||||
|
)
|
||||||
|
|
||||||
|
scores = torch.softmax(gating_output, dim=-1)
|
||||||
|
expected_ids = torch.topk(
|
||||||
|
scores + correction_bias.unsqueeze(0), k=8, dim=-1
|
||||||
|
).indices
|
||||||
|
expected_weights = scores.gather(1, topk_ids.to(torch.int64))
|
||||||
|
expected_weights /= expected_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
torch.sort(topk_ids.to(torch.int64), dim=-1).values.tolist(),
|
||||||
|
torch.sort(expected_ids, dim=-1).values.tolist(),
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(topk_weights, expected_weights)
|
||||||
|
|
||||||
|
def test_topk_softmax_with_correction_bias(self):
|
||||||
|
"""Bias must affect expert selection without becoming a routing weight."""
|
||||||
|
for num_tokens, num_experts, topk, with_bias, renormalize in itertools.product(
|
||||||
|
[1, 17, 128],
|
||||||
|
[16, 128, 384, 512],
|
||||||
|
[1, 2, 4, 8],
|
||||||
|
[False, True],
|
||||||
|
[False, True],
|
||||||
|
):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
hidden_states = torch.randn((num_tokens, 16), dtype=torch.bfloat16)
|
||||||
|
gating_output = torch.randn((num_tokens, num_experts), dtype=torch.bfloat16)
|
||||||
|
correction_bias = torch.randn(num_experts) if with_bias else None
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
scores = torch.softmax(gating_output.float(), dim=-1)
|
||||||
|
scores_for_choice = scores
|
||||||
|
if correction_bias is not None:
|
||||||
|
scores_for_choice = scores_for_choice + correction_bias.unsqueeze(0)
|
||||||
|
expected_choice_scores = torch.topk(
|
||||||
|
scores_for_choice, k=topk, dim=-1, sorted=True
|
||||||
|
).values
|
||||||
|
selected_choice_scores = torch.sort(
|
||||||
|
scores_for_choice.gather(1, topk_ids.to(torch.int64)),
|
||||||
|
dim=-1,
|
||||||
|
descending=True,
|
||||||
|
).values
|
||||||
|
|
||||||
|
expected_weights = scores.gather(1, topk_ids.to(torch.int64))
|
||||||
|
if renormalize:
|
||||||
|
expected_weights = expected_weights / expected_weights.sum(
|
||||||
|
dim=-1, keepdim=True
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(topk_ids.dtype, torch.int32)
|
||||||
|
self.assertEqual(topk_weights.dtype, torch.float32)
|
||||||
|
self.assertTrue(torch.all((topk_ids >= 0) & (topk_ids < num_experts)))
|
||||||
|
sorted_ids = torch.sort(topk_ids, dim=-1).values
|
||||||
|
self.assertTrue(torch.all(sorted_ids[:, 1:] != sorted_ids[:, :-1]))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
selected_choice_scores,
|
||||||
|
expected_choice_scores,
|
||||||
|
atol=1e-4,
|
||||||
|
rtol=1e-4,
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
topk_weights, expected_weights, atol=1e-4, rtol=1e-4
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestCustomTopK(CustomTestCase):
|
class TestCustomTopK(CustomTestCase):
|
||||||
def _run_single_test(
|
def _run_single_test(
|
||||||
@@ -251,6 +335,79 @@ class TestCustomTopK(CustomTestCase):
|
|||||||
123, 32, 1, False, torch.bfloat16, native_custom_f, fused_custom_f
|
123, 32, 1, False, torch.bfloat16, native_custom_f, fused_custom_f
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_topk_sigmoid_with_correction_bias(self):
|
||||||
|
"""Biased scores must select experts while returned weights stay unbiased."""
|
||||||
|
for num_tokens, num_experts, topk, with_bias, renormalize in itertools.product(
|
||||||
|
[1, 17, 128],
|
||||||
|
[16, 128, 256, 384, 512],
|
||||||
|
[1, 2, 4, 8],
|
||||||
|
[False, True],
|
||||||
|
[False, True],
|
||||||
|
):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
hidden_states = torch.randn((num_tokens, 16), dtype=torch.bfloat16)
|
||||||
|
gating_output = torch.randn((num_tokens, num_experts), dtype=torch.bfloat16)
|
||||||
|
correction_bias = torch.randn(num_experts) if with_bias else None
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
scores = torch.sigmoid(gating_output.float())
|
||||||
|
scores_for_choice = scores
|
||||||
|
if correction_bias is not None:
|
||||||
|
scores_for_choice = scores_for_choice + correction_bias.unsqueeze(0)
|
||||||
|
|
||||||
|
expected_choice_scores = torch.topk(
|
||||||
|
scores_for_choice, k=topk, dim=-1
|
||||||
|
).values
|
||||||
|
selected_choice_scores = torch.sort(
|
||||||
|
scores_for_choice.gather(1, topk_ids.to(torch.int64)),
|
||||||
|
dim=-1,
|
||||||
|
descending=True,
|
||||||
|
).values
|
||||||
|
|
||||||
|
expected_weights = scores.gather(1, topk_ids.to(torch.int64))
|
||||||
|
if renormalize:
|
||||||
|
expected_weights /= expected_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
self.assertEqual(topk_ids.dtype, torch.int32)
|
||||||
|
self.assertEqual(topk_weights.dtype, torch.float32)
|
||||||
|
self.assertTrue(torch.equal(selected_choice_scores, expected_choice_scores))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
topk_weights, expected_weights, atol=1e-4, rtol=1e-4
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_topk_sigmoid_mixed_input_dtypes(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
hidden_states = torch.randn((17, 16), dtype=torch.bfloat16)
|
||||||
|
gating_output = torch.randn((17, 256), dtype=torch.float32)
|
||||||
|
|
||||||
|
topk_weights, topk_ids = torch.ops.sgl_kernel.topk_sigmoid_cpu(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
gating_output=gating_output,
|
||||||
|
topk=8,
|
||||||
|
renormalize=True,
|
||||||
|
correction_bias=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
scores = torch.sigmoid(gating_output)
|
||||||
|
expected_ids = torch.topk(scores, k=8, dim=-1).indices
|
||||||
|
expected_weights = scores.gather(1, topk_ids.to(torch.int64))
|
||||||
|
expected_weights /= expected_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
torch.equal(
|
||||||
|
torch.sort(topk_ids.to(torch.int64), dim=-1).values,
|
||||||
|
torch.sort(expected_ids, dim=-1).values,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(topk_weights, expected_weights, atol=1e-5, rtol=1e-5)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user