diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index bfb32babc..ef7749c9a 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -929,7 +929,6 @@ def biased_grouped_topk_cpu( routed_scaling_factor: Optional[float] = None, apply_routed_scaling_factor_on_output: Optional[bool] = False, ): - assert not apply_routed_scaling_factor_on_output, "Not implemented" return torch.ops.sgl_kernel.biased_grouped_topk_cpu( hidden_states, gating_output, @@ -939,7 +938,7 @@ def biased_grouped_topk_cpu( num_expert_group, topk_group, num_fused_shared_experts, - routed_scaling_factor, + routed_scaling_factor if apply_routed_scaling_factor_on_output else None, # num_token_non_padded must be None since it is not supported in kernel num_token_non_padded=None, ) diff --git a/sgl-kernel/csrc/cpu/common.h b/sgl-kernel/csrc/cpu/common.h index 48f25e21c..65cdb3e86 100644 --- a/sgl-kernel/csrc/cpu/common.h +++ b/sgl-kernel/csrc/cpu/common.h @@ -74,43 +74,44 @@ namespace { } \ }() +// Helper MICRO for CPU_DISPATCH_FLOATING_TYPES_EXT: +// TYPE1: the primary dtype (input, output, weight); +// TYPE2: defined as PARAM_T input +#define CPU_DISPATCH_TYPE1_WITH_PARAM(TYPE1, PARAM_T, ...) \ + switch (TYPE1) { \ + case at::ScalarType::BFloat16: { \ + using scalar_t = at::BFloat16; \ + using param_t = PARAM_T; \ + return __VA_ARGS__(); \ + } \ + case at::ScalarType::Half: { \ + using scalar_t = at::Half; \ + using param_t = PARAM_T; \ + return __VA_ARGS__(); \ + } \ + case at::ScalarType::Float: { \ + using scalar_t = float; \ + using param_t = PARAM_T; \ + return __VA_ARGS__(); \ + } \ + default: \ + TORCH_CHECK(false, "Unsupported floating data type."); \ + } + // dispatch with mixed dtypes (TYPE1, TYPE2): // TYPE1: the primary dtype (input, output, weight); // TYPE2: the secondary dtype (bias, etc.). -#define CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(TYPE1, TYPE2, ...) \ - [&] { \ - if (TYPE2 == at::kFloat) { \ - switch (TYPE1) { \ - case at::ScalarType::BFloat16: { \ - using scalar_t = at::BFloat16; \ - using param_t = float; \ - return __VA_ARGS__(); \ - } \ - case at::ScalarType::Half: { \ - using scalar_t = at::Half; \ - using param_t = float; \ - return __VA_ARGS__(); \ - } \ - default: \ - TORCH_CHECK(false, "Unsupported floating data type.\n"); \ - } \ - } else { \ - TORCH_CHECK(TYPE1 == TYPE2); \ - switch (TYPE1) { \ - case at::ScalarType::BFloat16: { \ - using scalar_t = at::BFloat16; \ - using param_t = at::BFloat16; \ - return __VA_ARGS__(); \ - } \ - case at::ScalarType::Half: { \ - using scalar_t = at::Half; \ - using param_t = at::Half; \ - return __VA_ARGS__(); \ - } \ - default: \ - TORCH_CHECK(false, "Unsupported floating data type.\n"); \ - } \ - } \ +#define CPU_DISPATCH_FLOATING_TYPES_EXT(TYPE1, TYPE2, ...) \ + [&] { \ + if (TYPE2 == at::kFloat) { \ + CPU_DISPATCH_TYPE1_WITH_PARAM(TYPE1, float, __VA_ARGS__) \ + } else if (TYPE2 == at::ScalarType::BFloat16) { \ + CPU_DISPATCH_TYPE1_WITH_PARAM(TYPE1, at::BFloat16, __VA_ARGS__) \ + } else if (TYPE2 == at::ScalarType::Half) { \ + CPU_DISPATCH_TYPE1_WITH_PARAM(TYPE1, at::Half, __VA_ARGS__) \ + } else { \ + TORCH_CHECK(false, "Unsupported floating data type."); \ + } \ }() #define UNUSED(x) (void)(x) diff --git a/sgl-kernel/csrc/cpu/topk.cpp b/sgl-kernel/csrc/cpu/topk.cpp index 100e87a7c..d3580b261 100644 --- a/sgl-kernel/csrc/cpu/topk.cpp +++ b/sgl-kernel/csrc/cpu/topk.cpp @@ -136,7 +136,7 @@ void grouped_topk_kernel_impl( }); } -template +template , int> = 0> inline void sigmoid(float* __restrict__ out, const scalar_t* __restrict__ input) { using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; @@ -157,6 +157,18 @@ inline void sigmoid(float* __restrict__ out, const scalar_t* __restrict__ input) } } +template , int> = 0> +inline void sigmoid(float* __restrict__ out, const float* __restrict__ input) { + using fVec = at::vec::Vectorized; + const fVec one = fVec(1.f); + constexpr int kVecSize = fVec::size(); + for (int d = 0; d < SIZE; d += kVecSize) { + fVec in_fvec = fVec::loadu(input + d); + in_fvec = one / (one + in_fvec.neg().exp_u20()); + in_fvec.store(out + d); + } +} + template void topk_sigmoid_kernel_impl( float* __restrict__ topk_weights, @@ -250,12 +262,11 @@ void topk_softmax_kernel_impl( }); } -template +template inline void apply_bias(float* __restrict__ scores2, const float* __restrict__ scores, const param_t* __restrict__ bias) { using fVec = at::vec::Vectorized; - using bVec = at::vec::Vectorized; - auto vec_size = bVec::size(); + auto vec_size = fVec::size() * 2; int d = 0; for (; d <= SIZE - vec_size; d += vec_size) { fVec bias0, bias1, x0, x1; @@ -275,14 +286,16 @@ template void biased_grouped_topk_kernel_impl( float* __restrict__ topk_weights, int32_t* __restrict__ topk_ids, - const scalar_t* __restrict__ gating_output, + scalar_t* __restrict__ gating_output, const param_t* __restrict__ bias, + float scaling_factor_value, int64_t num_tokens, int64_t num_groups, int64_t topk_group, bool renormalize) { using Vec = at::vec::Vectorized; + bool apply_scaling_factor = scaling_factor_value != 1.0f; const int64_t num_experts_per_group = NUM_EXPERTS / num_groups; at::parallel_for(0, num_tokens, 0, [&](int64_t begin, int64_t end) { // scores: sigmoid @@ -297,8 +310,7 @@ void biased_grouped_topk_kernel_impl( for (int64_t i = begin; i < end; ++i) { // do sigmoid to get scores sigmoid(scores, gating_output + i * NUM_EXPERTS); - - apply_bias(scores2, scores, bias); + apply_bias(scores2, scores, bias); for (int64_t g = 0; g < num_groups; ++g) { // find the max @@ -356,23 +368,35 @@ void biased_grouped_topk_kernel_impl( } #if defined(CPU_CAPABILITY_AVX512) - if (renormalize) { + if (renormalize || apply_scaling_factor) { __mmask16 mask = (1ULL << TOPK) - 1; __m512 x = _mm512_maskz_loadu_ps(mask, topk_weights + i * TOPK); - float sum = _mm512_reduce_add_ps(x); - __m512 vscale = _mm512_set1_ps(1.f / sum); - __m512 y = _mm512_mul_ps(x, vscale); - _mm512_mask_storeu_ps(topk_weights + i * TOPK, mask, y); + if (renormalize) { + float sum = _mm512_reduce_add_ps(x); + __m512 vscale = _mm512_set1_ps(scaling_factor_value / sum); + __m512 y = _mm512_mul_ps(x, vscale); + _mm512_mask_storeu_ps(topk_weights + i * TOPK, mask, y); + } else { + __m512 vscale = _mm512_set1_ps(scaling_factor_value); + __m512 y = _mm512_mul_ps(x, vscale); + _mm512_mask_storeu_ps(topk_weights + i * TOPK, mask, y); + } } #else - 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; + if (renormalize || apply_scaling_factor){ + if (renormalize) { + float sum = 0.f; + for (int64_t j = 0; j < TOPK; ++j) { + sum += topk_weights[i * TOPK + j]; + } + float scale = scaling_factor_value / sum; + for (int64_t j = 0; j < TOPK; ++j) { + topk_weights[i * TOPK + j] *= scale; + } + }else{ + for (int64_t j = 0; j < TOPK; ++j) { + topk_weights[i * TOPK + j] *= scaling_factor_value; + } } } #endif @@ -415,6 +439,7 @@ void biased_grouped_topk_kernel_impl( topk_ids.data_ptr(), \ gating_output.data_ptr(), \ correction_bias.data_ptr(), \ + scaling_factor_value, \ num_tokens, \ num_expert_group, \ topk_group, \ @@ -625,7 +650,7 @@ std::tuple biased_grouped_topk_cpu( int64_t num_fused_shared_experts, std::optional routed_scaling_factor, std::optional num_token_non_padded) { - // TODO: Will support num_fused_shared_experts, routed_scaling_factor and num_token_non_padded. + // TODO: Will support num_fused_shared_experts and num_token_non_padded. // For now, we just check them as default value. TORCH_CHECK( num_fused_shared_experts == 0, @@ -635,26 +660,30 @@ std::tuple biased_grouped_topk_cpu( !num_token_non_padded.has_value(), "num_token_non_padded must be None default value, got: ", num_token_non_padded.value()); - RECORD_FUNCTION( "sgl-kernel::biased_grouped_topk_cpu", std::vector({hidden_states, gating_output, correction_bias})); CHECK_INPUT(gating_output); CHECK_INPUT(correction_bias); - 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(correction_bias.numel() == num_experts, "Bias shape mismatch"); 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)); + float scaling_factor_value = routed_scaling_factor.has_value() ? routed_scaling_factor.value() : 1.0f; - CPU_DISPATCH_REDUCED_FLOATING_TYPES_EXT(st, correction_bias.scalar_type(), "biased_grouped_topk_kernel", [&] { + CPU_DISPATCH_FLOATING_TYPES_EXT(st, correction_bias.scalar_type(), "biased_grouped_topk_kernel", [&] { TORCH_CHECK(topk == 8, "Unexpected topk: ", topk); switch (num_experts) { + case 128: + LAUNCH_BIASED_GROUPED_TOPK_KERNEL(128, 8); + break; + case 192: + LAUNCH_BIASED_GROUPED_TOPK_KERNEL(192, 8); + break; case 256: LAUNCH_BIASED_GROUPED_TOPK_KERNEL(256, 8); break; diff --git a/test/srt/cpu/test_topk.py b/test/srt/cpu/test_topk.py index 9f3dfc1b4..c3c96af82 100644 --- a/test/srt/cpu/test_topk.py +++ b/test/srt/cpu/test_topk.py @@ -64,13 +64,22 @@ class TestGroupedTopK(CustomTestCase): # DeepSeek V2/V3/R1 uses biased_grouped_top class TestBiasedGroupedTopK(CustomTestCase): def _run_single_test( - self, M, E, G, topk, topk_group, renormalize, dtype, bias_dtype + self, + M, + E, + G, + topk, + topk_group, + renormalize, + gating_dtype, + bias_dtype, + routed_scaling_factor, ): - torch.manual_seed(1234) + torch.manual_seed(1024) # expand gating_output by M, otherwise bfloat16 fall into same value aftering truncating - hidden_states = torch.randn(M, 100, dtype=dtype) - gating_output = torch.randn(M, E, dtype=dtype) * 2 * M + hidden_states = torch.randn(M, 100, dtype=torch.bfloat16) + gating_output = torch.randn(M, E, dtype=gating_dtype) * 2 * M correction_bias = torch.randn(E, dtype=bias_dtype) ref_topk_weights, ref_topk_ids = native_biased_grouped_topk( @@ -82,7 +91,11 @@ class TestBiasedGroupedTopK(CustomTestCase): G, topk_group, ) - + ref_topk_weights = ( + ref_topk_weights * routed_scaling_factor + if routed_scaling_factor is not None + else ref_topk_weights + ) # fused version topk_weights, topk_ids = torch.ops.sgl_kernel.biased_grouped_topk_cpu( hidden_states, @@ -93,7 +106,7 @@ class TestBiasedGroupedTopK(CustomTestCase): G, topk_group, 0, - None, + routed_scaling_factor, None, ) @@ -104,11 +117,22 @@ class TestBiasedGroupedTopK(CustomTestCase): torch.testing.assert_close(res, ref) def test_biased_grouped_topk(self): - for renormalize in [True, False]: + for renormalize in [False]: for bias_dtype in [torch.float32, torch.bfloat16]: - self._run_single_test( - 122, 256, 8, 8, 2, renormalize, torch.bfloat16, bias_dtype - ) + for gating_dtype in [torch.float32, torch.bfloat16]: + for routed_scaling_factor in [None, 1.125]: + for E_num in [128, 192, 256, 384]: + self._run_single_test( + 34, + E_num, + 8, + 8, + 2, + renormalize, + gating_dtype, + bias_dtype, + routed_scaling_factor, + ) class TestTopK(CustomTestCase):