From 939971e0b1a91700beed7f8fc2250e0737681e5f Mon Sep 17 00:00:00 2001 From: "Jiang, Yanbing" Date: Wed, 17 Sep 2025 10:56:39 +0000 Subject: [PATCH 1/2] Add Qwen3NextRMSNormGated kernel --- sgl-kernel/csrc/cpu/norm.cpp | 113 ++++++++++++++++++++ sgl-kernel/csrc/cpu/torch_extension_cpu.cpp | 5 + test/srt/cpu/test_norm.py | 46 ++++++++ 3 files changed, 164 insertions(+) diff --git a/sgl-kernel/csrc/cpu/norm.cpp b/sgl-kernel/csrc/cpu/norm.cpp index 2c4e1f38d0bc..74769369477e 100644 --- a/sgl-kernel/csrc/cpu/norm.cpp +++ b/sgl-kernel/csrc/cpu/norm.cpp @@ -221,6 +221,85 @@ void fused_add_rmsnorm_kernel_impl( }); } +template +void qwen3_next_rmsnorm_gated_kernel_impl( + scalar_t* __restrict__ output, + const scalar_t* __restrict__ input, + const scalar_t* __restrict__ weight, + const scalar_t* __restrict__ gate, + int64_t batch_size, + int64_t hidden_size, + int64_t input_strideN, + float eps = 1e-5) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + const fVec one = fVec(1.f); + + constexpr int kVecSize = bVec::size(); + at::parallel_for(0, batch_size, 0, [&](int64_t begin, int64_t end) { + for (int64_t i = begin; i < end; ++i) { + // local ptrs + scalar_t* __restrict__ out_ptr = output + i * hidden_size; + const scalar_t* __restrict__ input_ptr = input + i * input_strideN; + const scalar_t* __restrict__ gate_ptr = gate + i * hidden_size; + + fVec sum_fvec = fVec(float(0)); + float sum_val = float(0); + + int64_t d; +#pragma GCC unroll 4 + for (d = 0; d <= hidden_size - kVecSize; d += kVecSize) { + bVec x_bvec = bVec::loadu(input_ptr + d); + fVec x_fvec0, x_fvec1; + std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); + + sum_fvec += x_fvec0 * x_fvec0; + sum_fvec += x_fvec1 * x_fvec1; + } +#pragma GCC unroll 4 + for (; d < hidden_size; ++d) { + float x_val = static_cast(input_ptr[d]); + sum_val += x_val * x_val; + } + + sum_val += vec_reduce_sum(sum_fvec); + float rsqrt_var = float(1) / std::sqrt(sum_val / hidden_size + eps); + const fVec scale_fvec = fVec(rsqrt_var); + +#pragma GCC unroll 4 + for (d = 0; d <= hidden_size - kVecSize; d += kVecSize) { + bVec x_bvec = bVec::loadu(input_ptr + d); + fVec x_fvec0, x_fvec1; + std::tie(x_fvec0, x_fvec1) = at::vec::convert_to_float(x_bvec); + + bVec w_bvec = bVec::loadu(weight + d); + fVec w_fvec0, w_fvec1; + std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec); + + bVec g_bvec = bVec::loadu(gate_ptr + d); + fVec g_fvec0, g_fvec1; + std::tie(g_fvec0, g_fvec1) = at::vec::convert_to_float(g_bvec); + g_fvec0 = g_fvec0 / (one + g_fvec0.neg().exp_u20()); + g_fvec1 = g_fvec1 / (one + g_fvec1.neg().exp_u20()); + + x_fvec0 = x_fvec0 * scale_fvec * w_fvec0 * g_fvec0; + x_fvec1 = x_fvec1 * scale_fvec * w_fvec1 * g_fvec1; + + bVec out_bvec = convert_from_float_ext(x_fvec0, x_fvec1); + out_bvec.store(out_ptr + d); + } +#pragma GCC unroll 4 + for (; d < hidden_size; ++d) { + float x_val = static_cast(input_ptr[d]); + float w_val = static_cast(weight[d]); + float g_val = static_cast(gate_ptr[d]); + + out_ptr[d] = static_cast(x_val * rsqrt_var * w_val * g_val / (1.f + std::exp(-g_val))); + } + } + }); +} + } // anonymous namespace // input : {batch_size, hidden_size} @@ -267,6 +346,40 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { return output; } +// input : {batch_size, hidden_size} +// weight: {hidden_size} +// gate: {batch_size, hidden_size} +at::Tensor qwen3_next_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps) { + RECORD_FUNCTION("sgl-kernel::qwen3_next_rmsnorm_gated_cpu", std::vector({input, weight, gate})); + + CHECK_LAST_DIM_CONTIGUOUS_INPUT(input); + CHECK_INPUT(weight); + CHECK_INPUT(gate); + CHECK_DIM(2, input); + CHECK_DIM(1, weight); + CHECK_DIM(2, gate); + CHECK_EQ(input.size(1), weight.size(0)); + int64_t batch_size = input.size(0); + int64_t hidden_size = input.size(1); + CHECK_EQ(input.size(0), gate.size(0)); + CHECK_EQ(input.size(1), gate.size(1)); + at::Tensor output = at::empty_like(input); + int64_t input_strideN = input.stride(0); + + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "qwen3_next_rmsnorm_gated_kernel", [&] { + qwen3_next_rmsnorm_gated_kernel_impl( + output.data_ptr(), + input.data_ptr(), + weight.data_ptr(), + gate.data_ptr(), + batch_size, + hidden_size, + input_strideN, + eps); + }); + return output; +} + // input : {batch_size, hidden_size} // residual: {batch_size, hidden_size} // weight : {hidden_size} diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index 2c8d9e3ececc..f35605b71edb 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -33,6 +33,9 @@ at::Tensor l2norm_cpu(at::Tensor& input, double eps); // rmsnorm at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); +// qwen3_next_rmsnorm_gated +at::Tensor qwen3_next_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps); + // fused_add_rmsnorm void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps); @@ -247,6 +250,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { m.impl("rmsnorm_cpu", torch::kCPU, &rmsnorm_cpu); m.def("l2norm_cpu(Tensor input, float eps) -> Tensor"); m.impl("l2norm_cpu", torch::kCPU, &l2norm_cpu); + m.def("qwen3_next_rmsnorm_gated_cpu(Tensor input, Tensor weight, Tensor gate, float eps) -> Tensor"); + m.impl("qwen3_next_rmsnorm_gated_cpu", torch::kCPU, &qwen3_next_rmsnorm_gated_cpu); m.def("fused_add_rmsnorm_cpu(Tensor(a!) input, Tensor residual, Tensor weight, float eps) -> ()"); m.impl("fused_add_rmsnorm_cpu", torch::kCPU, &fused_add_rmsnorm_cpu); diff --git a/test/srt/cpu/test_norm.py b/test/srt/cpu/test_norm.py index 75bacb198c76..621c8304b70c 100644 --- a/test/srt/cpu/test_norm.py +++ b/test/srt/cpu/test_norm.py @@ -86,5 +86,51 @@ def test_norm(self): self._l2norm_test(*params) +class TestQwen3NextRMSNormGated(CustomTestCase): + M = [4096, 1024] + N = [4096, 4096 + 13] + dtype = [torch.float16, torch.bfloat16] + + def _forward_native( + self, + hidden_states: torch.Tensor, + weight: torch.Tensor, + variance_epsilon: float = 1e-6, + gate: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + input_dtype = hidden_states.dtype + hidden_states = hidden_states.to(torch.float32) + variance = hidden_states.pow(2).mean(-1, keepdim=True) + # Norm before gate + hidden_states = hidden_states * torch.rsqrt(variance + variance_epsilon) + hidden_states = weight * hidden_states.to(input_dtype) + hidden_states = hidden_states * torch.nn.functional.silu(gate.to(torch.float32)) + + return hidden_states.to(input_dtype) + + def _norm_test(self, m, n, dtype): + + x = torch.randn([m, n], dtype=dtype) + x = make_non_contiguous(x) + batch_size = x.size(0) + hidden_size = x.size(-1) + weight = torch.randn(hidden_size, dtype=dtype) + variance_epsilon = 1e-6 + gate = torch.randn([batch_size, hidden_size], dtype=dtype) + + out = torch.ops.sgl_kernel.qwen3_next_rmsnorm_gated_cpu( + x, weight, gate, variance_epsilon + ) + ref_out = self._forward_native(x, weight, variance_epsilon, gate) + + atol = rtol = precision[ref_out.dtype] * 2 + torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol) + + def test_norm(self): + for params in itertools.product(self.M, self.N, self.dtype): + with self.subTest(m=params[0], n=params[1], dtype=params[2]): + self._norm_test(*params) + + if __name__ == "__main__": unittest.main() From b39b00777edca69e4a202f112a34d473d7c6f61b Mon Sep 17 00:00:00 2001 From: "Jiang, Yanbing" Date: Fri, 7 Nov 2025 13:48:39 +0000 Subject: [PATCH 2/2] Update func name --- sgl-kernel/csrc/cpu/norm.cpp | 10 +++++----- sgl-kernel/csrc/cpu/torch_extension_cpu.cpp | 6 +++--- test/srt/cpu/test_norm.py | 4 ++-- 3 files changed, 10 insertions(+), 10 deletions(-) diff --git a/sgl-kernel/csrc/cpu/norm.cpp b/sgl-kernel/csrc/cpu/norm.cpp index 74769369477e..239ee1b2283b 100644 --- a/sgl-kernel/csrc/cpu/norm.cpp +++ b/sgl-kernel/csrc/cpu/norm.cpp @@ -222,7 +222,7 @@ void fused_add_rmsnorm_kernel_impl( } template -void qwen3_next_rmsnorm_gated_kernel_impl( +void fused_rmsnorm_gated_kernel_impl( scalar_t* __restrict__ output, const scalar_t* __restrict__ input, const scalar_t* __restrict__ weight, @@ -349,8 +349,8 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { // input : {batch_size, hidden_size} // weight: {hidden_size} // gate: {batch_size, hidden_size} -at::Tensor qwen3_next_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps) { - RECORD_FUNCTION("sgl-kernel::qwen3_next_rmsnorm_gated_cpu", std::vector({input, weight, gate})); +at::Tensor fused_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps) { + RECORD_FUNCTION("sgl-kernel::fused_rmsnorm_gated_cpu", std::vector({input, weight, gate})); CHECK_LAST_DIM_CONTIGUOUS_INPUT(input); CHECK_INPUT(weight); @@ -366,8 +366,8 @@ at::Tensor qwen3_next_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, a at::Tensor output = at::empty_like(input); int64_t input_strideN = input.stride(0); - AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "qwen3_next_rmsnorm_gated_kernel", [&] { - qwen3_next_rmsnorm_gated_kernel_impl( + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "fused_rmsnorm_gated_kernel", [&] { + fused_rmsnorm_gated_kernel_impl( output.data_ptr(), input.data_ptr(), weight.data_ptr(), diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index f35605b71edb..e00bc319b794 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -34,7 +34,7 @@ at::Tensor l2norm_cpu(at::Tensor& input, double eps); at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); // qwen3_next_rmsnorm_gated -at::Tensor qwen3_next_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps); +at::Tensor fused_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Tensor& gate, double eps); // fused_add_rmsnorm void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps); @@ -250,8 +250,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { m.impl("rmsnorm_cpu", torch::kCPU, &rmsnorm_cpu); m.def("l2norm_cpu(Tensor input, float eps) -> Tensor"); m.impl("l2norm_cpu", torch::kCPU, &l2norm_cpu); - m.def("qwen3_next_rmsnorm_gated_cpu(Tensor input, Tensor weight, Tensor gate, float eps) -> Tensor"); - m.impl("qwen3_next_rmsnorm_gated_cpu", torch::kCPU, &qwen3_next_rmsnorm_gated_cpu); + m.def("fused_rmsnorm_gated_cpu(Tensor input, Tensor weight, Tensor gate, float eps) -> Tensor"); + m.impl("fused_rmsnorm_gated_cpu", torch::kCPU, &fused_rmsnorm_gated_cpu); m.def("fused_add_rmsnorm_cpu(Tensor(a!) input, Tensor residual, Tensor weight, float eps) -> ()"); m.impl("fused_add_rmsnorm_cpu", torch::kCPU, &fused_add_rmsnorm_cpu); diff --git a/test/srt/cpu/test_norm.py b/test/srt/cpu/test_norm.py index 621c8304b70c..1dcba48e3bc5 100644 --- a/test/srt/cpu/test_norm.py +++ b/test/srt/cpu/test_norm.py @@ -86,7 +86,7 @@ def test_norm(self): self._l2norm_test(*params) -class TestQwen3NextRMSNormGated(CustomTestCase): +class TestFusedRMSNormGated(CustomTestCase): M = [4096, 1024] N = [4096, 4096 + 13] dtype = [torch.float16, torch.bfloat16] @@ -118,7 +118,7 @@ def _norm_test(self, m, n, dtype): variance_epsilon = 1e-6 gate = torch.randn([batch_size, hidden_size], dtype=dtype) - out = torch.ops.sgl_kernel.qwen3_next_rmsnorm_gated_cpu( + out = torch.ops.sgl_kernel.fused_rmsnorm_gated_cpu( x, weight, gate, variance_epsilon ) ref_out = self._forward_native(x, weight, variance_epsilon, gate)