Add fused_rmsnorm_gated_cpu kernel for CPU to support Qwen3-Next (#11577)
This commit is contained in:
@@ -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
|
} // anonymous namespace
|
||||||
|
|
||||||
// input : {batch_size, hidden_size}
|
// input : {batch_size, hidden_size}
|
||||||
@@ -267,6 +346,40 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) {
|
|||||||
return output;
|
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}
|
// input : {batch_size, hidden_size}
|
||||||
// residual: {batch_size, hidden_size}
|
// residual: {batch_size, hidden_size}
|
||||||
// weight : {hidden_size}
|
// weight : {hidden_size}
|
||||||
|
|||||||
@@ -33,6 +33,9 @@ at::Tensor l2norm_cpu(at::Tensor& input, double eps);
|
|||||||
// rmsnorm
|
// rmsnorm
|
||||||
at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps);
|
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
|
// fused_add_rmsnorm
|
||||||
void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps);
|
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.impl("rmsnorm_cpu", torch::kCPU, &rmsnorm_cpu);
|
||||||
m.def("l2norm_cpu(Tensor input, float eps) -> Tensor");
|
m.def("l2norm_cpu(Tensor input, float eps) -> Tensor");
|
||||||
m.impl("l2norm_cpu", torch::kCPU, &l2norm_cpu);
|
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.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.impl("fused_add_rmsnorm_cpu", torch::kCPU, &fused_add_rmsnorm_cpu);
|
||||||
|
|
||||||
|
|||||||
@@ -85,5 +85,51 @@ class TestNorm(CustomTestCase):
|
|||||||
self._l2norm_test(*params)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user