diff --git a/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu b/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu index 7afed8c8f..ba12f4dd8 100644 --- a/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu +++ b/sgl-kernel/csrc/moe/moe_topk_softmax_kernels.cu @@ -44,6 +44,8 @@ using MaxReduceOp = cub::Max; using MinReduceOp = cub::Min; #endif +using cub_kvp = cub::KeyValuePair; + /// Aligned array type template < typename T, @@ -143,9 +145,122 @@ __launch_bounds__(TPB) __global__ void moeSoftmax( } } +namespace moe { +struct TopKPair { + static const int PAIR = 2; + static const int MAX_INDEX = 0; + cub_kvp max; + cub_kvp secondMax; + + __device__ TopKPair() {} + __device__ TopKPair(cub_kvp max, cub_kvp secondMax) : max(max), secondMax(secondMax) {} +}; + +struct TopKPairArgMax { + __device__ TopKPairArgMax() {} + __device__ __forceinline__ TopKPair operator()(const TopKPair& candidate1, const TopKPair& candidate2) const { + cub_kvp globalMax, globalSecondMax; + + // Determine the global maximum + if (candidate1.max.value > candidate2.max.value) { + globalMax = candidate1.max; + } else { + globalMax = candidate2.max; + } + + // Determine the global second maximum + if (globalMax.key == candidate1.max.key) { + // If candidate1 contributed the max, compare its secondMax with candidate2's max + globalSecondMax = (candidate1.secondMax.value > candidate2.max.value) ? candidate1.secondMax : candidate2.max; + } else { + // If candidate2 contributed the max, compare its secondMax with candidate1's max + globalSecondMax = (candidate2.secondMax.value > candidate1.max.value) ? candidate2.secondMax : candidate1.max; + } + return TopKPair(globalMax, globalSecondMax); + } +}; +} // namespace moe + +template +__launch_bounds__(TPB) __global__ void moeTopKFast( + float* inputs_after_softmax, + const bool* finished, + float* output, + int* indices, + const int num_experts, + const int k, + const int start_expert, + const int end_expert, + const bool renormalize) { + using namespace moe; + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage tmpStorage; + TopKPair thread_pair; + + const int block_row = blockIdx.x; + + const bool row_is_active = finished ? !finished[block_row] : true; + const int thread_read_offset = blockIdx.x * num_experts; + float row_sum_for_renormalize = 0; + // Each loop finds the top 2 elements, + // thus requiring only ⌈k/2⌉ loops (calculated as (k + 1) / 2). + for (int k_idx = 0; k_idx < (k + TopKPair::PAIR - 1) / TopKPair::PAIR; ++k_idx) { + // Initializing the top 2 elements by the minimum value. + thread_pair.max.key = 0; + thread_pair.max.value = -1.f; + thread_pair.secondMax.key = 0; + thread_pair.secondMax.value = -1.f; + + cub_kvp inp_kvp; + for (int expert = threadIdx.x; expert < num_experts; expert += TPB) { + const int idx = thread_read_offset + expert; + inp_kvp.key = expert; + inp_kvp.value = inputs_after_softmax[idx]; + // updating the thread_pair according to inp_kvp's value + if (inp_kvp.value > thread_pair.max.value) { + thread_pair.secondMax = thread_pair.max; + thread_pair.max = inp_kvp; + } else if (inp_kvp.value > thread_pair.secondMax.value) { + thread_pair.secondMax = inp_kvp; + } + } + + TopKPairArgMax reducer; + const TopKPair result_pair = BlockReduce(tmpStorage).Reduce(thread_pair, reducer); + if (threadIdx.x == 0) { +#pragma unroll + // updating 2 elements to the result. + for (int i = 0; i < TopKPair::PAIR; i++) { + if (k_idx * 2 + i >= k) break; + cub_kvp result = (i == TopKPair::MAX_INDEX) ? result_pair.max : result_pair.secondMax; + int expert = result.key; + bool node_uses_expert = expert >= start_expert && expert < end_expert; + bool should_process_row = row_is_active && node_uses_expert; + // The inputs_after_softmax is modified in-place to avoid unnecessary loops for finding the top k-1 value. + // 1.f represents the minimum value. + inputs_after_softmax[thread_read_offset + expert] = -1.f; + int idx = k * block_row + k_idx * 2 + i; + output[idx] = result.value; + indices[idx] = should_process_row ? (expert - start_expert) : num_experts; + assert(indices[idx] >= 0); + row_sum_for_renormalize += result.value; + } + } + __syncthreads(); + } + + if (renormalize && threadIdx.x == 0) { + float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize; + for (int k_idx = 0; k_idx < k; ++k_idx) { + const int idx = k * block_row + k_idx; + output[idx] = output[idx] * row_sum_for_renormalize_inv; + } + } +} + template __launch_bounds__(TPB) __global__ void moeTopK( - const float* inputs_after_softmax, + float* inputs_after_softmax, const bool* finished, float* output, int* indices, @@ -175,15 +290,6 @@ __launch_bounds__(TPB) __global__ void moeTopK( const int idx = thread_read_offset + expert; inp_kvp.key = expert; inp_kvp.value = inputs_after_softmax[idx]; - - for (int prior_k = 0; prior_k < k_idx; ++prior_k) { - const int prior_winning_expert = indices[k * block_row + prior_k]; - - if (prior_winning_expert == expert) { - inp_kvp = thread_kvp; - } - } - thread_kvp = arg_max(inp_kvp, thread_kvp); } @@ -199,6 +305,9 @@ __launch_bounds__(TPB) __global__ void moeTopK( indices[idx] = should_process_row ? (expert - start_expert) : num_experts; assert(indices[idx] >= 0); row_sum_for_renormalize += result_kvp.value; + // The inputs_after_softmax is modified in-place to avoid unnecessary loops for finding the top k-1 value. + // 1.f represents the minimum value. + inputs_after_softmax[thread_read_offset + expert] = -1.f; } __syncthreads(); } @@ -592,8 +701,15 @@ void topkGatingSoftmaxKernelLauncher( static constexpr int TPB = 256; moeSoftmax<<>>( gating_output, nullptr, softmax_workspace, num_experts, moe_softcapping, correction_bias); - moeTopK<<>>( - softmax_workspace, nullptr, topk_weights, topk_indices, num_experts, topk, 0, num_experts, renormalize); + if (topk == 1) { + // Note: As an optimization for better performance, + // the softmax_workspace is overwritten in-place by both moeTopK and moeTopKFast. + moeTopK<<>>( + softmax_workspace, nullptr, topk_weights, topk_indices, num_experts, topk, 0, num_experts, renormalize); + } else { + moeTopKFast<<>>( + softmax_workspace, nullptr, topk_weights, topk_indices, num_experts, topk, 0, num_experts, renormalize); + } } } } diff --git a/sgl-kernel/tests/test_moe_topk_softmax.py b/sgl-kernel/tests/test_moe_topk_softmax.py index b9a802c51..d6441a030 100644 --- a/sgl-kernel/tests/test_moe_topk_softmax.py +++ b/sgl-kernel/tests/test_moe_topk_softmax.py @@ -5,6 +5,50 @@ import torch from sgl_kernel import topk_softmax +def compare_topk_values(gating_output, topk_indices_ref, topk_indices): + values_ref = torch.gather(gating_output, 1, topk_indices_ref) + values = torch.gather(gating_output, 1, topk_indices) + return torch.equal(values_ref, values) + + +@pytest.mark.parametrize( + "num_tokens, num_experts, topk", + list( + itertools.product( + [1, 16, 128, 512, 1024, 2048], # num_tokens + [512], # num_experts + [1, 2, 3, 4, 5, 8], # topk + ) + ), +) +def test_topkfast_softmax(num_tokens, num_experts, topk): + gating_output = torch.randn( + (num_tokens, num_experts), dtype=torch.float32, device="cuda" + ) + + topk_weights = torch.empty((num_tokens, topk), dtype=torch.float32, device="cuda") + topk_indices = torch.empty((num_tokens, topk), dtype=torch.int32, device="cuda") + + topk_softmax( + topk_weights, + topk_indices, + gating_output, + ) + + # Native torch implementation + softmax_output = torch.softmax(gating_output, dim=-1) + topk_weights_ref, topk_indices_ref = torch.topk(softmax_output, topk, dim=-1) + + # Verify the top-k weights and indices match the torch native ones + assert torch.allclose( + topk_weights_ref, topk_weights, atol=1e-3, rtol=1e-3 + ), f"Weights mismatch: torch={topk_indices_ref} vs SGLang={topk_weights}" + + assert compare_topk_values( + gating_output, topk_indices_ref.int(), topk_indices + ), f"Values at the two indices are not equal: torch={topk_indices_ref}, SGLang={topk_indices}, values={gating_output}" + + @pytest.mark.parametrize( "num_tokens, num_experts, topk", list( @@ -38,9 +82,9 @@ def test_topk_softmax(num_tokens, num_experts, topk): topk_weights_ref, topk_weights, atol=1e-3, rtol=1e-3 ), f"Weights mismatch: torch={topk_indices_ref} vs SGLang={topk_weights}" - assert torch.allclose( - topk_indices_ref.int(), topk_indices, atol=0, rtol=0 - ), f"Indices mismatch: torch={topk_indices_ref}, SGLang={topk_indices}" + assert compare_topk_values( + gating_output, topk_indices_ref.int(), topk_indices + ), f"Values at the two indices are not equal: torch={topk_indices_ref}, SGLang={topk_indices}, values={gating_output}" @pytest.mark.parametrize( @@ -48,7 +92,7 @@ def test_topk_softmax(num_tokens, num_experts, topk): list( itertools.product( [1, 16, 128, 512, 1024, 2048], # num_tokens - [4, 8, 16, 32, 64, 128, 256], # num_experts + [4, 8, 16, 32, 64, 128, 256, 512], # num_experts [1, 2, 4], # topk [torch.float16, torch.bfloat16, torch.float32], # dtype ) @@ -81,9 +125,9 @@ def test_topk_softmax_dtype_regression(num_tokens, num_experts, topk, dtype): topk_weights_ref, topk_weights, atol=1e-3, rtol=1e-3 ), f"Weights mismatch: SGLang old interface={topk_indices_ref} vs SGLang new interface={topk_weights}" - assert torch.allclose( - topk_indices_ref.int(), topk_indices, atol=0, rtol=0 - ), f"Indices mismatch: SGLang old interface={topk_indices_ref}, SGLang new interface={topk_indices}" + assert compare_topk_values( + gating_output, topk_indices_ref.int(), topk_indices + ), f"Values at the two indices are not equal: torch={topk_indices_ref}, SGLang={topk_indices}, values={gating_output}" @pytest.mark.parametrize( @@ -91,7 +135,7 @@ def test_topk_softmax_dtype_regression(num_tokens, num_experts, topk, dtype): list( itertools.product( [1, 16, 128, 512, 1024, 2048], # num_tokens - [4, 8, 16, 32, 64, 128, 256], # num_experts + [4, 8, 16, 32, 64, 128, 256, 512], # num_experts [1, 2, 4], # topk ) ), @@ -130,9 +174,9 @@ def test_topk_softmax_renormalize(num_tokens, num_experts, topk): topk_weights_ref, topk_weights, atol=1e-3, rtol=1e-3 ), f"Weights mismatch: SGLang w/o fused renormalize={topk_indices_ref} vs SGLang w/ fused renormalize={topk_weights}" - assert torch.allclose( - topk_indices_ref.int(), topk_indices, atol=0, rtol=0 - ), f"Indices mismatch: SGLang w/o fused renormalize={topk_indices_ref}, SGLang w/ fused renormalize={topk_indices}" + assert compare_topk_values( + gating_output, topk_indices_ref.int(), topk_indices + ), f"Values at the two indices are not equal: torch={topk_indices_ref}, SGLang={topk_indices}, values={gating_output}" if __name__ == "__main__": diff --git a/test/srt/quant/test_block_int8.py b/test/srt/quant/test_block_int8.py index 1dc1edad5..eefab0265 100644 --- a/test/srt/quant/test_block_int8.py +++ b/test/srt/quant/test_block_int8.py @@ -94,7 +94,7 @@ def native_w8a8_block_int8_matmul(A, B, As, Bs, block_size, output_dtype=torch.f # For test -def torch_w8a8_block_int8_moe(a, w1, w2, w1_s, w2_s, score, topk, block_shape): +def torch_w8a8_block_int8_moe(a, w1, w2, w1_s, w2_s, topk_output, topk, block_shape): """This function performs fused moe with block-wise quantization using native torch.""" set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) @@ -102,8 +102,10 @@ def torch_w8a8_block_int8_moe(a, w1, w2, w1_s, w2_s, score, topk, block_shape): B, D = a.shape a = a.view(B, -1, D).repeat(1, topk, 1).reshape(-1, D) out = torch.zeros(B * topk, w2.shape[1], dtype=a.dtype, device=a.device) - score = torch.softmax(score, dim=-1, dtype=torch.float32) - topk_weight, topk_ids = torch.topk(score, topk) + # Use topk_output instead of torch.topk for consistent equal-value handling + # moeTopK kernel and torch.topk may differ in tie-breaking for equal values + topk_weight, topk_ids = topk_output.topk_weights, topk_output.topk_ids + topk_weight = topk_weight.view(-1) topk_ids = topk_ids.view(-1) @@ -183,7 +185,7 @@ class TestW8A8BlockINT8FusedMoE(CustomTestCase): with torch.inference_mode(): ref_out = torch_w8a8_block_int8_moe( - a, w1, w2, w1_s, w2_s, score, topk, block_size + a, w1, w2, w1_s, w2_s, topk_output, topk, block_size ) out = fused_moe( a,