15 const std::optional<Tensor> weight,
float epsilon)
21 assert(input.ndim() >= 2 &&
22 "`FusedAddRmsNorm` requires `input` to have at least 2 dimensions");
24 "`FusedAddRmsNorm` requires a non-empty normalized dimension");
25 assert(input.shape() == residual.shape() &&
26 "`FusedAddRmsNorm` requires `input` and `residual` to have the same "
28 assert(input.dtype() == residual.dtype() &&
29 "`FusedAddRmsNorm` requires `input` and `residual` to have the same "
31 assert(input.stride(-1) == 1 &&
32 "`FusedAddRmsNorm` requires the last dimension of `input` to be "
34 assert(residual.stride(-1) == 1 &&
35 "`FusedAddRmsNorm` requires the last dimension of `residual` to be "
38 for (Tensor::Size i = 0; i + 2 < input.ndim(); ++i) {
39 assert(input.stride(i) == input.size(i + 1) * input.stride(i + 1) &&
40 "`FusedAddRmsNorm` requires `input` rows to have a uniform "
42 assert(residual.stride(i) ==
43 residual.size(i + 1) * residual.stride(i + 1) &&
44 "`FusedAddRmsNorm` requires `residual` rows to have a uniform "
48 if (weight.has_value()) {
49 assert(weight->ndim() == 1 && weight->size(0) ==
dim_ &&
50 "`FusedAddRmsNorm` requires 1D `weight` with size equal to the "
51 "normalized dimension");
52 assert(weight->dtype() == input.dtype() &&
53 "`FusedAddRmsNorm` requires `input` and `weight` to have the "
55 assert(weight->stride(0) == 1 &&
56 "`FusedAddRmsNorm` requires `weight` to be contiguous");