From 2424303dfb96047266871702b9219a953aae4c0d Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Tue, 19 May 2026 09:26:10 +0800 Subject: [PATCH] [codex] Optimize hidden-size 512 RMSNorm dispatch (#24710) Co-authored-by: Codex --- .../jit_kernel/csrc/elementwise/rmsnorm.cuh | 50 +++++++++++-------- python/sglang/jit_kernel/norm.py | 2 + .../sglang/jit_kernel/tests/test_rmsnorm.py | 2 +- 3 files changed, 33 insertions(+), 21 deletions(-) diff --git a/python/sglang/jit_kernel/csrc/elementwise/rmsnorm.cuh b/python/sglang/jit_kernel/csrc/elementwise/rmsnorm.cuh index 2e1edd692..4dba06914 100644 --- a/python/sglang/jit_kernel/csrc/elementwise/rmsnorm.cuh +++ b/python/sglang/jit_kernel/csrc/elementwise/rmsnorm.cuh @@ -85,17 +85,22 @@ __global__ __launch_bounds__(kDim / 16) void rmsnorm_cta_double(const RMSNormPar } sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto warp_id = threadIdx.x / kWarpThreads; - smem[warp_id] = sum_of_squares; - __syncthreads(); - if (warp_id == 0) { - const auto tx = threadIdx.x; - const auto local_sum = tx < kNumWarps ? smem[tx] : 0.0f; - sum_of_squares = warp::reduce_sum(local_sum); - smem[tx] = math::rsqrt(sum_of_squares / kDim + eps); + float norm_factor; + if constexpr (kNumWarps == 1) { + norm_factor = math::rsqrt(sum_of_squares / kDim + eps); + } else { + const auto warp_id = threadIdx.x / kWarpThreads; + smem[warp_id] = sum_of_squares; + __syncthreads(); + if (warp_id == 0) { + const auto tx = threadIdx.x; + const auto local_sum = tx < kNumWarps ? smem[tx] : 0.0f; + sum_of_squares = warp::reduce_sum(local_sum); + smem[tx] = math::rsqrt(sum_of_squares / kDim + eps); + } + __syncthreads(); + norm_factor = smem[warp_id]; } - __syncthreads(); - const float norm_factor = smem[warp_id]; Storage output_first, output_second; #pragma unroll @@ -147,17 +152,22 @@ __global__ __launch_bounds__(kDim / 16) void rmsnorm_cta_wide(const RMSNormParam } sum_of_squares = warp::reduce_sum(sum_of_squares); - const auto warp_id = threadIdx.x / kWarpThreads; - smem[warp_id] = sum_of_squares; - __syncthreads(); - if (warp_id == 0) { - const auto tx = threadIdx.x; - const auto local_sum = tx < kNumWarps ? smem[tx] : 0.0f; - sum_of_squares = warp::reduce_sum(local_sum); - smem[tx] = math::rsqrt(sum_of_squares / kDim + eps); + float norm_factor; + if constexpr (kNumWarps == 1) { + norm_factor = math::rsqrt(sum_of_squares / kDim + eps); + } else { + const auto warp_id = threadIdx.x / kWarpThreads; + smem[warp_id] = sum_of_squares; + __syncthreads(); + if (warp_id == 0) { + const auto tx = threadIdx.x; + const auto local_sum = tx < kNumWarps ? smem[tx] : 0.0f; + sum_of_squares = warp::reduce_sum(local_sum); + smem[tx] = math::rsqrt(sum_of_squares / kDim + eps); + } + __syncthreads(); + norm_factor = smem[warp_id]; } - __syncthreads(); - const float norm_factor = smem[warp_id]; Storage output_vec; #pragma unroll diff --git a/python/sglang/jit_kernel/norm.py b/python/sglang/jit_kernel/norm.py index 25b4a5f2c..f1e19fea0 100644 --- a/python/sglang/jit_kernel/norm.py +++ b/python/sglang/jit_kernel/norm.py @@ -46,6 +46,8 @@ def _is_supported_rmsnorm_hidden_size(d: int) -> bool: def _rmsnorm_kernel_class(hidden_size: int) -> str: if hidden_size in _RMSNORM_WARP_SIZES: return "RMSNormWarpKernel" + if hidden_size == 512: + return "RMSNormHalfKernel" if hidden_size >= _RMSNORM_HALF_BLOCK_MIN_SIZE: if hidden_size % 512 == 0: return "RMSNormHalfKernel" diff --git a/python/sglang/jit_kernel/tests/test_rmsnorm.py b/python/sglang/jit_kernel/tests/test_rmsnorm.py index 38b2efc90..37539ee19 100644 --- a/python/sglang/jit_kernel/tests/test_rmsnorm.py +++ b/python/sglang/jit_kernel/tests/test_rmsnorm.py @@ -88,7 +88,7 @@ def test_rmsnorm_hidden_size_support(hidden_size: int) -> None: (64, "RMSNormWarpKernel"), (128, "RMSNormWarpKernel"), (256, "RMSNormWarpKernel"), - (512, "RMSNormKernel"), + (512, "RMSNormHalfKernel"), (1536, "RMSNormKernel"), (2048, "RMSNormHalfKernel"), (2304, "RMSNormKernel"), # NOTE: not 512 aligned