19 std::optional<float> eps,
Tensor out,
Tensor residual_out)
20 : input_shape_{input.shape()},
21 out_shape_{out.shape()},
22 input_strides_{input.strides()},
23 residual_strides_{residual.strides()},
24 out_strides_{out.strides()},
25 residual_out_strides_{residual_out.strides()},
26 eps_{eps.value_or(1e-6f)},
29 batch_size_{ndim_ == 2 ? out.size(-2) : out.size(-3)},
30 nhead_{ndim_ == 2 ? 1 : out.size(-2)} {
31 assert((ndim_ == 2 || ndim_ == 3) &&
32 "`AddRmsNorm` supports 2D or 3D tensors only");
33 assert(input.shape() == out.shape() &&
34 "`AddRmsNorm` requires `input` and `out` to have the same shape");
35 assert(input.shape() == residual.shape() &&
36 "`AddRmsNorm` requires `input` and `residual` to have the same "
38 assert(input.shape() == residual_out.shape() &&
39 "`AddRmsNorm` requires `input` and `residual_out` to have the "
41 assert(weight.ndim() == 1 && weight.size(-1) == dim_ &&
42 "`AddRmsNorm` requires 1D `weight` with size equal to the "
43 "normalized dimension");
44 assert(input.dtype() == out.dtype() &&
45 "`AddRmsNorm` requires `input` and `out` to have the same dtype");
46 assert(input.dtype() == residual.dtype() &&
47 "`AddRmsNorm` requires `input` and `residual` to have the same "
49 assert(input.dtype() == residual_out.dtype() &&
50 "`AddRmsNorm` requires `input` and `residual_out` to have the same "
53 assert(input.stride(-1) == 1 &&
54 "`AddRmsNorm` requires the last dimension of `input` to be "
56 assert(residual.stride(-1) == 1 &&
57 "`AddRmsNorm` requires the last dimension of `residual` to be "
59 assert(out.stride(-1) == 1 &&
60 "`AddRmsNorm` requires the last dimension of `out` to be "
62 assert(residual_out.stride(-1) == 1 &&
63 "`AddRmsNorm` requires the last dimension of `residual_out` to be "
65 assert(weight.stride(-1) == 1 &&
66 "`AddRmsNorm` requires the last dimension of `weight` to be "