[CPU] Add Gemma3RMSNorm kernel in sgl-kernel and add ut (#9324)
This commit is contained in:
@@ -408,6 +408,22 @@ class GemmaRMSNorm(CustomOp):
|
|||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
return self._forward_impl(x, residual)
|
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(
|
def forward_npu(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
@@ -445,6 +461,11 @@ class Gemma3RMSNorm(CustomOp):
|
|||||||
output = output * (1.0 + self.weight.float())
|
output = output * (1.0 + self.weight.float())
|
||||||
return output.type_as(x)
|
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):
|
def forward_cuda(self, x):
|
||||||
return self.forward_native(x)
|
return self.forward_native(x)
|
||||||
|
|
||||||
|
|||||||
@@ -65,7 +65,7 @@ void l2norm_kernel_impl(
|
|||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
template <typename scalar_t>
|
template <typename scalar_t, typename func_t, typename vec_func_t>
|
||||||
void rmsnorm_kernel_impl(
|
void rmsnorm_kernel_impl(
|
||||||
scalar_t* __restrict__ output,
|
scalar_t* __restrict__ output,
|
||||||
const scalar_t* __restrict__ input,
|
const scalar_t* __restrict__ input,
|
||||||
@@ -73,6 +73,8 @@ void rmsnorm_kernel_impl(
|
|||||||
int64_t batch_size,
|
int64_t batch_size,
|
||||||
int64_t hidden_size,
|
int64_t hidden_size,
|
||||||
int64_t input_strideN,
|
int64_t input_strideN,
|
||||||
|
const func_t& f,
|
||||||
|
const vec_func_t& vf,
|
||||||
float eps = 1e-5) {
|
float eps = 1e-5) {
|
||||||
using bVec = at::vec::Vectorized<scalar_t>;
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
using fVec = at::vec::Vectorized<float>;
|
using fVec = at::vec::Vectorized<float>;
|
||||||
@@ -117,8 +119,8 @@ void rmsnorm_kernel_impl(
|
|||||||
fVec w_fvec0, w_fvec1;
|
fVec w_fvec0, w_fvec1;
|
||||||
std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec);
|
std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec);
|
||||||
|
|
||||||
x_fvec0 = x_fvec0 * scale_fvec * w_fvec0;
|
x_fvec0 = x_fvec0 * scale_fvec * vf(w_fvec0);
|
||||||
x_fvec1 = x_fvec1 * scale_fvec * w_fvec1;
|
x_fvec1 = x_fvec1 * scale_fvec * vf(w_fvec1);
|
||||||
|
|
||||||
bVec out_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
|
bVec out_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
|
||||||
out_bvec.store(out_ptr + d);
|
out_bvec.store(out_ptr + d);
|
||||||
@@ -127,13 +129,93 @@ void rmsnorm_kernel_impl(
|
|||||||
for (; d < hidden_size; ++d) {
|
for (; d < hidden_size; ++d) {
|
||||||
float x_val = static_cast<float>(input_ptr[d]);
|
float x_val = static_cast<float>(input_ptr[d]);
|
||||||
float w_val = static_cast<float>(weight[d]);
|
float w_val = static_cast<float>(weight[d]);
|
||||||
out_ptr[d] = static_cast<scalar_t>(x_val * rsqrt_var * w_val);
|
out_ptr[d] = static_cast<scalar_t>(x_val * rsqrt_var * f(w_val));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
template <typename scalar_t>
|
template <typename scalar_t>
|
||||||
|
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<scalar_t>;
|
||||||
|
using fVec = at::vec::Vectorized<float>;
|
||||||
|
|
||||||
|
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<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);
|
||||||
|
|
||||||
|
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<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]);
|
||||||
|
out_ptr[d] = static_cast<scalar_t>(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 <typename scalar_t, typename func_t, typename vec_func_t>
|
||||||
void fused_add_rmsnorm_kernel_impl(
|
void fused_add_rmsnorm_kernel_impl(
|
||||||
scalar_t* __restrict__ input,
|
scalar_t* __restrict__ input,
|
||||||
scalar_t* __restrict__ residual,
|
scalar_t* __restrict__ residual,
|
||||||
@@ -142,6 +224,8 @@ void fused_add_rmsnorm_kernel_impl(
|
|||||||
int64_t batch_size,
|
int64_t batch_size,
|
||||||
int64_t hidden_size,
|
int64_t hidden_size,
|
||||||
int64_t input_strideN,
|
int64_t input_strideN,
|
||||||
|
const func_t& f,
|
||||||
|
const vec_func_t& vf,
|
||||||
float eps = 1e-5) {
|
float eps = 1e-5) {
|
||||||
using bVec = at::vec::Vectorized<scalar_t>;
|
using bVec = at::vec::Vectorized<scalar_t>;
|
||||||
using fVec = at::vec::Vectorized<float>;
|
using fVec = at::vec::Vectorized<float>;
|
||||||
@@ -207,14 +291,14 @@ void fused_add_rmsnorm_kernel_impl(
|
|||||||
fVec w_fvec0, w_fvec1;
|
fVec w_fvec0, w_fvec1;
|
||||||
std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec);
|
std::tie(w_fvec0, w_fvec1) = at::vec::convert_to_float(w_bvec);
|
||||||
|
|
||||||
x_fvec0 = x_fvec0 * scale_fvec * w_fvec0;
|
x_fvec0 = x_fvec0 * scale_fvec * vf(w_fvec0);
|
||||||
x_fvec1 = x_fvec1 * scale_fvec * w_fvec1;
|
x_fvec1 = x_fvec1 * scale_fvec * vf(w_fvec1);
|
||||||
bVec x_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
|
bVec x_bvec = convert_from_float_ext<scalar_t>(x_fvec0, x_fvec1);
|
||||||
x_bvec.store(input_ptr + d);
|
x_bvec.store(input_ptr + d);
|
||||||
}
|
}
|
||||||
#pragma GCC unroll 4
|
#pragma GCC unroll 4
|
||||||
for (; d < hidden_size; ++d) {
|
for (; d < hidden_size; ++d) {
|
||||||
float x_val = buffer_ptr[d] * rsqrt_var * static_cast<float>(weight[d]);
|
float x_val = buffer_ptr[d] * rsqrt_var * static_cast<float>(f(weight[d]));
|
||||||
input_ptr[d] = x_val;
|
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);
|
int64_t input_strideN = input.stride(0);
|
||||||
|
|
||||||
AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "rmsnorm_kernel", [&] {
|
AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "rmsnorm_kernel", [&] {
|
||||||
|
using Vec = at::vec::Vectorized<float>;
|
||||||
rmsnorm_kernel_impl<scalar_t>(
|
rmsnorm_kernel_impl<scalar_t>(
|
||||||
output.data_ptr<scalar_t>(),
|
output.data_ptr<scalar_t>(),
|
||||||
input.data_ptr<scalar_t>(),
|
input.data_ptr<scalar_t>(),
|
||||||
@@ -451,6 +536,8 @@ at::Tensor rmsnorm_cpu(at::Tensor& input, at::Tensor& weight, double eps) {
|
|||||||
batch_size,
|
batch_size,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
input_strideN,
|
input_strideN,
|
||||||
|
[](float x) { return x; },
|
||||||
|
[](Vec x) { return x; },
|
||||||
eps);
|
eps);
|
||||||
});
|
});
|
||||||
return output;
|
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<c10::IValue>({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<float>;
|
||||||
|
Vec one_vec = Vec(float(1));
|
||||||
|
rmsnorm_kernel_impl<scalar_t>(
|
||||||
|
output.data_ptr<scalar_t>(),
|
||||||
|
input.data_ptr<scalar_t>(),
|
||||||
|
weight.data_ptr<scalar_t>(),
|
||||||
|
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<c10::IValue>({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<float>;
|
||||||
|
Vec one_vec = Vec(float(1));
|
||||||
|
rmsnorm_kernel_impl<scalar_t>(
|
||||||
|
output.data_ptr<scalar_t>(),
|
||||||
|
input.data_ptr<scalar_t>(),
|
||||||
|
weight.data_ptr<scalar_t>(),
|
||||||
|
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<scalar_t>(
|
||||||
|
output.data_ptr<scalar_t>(),
|
||||||
|
input.data_ptr<scalar_t>(),
|
||||||
|
weight.data_ptr<scalar_t>(),
|
||||||
|
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}
|
// input : {batch_size, hidden_size}
|
||||||
// weight: {hidden_size}
|
// weight: {hidden_size}
|
||||||
// gate: {batch_size, 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::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", [&] {
|
AT_DISPATCH_REDUCED_FLOATING_TYPES(input.scalar_type(), "fused_add_rmsnorm_kernel", [&] {
|
||||||
|
using Vec = at::vec::Vectorized<float>;
|
||||||
fused_add_rmsnorm_kernel_impl<scalar_t>(
|
fused_add_rmsnorm_kernel_impl<scalar_t>(
|
||||||
input.data_ptr<scalar_t>(),
|
input.data_ptr<scalar_t>(),
|
||||||
residual.data_ptr<scalar_t>(),
|
residual.data_ptr<scalar_t>(),
|
||||||
@@ -551,6 +730,48 @@ void fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor&
|
|||||||
batch_size,
|
batch_size,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
input_strideN,
|
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<c10::IValue>({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<float>;
|
||||||
|
Vec one_vec = Vec(float(1));
|
||||||
|
fused_add_rmsnorm_kernel_impl<scalar_t>(
|
||||||
|
input.data_ptr<scalar_t>(),
|
||||||
|
residual.data_ptr<scalar_t>(),
|
||||||
|
weight.data_ptr<scalar_t>(),
|
||||||
|
buffer.data_ptr<float>(),
|
||||||
|
batch_size,
|
||||||
|
hidden_size,
|
||||||
|
input_strideN,
|
||||||
|
[](float x) { return x + 1; },
|
||||||
|
[one_vec](Vec x) { return x + one_vec; },
|
||||||
eps);
|
eps);
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -32,6 +32,8 @@ 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);
|
||||||
|
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
|
// layernorm
|
||||||
void layernorm_cpu(at::Tensor& input, at::Tensor& weight, double eps);
|
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
|
// 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);
|
||||||
|
void gemma_fused_add_rmsnorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps);
|
||||||
|
|
||||||
// fused_add_layernorm
|
// fused_add_layernorm
|
||||||
void fused_add_layernorm_cpu(at::Tensor& input, at::Tensor& residual, at::Tensor& weight, double eps);
|
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
|
// norm
|
||||||
m.def("rmsnorm_cpu(Tensor input, Tensor weight, float eps) -> Tensor");
|
m.def("rmsnorm_cpu(Tensor input, Tensor weight, float eps) -> Tensor");
|
||||||
m.impl("rmsnorm_cpu", torch::kCPU, &rmsnorm_cpu);
|
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.def("layernorm_cpu(Tensor(a!) input, Tensor weight, float eps) -> ()");
|
||||||
m.impl("layernorm_cpu", torch::kCPU, &layernorm_cpu);
|
m.impl("layernorm_cpu", torch::kCPU, &layernorm_cpu);
|
||||||
m.def("l2norm_cpu(Tensor input, float eps) -> Tensor");
|
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.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);
|
||||||
|
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.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);
|
m.impl("fused_add_layernorm_cpu", torch::kCPU, &fused_add_layernorm_cpu);
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,35 @@ class TestNorm(CustomTestCase):
|
|||||||
else:
|
else:
|
||||||
return x, residual
|
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):
|
def _norm_test(self, m, n, dtype):
|
||||||
|
|
||||||
x = torch.randn([m, n], dtype=dtype)
|
x = torch.randn([m, n], dtype=dtype)
|
||||||
@@ -78,11 +107,58 @@ class TestNorm(CustomTestCase):
|
|||||||
atol = rtol = precision[ref_out.dtype]
|
atol = rtol = precision[ref_out.dtype]
|
||||||
torch.testing.assert_close(ref_out, out, atol=atol, rtol=rtol)
|
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):
|
def test_norm(self):
|
||||||
for params in itertools.product(self.M, self.N, self.dtype):
|
for params in itertools.product(self.M, self.N, self.dtype):
|
||||||
with self.subTest(m=params[0], n=params[1], dtype=params[2]):
|
with self.subTest(m=params[0], n=params[1], dtype=params[2]):
|
||||||
self._norm_test(*params)
|
self._norm_test(*params)
|
||||||
self._l2norm_test(*params)
|
self._l2norm_test(*params)
|
||||||
|
self._gemma_rmsnorm_test(*params)
|
||||||
|
self._gemma3_rmsnorm_test(*params)
|
||||||
|
|
||||||
|
|
||||||
class TestFusedRMSNormGated(CustomTestCase):
|
class TestFusedRMSNormGated(CustomTestCase):
|
||||||
|
|||||||
Reference in New Issue
Block a user