diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index e543a63ee..68f94a090 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -408,6 +408,22 @@ class GemmaRMSNorm(CustomOp): ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: return self._forward_impl(x, residual) + def forward_cpu( + self, + x: torch.Tensor, + residual: Optional[torch.Tensor] = None, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + if _is_cpu_amx_available: + if residual is not None: + torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu( + x, residual, self.weight.data, self.variance_epsilon + ) + return x, residual + return torch.ops.sgl_kernel.gemma_rmsnorm_cpu( + x, self.weight.data, self.variance_epsilon + ) + return self.forward_native(x, residual) + def forward_npu( self, x: torch.Tensor, @@ -445,6 +461,11 @@ class Gemma3RMSNorm(CustomOp): output = output * (1.0 + self.weight.float()) return output.type_as(x) + def forward_cpu(self, x): + if _is_cpu_amx_available and x.stride(-1) == 1: + return torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, self.weight, self.eps) + return self.forward_native(x) + def forward_cuda(self, x): return self.forward_native(x) diff --git a/sgl-kernel/csrc/cpu/norm.cpp b/sgl-kernel/csrc/cpu/norm.cpp index 580c51fe6..d822c0d44 100644 --- a/sgl-kernel/csrc/cpu/norm.cpp +++ b/sgl-kernel/csrc/cpu/norm.cpp @@ -65,7 +65,7 @@ void l2norm_kernel_impl( } }); } -template +template void rmsnorm_kernel_impl( scalar_t* __restrict__ output, const scalar_t* __restrict__ input, @@ -73,6 +73,8 @@ void rmsnorm_kernel_impl( int64_t batch_size, int64_t hidden_size, int64_t input_strideN, + const func_t& f, + const vec_func_t& vf, float eps = 1e-5) { using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; @@ -117,8 +119,8 @@ void rmsnorm_kernel_impl( fVec w_fvec0, w_fvec1; std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec); - x_fvec0 = x_fvec0 * scale_fvec * w_fvec0; - x_fvec1 = x_fvec1 * scale_fvec * w_fvec1; + x_fvec0 = x_fvec0 * scale_fvec * vf(w_fvec0); + x_fvec1 = x_fvec1 * scale_fvec * vf(w_fvec1); bVec out_bvec = convert_from_float_ext(x_fvec0, x_fvec1); out_bvec.store(out_ptr + d); @@ -127,13 +129,93 @@ void rmsnorm_kernel_impl( for (; d < hidden_size; ++d) { float x_val = static_cast(input_ptr[d]); float w_val = static_cast(weight[d]); - out_ptr[d] = static_cast(x_val * rsqrt_var * w_val); + out_ptr[d] = static_cast(x_val * rsqrt_var * f(w_val)); } } }); } template +void gemma3_rmsnorm_kernel_4d_impl( + scalar_t* __restrict__ output, + const scalar_t* __restrict__ input, + const scalar_t* __restrict__ weight, + int64_t batch_size, + int64_t num_head, + int64_t seq_len, + int64_t hidden_size, + int64_t input_strideB, + int64_t input_strideH, + int64_t input_strideS, + int64_t output_strideB, + int64_t output_strideH, + int64_t output_strideS, + float eps = 1e-5) { + using bVec = at::vec::Vectorized; + using fVec = at::vec::Vectorized; + + constexpr int kVecSize = bVec::size(); + at::parallel_for(0, batch_size * num_head * seq_len, 0, [&](int64_t begin, int64_t end) { + int64_t bi{0}, hi{0}, si{0}; + data_index_init(begin, bi, batch_size, hi, num_head, si, seq_len); + for (int64_t i = begin; i < end; ++i) { + // local ptrs + scalar_t* __restrict__ out_ptr = output + bi * output_strideB + hi * output_strideH + si * output_strideS; + const scalar_t* __restrict__ input_ptr = input + bi * input_strideB + hi * input_strideH + si * input_strideS; + + fVec sum_fvec = fVec(float(0)); + float sum_val = float(0); + fVec one_fvec = fVec(float(1)); + + 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); + + x_fvec0 = x_fvec0 * scale_fvec * (w_fvec0 + one_fvec); + x_fvec1 = x_fvec1 * scale_fvec * (w_fvec1 + one_fvec); + + 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]); + out_ptr[d] = static_cast(x_val * rsqrt_var * (w_val + 1)); + } + // move to the next index + data_index_step(bi, batch_size, hi, num_head, si, seq_len); + } + }); +} + +template void fused_add_rmsnorm_kernel_impl( scalar_t* __restrict__ input, scalar_t* __restrict__ residual, @@ -142,6 +224,8 @@ void fused_add_rmsnorm_kernel_impl( int64_t batch_size, int64_t hidden_size, int64_t input_strideN, + const func_t& f, + const vec_func_t& vf, float eps = 1e-5) { using bVec = at::vec::Vectorized; using fVec = at::vec::Vectorized; @@ -207,14 +291,14 @@ void fused_add_rmsnorm_kernel_impl( fVec w_fvec0, w_fvec1; std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec); - x_fvec0 = x_fvec0 * scale_fvec * w_fvec0; - x_fvec1 = x_fvec1 * scale_fvec * w_fvec1; + x_fvec0 = x_fvec0 * scale_fvec * vf(w_fvec0); + x_fvec1 = x_fvec1 * scale_fvec * vf(w_fvec1); bVec x_bvec = convert_from_float_ext(x_fvec0, x_fvec1); x_bvec.store(input_ptr + d); } #pragma GCC unroll 4 for (; d < hidden_size; ++d) { - float x_val = buffer_ptr[d] * rsqrt_var * static_cast(weight[d]); + float x_val = buffer_ptr[d] * rsqrt_var * static_cast(f(weight[d])); input_ptr[d] = x_val; } } @@ -444,6 +528,7 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { int64_t input_strideN = input.stride(0); AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "rmsnorm_kernel", [&] { + using Vec = at::vec::Vectorized; rmsnorm_kernel_impl( output.data_ptr(), input.data_ptr(), @@ -451,6 +536,8 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { batch_size, hidden_size, input_strideN, + [](float x) { return x; }, + [](Vec x) { return x; }, eps); }); return output; @@ -485,6 +572,97 @@ void layernorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { }); } +at::Tensor gemma_rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { + RECORD_FUNCTION("sgl-kernel::gemma_rmsnorm_cpu", std::vector({input, weight})); + + CHECK_LAST_DIM_CONTIGUOUS_INPUT(input); + CHECK_INPUT(weight); + CHECK_DIM(2, input); + CHECK_DIM(1, weight); + CHECK_EQ(input.size(1), weight.size(0)); + int64_t batch_size = input.size(0); + int64_t hidden_size = input.size(1); + at::Tensor output = at::empty_like(input); + int64_t input_strideN = input.stride(0); + + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "gemma_rmsnorm_kernel", [&] { + using Vec = at::vec::Vectorized; + Vec one_vec = Vec(float(1)); + rmsnorm_kernel_impl( + output.data_ptr(), + input.data_ptr(), + weight.data_ptr(), + batch_size, + hidden_size, + input_strideN, + [](float x) { return x + 1; }, + [one_vec](Vec x) { return x + one_vec; }, + eps); + }); + return output; +} + +// input : {batch_size, hidden_size} or {batch_size, num_head, seq_len, head_dim} +// weight: {hidden_size} +at::Tensor gemma3_rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) { + RECORD_FUNCTION("sgl-kernel::gemma3_rmsnorm_cpu", std::vector({input, weight})); + + CHECK_LAST_DIM_CONTIGUOUS_INPUT(input); + CHECK_INPUT(weight); + TORCH_CHECK( + input.dim() == 2 || input.dim() == 4, "gemma3_rmsnorm_cpu: input must be 2D or 4D, got ", input.dim(), "D"); + CHECK_DIM(1, weight); + CHECK_EQ(input.size(-1), weight.size(0)); + int64_t batch_size = input.size(0); + int64_t hidden_size = weight.size(0); + at::Tensor output = at::empty_like(input); + if (input.dim() == 2) { + int64_t input_strideN = input.stride(0); + + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "gemma3_rmsnorm_kernel", [&] { + using Vec = at::vec::Vectorized; + Vec one_vec = Vec(float(1)); + rmsnorm_kernel_impl( + output.data_ptr(), + input.data_ptr(), + weight.data_ptr(), + batch_size, + hidden_size, + input_strideN, + [](float x) { return x + 1; }, + [one_vec](Vec x) { return x + one_vec; }, + eps); + }); + } else { + int64_t input_strideB = input.stride(0); + int64_t input_strideH = input.stride(1); + int64_t input_strideS = input.stride(2); + int64_t output_strideB = output.stride(0); + int64_t output_strideH = output.stride(1); + int64_t output_strideS = output.stride(2); + int64_t num_head = input.size(1); + int64_t seq_len = input.size(2); + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "gemma3_rmsnorm_kernel", [&] { + gemma3_rmsnorm_kernel_4d_impl( + output.data_ptr(), + input.data_ptr(), + weight.data_ptr(), + batch_size, + num_head, + seq_len, + hidden_size, + input_strideB, + input_strideH, + input_strideS, + output_strideB, + output_strideH, + output_strideS, + eps); + }); + } + return output; +} + // input : {batch_size, hidden_size} // weight: {hidden_size} // gate: {batch_size, hidden_size} @@ -543,6 +721,7 @@ void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& at::Tensor buffer = at::empty({num_threads, hidden_size}, input.options().dtype(at::kFloat)); AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "fused_add_rmsnorm_kernel", [&] { + using Vec = at::vec::Vectorized; fused_add_rmsnorm_kernel_impl( input.data_ptr(), residual.data_ptr(), @@ -551,6 +730,48 @@ void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& batch_size, hidden_size, input_strideN, + [](float x) { return x; }, + [](Vec x) { return x; }, + eps); + }); +} + +// input : {batch_size, hidden_size} +// residual: {batch_size, hidden_size} +// weight : {hidden_size} +void gemma_fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps) { + RECORD_FUNCTION("sgl-kernel::gemma_fused_add_rmsnorm_cpu", std::vector({input, residual, weight})); + CHECK_LAST_DIM_CONTIGUOUS_INPUT(input); + CHECK_INPUT(residual); + CHECK_INPUT(weight); + CHECK_DIM(2, input); + CHECK_DIM(2, residual); + CHECK_DIM(1, weight); + CHECK_EQ(input.size(0), residual.size(0)); + CHECK_EQ(input.size(1), residual.size(1)); + CHECK_EQ(input.size(1), weight.size(0)); + int64_t batch_size = input.size(0); + int64_t hidden_size = input.size(1); + int64_t input_strideN = input.stride(0); + + // allocate temp buffer to store x in float32 per thread + // TODO: implement a singleton for context + int64_t num_threads = at::get_num_threads(); + at::Tensor buffer = at::empty({num_threads, hidden_size}, input.options().dtype(at::kFloat)); + + AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "gemma_fused_add_rmsnorm_kernel", [&] { + using Vec = at::vec::Vectorized; + Vec one_vec = Vec(float(1)); + fused_add_rmsnorm_kernel_impl( + input.data_ptr(), + residual.data_ptr(), + weight.data_ptr(), + buffer.data_ptr(), + batch_size, + hidden_size, + input_strideN, + [](float x) { return x + 1; }, + [one_vec](Vec x) { return x + one_vec; }, eps); }); } diff --git a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp index bc8b8fcdd..428fa090d 100644 --- a/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp +++ b/sgl-kernel/csrc/cpu/torch_extension_cpu.cpp @@ -32,6 +32,8 @@ at::Tensor l2norm_cpu(at::Tensor& input, double eps); // rmsnorm at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); +at::Tensor gemma_rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); +at::Tensor gemma3_rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); // layernorm void layernorm_cpu(at::Tensor& input, at::Tensor& weight, double eps); @@ -41,6 +43,7 @@ at::Tensor fused_rmsnorm_gated_cpu(at::Tensor& input, at::Tensor& weight, at::Te // fused_add_rmsnorm void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps); +void gemma_fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps); // fused_add_layernorm void fused_add_layernorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps); @@ -330,6 +333,10 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { // norm m.def("rmsnorm_cpu(Tensor input, Tensor weight, float eps) -> Tensor"); m.impl("rmsnorm_cpu", torch::kCPU, &rmsnorm_cpu); + m.def("gemma_rmsnorm_cpu(Tensor input, Tensor weight, float eps) -> Tensor"); + m.impl("gemma_rmsnorm_cpu", torch::kCPU, &gemma_rmsnorm_cpu); + m.def("gemma3_rmsnorm_cpu(Tensor input, Tensor weight, float eps) -> Tensor"); + m.impl("gemma3_rmsnorm_cpu", torch::kCPU, &gemma3_rmsnorm_cpu); m.def("layernorm_cpu(Tensor(a!) input, Tensor weight, float eps) -> ()"); m.impl("layernorm_cpu", torch::kCPU, &layernorm_cpu); m.def("l2norm_cpu(Tensor input, float eps) -> Tensor"); @@ -338,6 +345,8 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { 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); + m.def("gemma_fused_add_rmsnorm_cpu(Tensor input, Tensor residual, Tensor weight, float eps) -> ()"); + m.impl("gemma_fused_add_rmsnorm_cpu", torch::kCPU, &gemma_fused_add_rmsnorm_cpu); m.def("fused_add_layernorm_cpu(Tensor(a!) input, Tensor residual, Tensor weight, float eps) -> ()"); m.impl("fused_add_layernorm_cpu", torch::kCPU, &fused_add_layernorm_cpu); diff --git a/test/srt/cpu/test_norm.py b/test/srt/cpu/test_norm.py index ade9a8c22..c118ab63e 100644 --- a/test/srt/cpu/test_norm.py +++ b/test/srt/cpu/test_norm.py @@ -36,6 +36,35 @@ class TestNorm(CustomTestCase): else: return x, residual + def _norm(self, x, eps): + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps) + + def _gemma3_rmsnorm_native( + self, x: torch.Tensor, weight: torch.Tensor, variance_epsilon: float = 1e-6 + ): + output = self._norm(x.float(), variance_epsilon) + output = output * (1.0 + weight.float()) + return output.type_as(x) + + def _gemma_rmsnorm_native( + self, + x: torch.Tensor, + weight: torch.Tensor, + variance_epsilon: float = 1e-6, + residual: Optional[torch.Tensor] = None, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + orig_dtype = x.dtype + if residual is not None: + x = x + residual + residual = x + + x = x.float() + variance = x.pow(2).mean(dim=-1, keepdim=True) + x = x * torch.rsqrt(variance + variance_epsilon) + x = x * (1.0 + weight.float()) + x = x.to(orig_dtype) + return x if residual is None else (x, residual) + def _norm_test(self, m, n, dtype): x = torch.randn([m, n], dtype=dtype) @@ -78,11 +107,58 @@ class TestNorm(CustomTestCase): atol = rtol = precision[ref_out.dtype] torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol) + def _gemma_rmsnorm_test(self, m, n, dtype): + + x = torch.randn([m, n], dtype=dtype) + x = make_non_contiguous(x) + hidden_size = x.size(-1) + weight = torch.randn(hidden_size, dtype=dtype) + variance_epsilon = 1e-6 + + out = torch.ops.sgl_kernel.gemma_rmsnorm_cpu(x, weight, variance_epsilon) + ref_out = self._gemma_rmsnorm_native(x, weight, variance_epsilon) + + atol = rtol = precision[ref_out.dtype] + torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol) + + ref_x = x.clone() + residual = torch.randn([m, hidden_size], dtype=dtype) + ref_residual = residual.clone() + + torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu( + x, residual, weight, variance_epsilon + ) + + ref_x, ref_residual = self._gemma_rmsnorm_native( + ref_x, weight, variance_epsilon, ref_residual + ) + + torch.testing.assert_close(x, ref_x, atol=atol, rtol=rtol) + torch.testing.assert_close(residual, ref_residual, atol=atol, rtol=rtol) + + def _gemma3_rmsnorm_test(self, m, n, dtype): + x_list = [ + torch.randn([m, n], dtype=dtype), + torch.randn([1, m, 2, n], dtype=dtype), + ] + for x in x_list: + x = make_non_contiguous(x) + hidden_size = x.size(-1) + weight = torch.randn(hidden_size, dtype=dtype) + variance_epsilon = 1e-6 + out = torch.ops.sgl_kernel.gemma3_rmsnorm_cpu(x, weight, variance_epsilon) + ref_out = self._gemma3_rmsnorm_native(x, weight, variance_epsilon) + + atol = rtol = precision[ref_out.dtype] + 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) self._l2norm_test(*params) + self._gemma_rmsnorm_test(*params) + self._gemma3_rmsnorm_test(*params) class TestFusedRMSNormGated(CustomTestCase):