diff --git a/python/sglang/kernels/aot/csrc/cpu/norm.cpp b/python/sglang/kernels/aot/csrc/cpu/norm.cpp index 76cce6a85..058441f04 100644 --- a/python/sglang/kernels/aot/csrc/cpu/norm.cpp +++ b/python/sglang/kernels/aot/csrc/cpu/norm.cpp @@ -467,6 +467,167 @@ void fused_qk_norm4d_kernel_impl( }); } +template +float sum_squares(const scalar_t* __restrict__ input, int64_t size) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + 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(input[i]); + sum_val += input_val * input_val; + } + return sum_val + vec_reduce_sum(sum_fvec); +} + +template +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 +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; + using fVec = at::vec::Vectorized; + 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(size * tp_world_size); + + if constexpr (NormTraits::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::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::has_weight) { + auto [w_fvec0, w_fvec1] = load_float_vec2(static_cast(params.weight) + i); + if constexpr (NormTraits::has_shift) { + w_fvec0 = NormTraits::apply_shift(w_fvec0, shift_fvec); + w_fvec1 = NormTraits::apply_shift(w_fvec1, shift_fvec); + } + x_fvec0 = NormTraits::apply_weight(x_fvec0, w_fvec0); + x_fvec1 = NormTraits::apply_weight(x_fvec1, w_fvec1); + } + if constexpr (NormTraits::has_bias) { + if (use_bias) { + auto [b_fvec0, b_fvec1] = load_float_vec2(static_cast(params.bias) + i); + x_fvec0 = NormTraits::apply_bias(x_fvec0, b_fvec0); + x_fvec1 = NormTraits::apply_bias(x_fvec1, b_fvec1); + } + } + + convert_from_float_ext(x_fvec0, x_fvec1).store(output + i); + } + + for (; i < size; ++i) { + float x_val = static_cast(input[i]); + if constexpr (NormTraits::has_mean) { + x_val = (x_val - mean) * scale; + } else { + x_val = x_val * scale; + } + if constexpr (NormTraits::has_weight) { + float w_val = static_cast(static_cast(params.weight)[i]); + if constexpr (NormTraits::has_shift) { + w_val = NormTraits::apply_shift(w_val, params.shift); + } + x_val = NormTraits::apply_weight(x_val, w_val); + } + if constexpr (NormTraits::has_bias) { + if (use_bias) { + const float b_val = static_cast(static_cast(params.bias)[i]); + x_val = NormTraits::apply_bias(x_val, b_val); + } + } + output[i] = static_cast(x_val); + } +} + +template +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::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::has_mean) { + q_sum = sum[b * 2]; + k_sum = sum[b * 2 + 1]; + } + apply_norm_from_stats( + 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( + 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 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(q_weight, {q.size(1)}, st); + CHECK_INPUT_SHAPE_DTYPE(k_weight, {k.size(1)}, st); + + NormParams q_params{q, static_cast(eps)}; + q_params.weight = q_weight.data_ptr(); + + NormParams k_params{k, static_cast(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( + q_out.data_ptr(), + k_out.data_ptr(), + nullptr, + q.data_ptr(), + k.data_ptr(), + 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( + sum_sq.data_ptr(), q.data_ptr(), k.data_ptr(), 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 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(q_weight, {q.size(1)}, st); + CHECK_INPUT_SHAPE_DTYPE(k_weight, {k.size(1)}, st); + CHECK_INPUT_SHAPE_DTYPE(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(eps)}; + q_params.weight = q_weight.data_ptr(); + NormParams k_params{k, static_cast(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( + q_out.data_ptr(), + k_out.data_ptr(), + q.data_ptr(), + k.data_ptr(), + nullptr, + sum_sq.data_ptr(), + 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 fused_qk_gemma_rmsnorm_cpu( diff --git a/python/sglang/kernels/aot/csrc/cpu/topk.cpp b/python/sglang/kernels/aot/csrc/cpu/topk.cpp index da5aae97a..168a70382 100644 --- a/python/sglang/kernels/aot/csrc/cpu/topk.cpp +++ b/python/sglang/kernels/aot/csrc/cpu/topk.cpp @@ -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::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) { + 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; - const int64_t num_experts_per_group = NUM_EXPERTS; + using elem_t = std::pair; at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) { alignas(64) float scores[NUM_EXPERTS]; - using elem_t = std::pair; - std::vector queue(num_experts_per_group); + alignas(64) elem_t queue[NUM_EXPERTS]; for (int64_t i = begin; i < end; ++i) { - at::vec::convert(gating_output + i * NUM_EXPERTS, scores, NUM_EXPERTS); + const scalar_t* token_logits = gating_output + i * NUM_EXPERTS; - float gmax = at::vec::reduce_all( - [](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(token_logits, scores, NUM_EXPERTS); + for (int32_t expert = 0; expert < NUM_EXPERTS; ++expert) { + queue[expert] = {scores[expert], expert}; + } + } else { + sigmoid(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; @@ -236,7 +256,8 @@ void topk_softmax_kernel_impl( softmax(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(), \ topk_ids.data_ptr(), \ gating_output.data_ptr(), \ + correction_bias_ptr, \ num_tokens, \ topk, \ renormalize); @@ -429,6 +452,7 @@ void biased_grouped_topk_kernel_impl( topk_weights.data_ptr(), \ topk_ids.data_ptr(), \ gating_output.data_ptr(), \ + correction_bias_ptr, \ num_tokens, \ topk, \ renormalize); @@ -447,21 +471,30 @@ void biased_grouped_topk_kernel_impl( } // anonymous namespace -std::tuple -topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) { +std::tuple topk_sigmoid_cpu( + at::Tensor& hidden_states, + at::Tensor& gating_output, + int64_t topk, + bool renormalize, + const std::optional& 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(correction_bias_tensor, {num_experts}, at::kFloat); + correction_bias_ptr = correction_bias_tensor.data_ptr(); + } 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 -topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize) { +std::tuple topk_softmax_cpu( + at::Tensor& hidden_states, + at::Tensor& gating_output, + int64_t topk, + bool renormalize, + const std::optional& 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(correction_bias_tensor, {num_experts}, at::kFloat); + correction_bias_ptr = correction_bias_tensor.data_ptr(); + } 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); diff --git a/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp b/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp index c2b375f36..51b83beb6 100644 --- a/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp +++ b/python/sglang/kernels/aot/csrc/cpu/torch_extension_cpu.cpp @@ -58,6 +58,19 @@ at::Tensor fused_add_layernorm_cpu( const std::optional& bias, double eps); +// fused_qk_rmsnorm +std::tuple 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 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 fused_qk_gemma_rmsnorm_cpu( const at::Tensor& q, @@ -157,10 +170,18 @@ void rotate_input_ids_cpu( const std::optional& select_index_opt); // topk -std::tuple -topk_sigmoid_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize); -std::tuple -topk_softmax_cpu(at::Tensor& hidden_states, at::Tensor& gating_output, int64_t topk, bool renormalize); +std::tuple topk_sigmoid_cpu( + at::Tensor& hidden_states, + at::Tensor& gating_output, + int64_t topk, + bool renormalize, + const std::optional& correction_bias); +std::tuple topk_softmax_cpu( + at::Tensor& hidden_states, + at::Tensor& gating_output, + int64_t topk, + bool renormalize, + const std::optional& correction_bias); std::tuple 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, " diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 7c57ba2f8..998eb90fd 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -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 diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 65f068706..eb066cdb3 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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) diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 6fb388ee1..29adf3ef4 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -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): diff --git a/test/registered/cpu/test_norm.py b/test/registered/cpu/test_norm.py index 0061908a6..bc13b1395 100644 --- a/test/registered/cpu/test_norm.py +++ b/test/registered/cpu/test_norm.py @@ -234,6 +234,87 @@ class TestFusedRMSNormGated: 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: def _forward_native( diff --git a/test/registered/cpu/test_topk.py b/test/registered/cpu/test_topk.py index 54b78c74e..461940bf1 100644 --- a/test/registered/cpu/test_topk.py +++ b/test/registered/cpu/test_topk.py @@ -1,3 +1,4 @@ +import itertools import unittest import torch @@ -207,6 +208,89 @@ class TestTopK(CustomTestCase): self._run_single_test(123, 256, 4, 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): def _run_single_test( @@ -251,6 +335,79 @@ class TestCustomTopK(CustomTestCase): 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__": unittest.main()