Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
113 changes: 113 additions & 0 deletions sgl-kernel/csrc/cpu/norm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -221,6 +221,85 @@ void fused_add_rmsnorm_kernel_impl(
});
}

template <typename scalar_t>
void fused_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<scalar_t>;
using fVec = at::vec::Vectorized<float>;
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<float>(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<scalar_t>(x_fvec0, x_fvec1);
out_bvec.store(out_ptr + d);
}
#pragma GCC unroll 4
for (; d < hidden_size; ++d) {
float x_val = static_cast<float>(input_ptr[d]);
float w_val = static_cast<float>(weight[d]);
float g_val = static_cast<float>(gate_ptr[d]);

out_ptr[d] = static_cast<scalar_t>(x_val * rsqrt_var * w_val * g_val / (1.f + std::exp(-g_val)));
}
}
});
}

} // anonymous namespace

// input : {batch_size, hidden_size}
Expand Down Expand Up @@ -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 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<c10::IValue>({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(), "fused_rmsnorm_gated_kernel", [&] {
fused_rmsnorm_gated_kernel_impl<scalar_t>(
output.data_ptr<scalar_t>(),
input.data_ptr<scalar_t>(),
weight.data_ptr<scalar_t>(),
gate.data_ptr<scalar_t>(),
batch_size,
hidden_size,
input_strideN,
eps);
});
return output;
}

// input : {batch_size, hidden_size}
// residual: {batch_size, hidden_size}
// weight : {hidden_size}
Expand Down
5 changes: 5 additions & 0 deletions sgl-kernel/csrc/cpu/torch_extension_cpu.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 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);

Expand Down Expand Up @@ -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("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);

Expand Down
46 changes: 46 additions & 0 deletions test/srt/cpu/test_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,5 +85,51 @@ def test_norm(self):
self._l2norm_test(*params)


class TestFusedRMSNormGated(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.fused_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()
Loading