InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
add_rms_norm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ADD_RMS_NORM_H_
2#define INFINI_OPS_BASE_ADD_RMS_NORM_H_
3
4#include <cstddef>
5#include <optional>
6
7#include "operator.h"
8#include "tensor.h"
9
10namespace infini::ops {
11
12// Legacy out-of-place fused add + RMSNorm interface.
15class [[deprecated("Use `FusedAddRmsNorm` instead.")]] AddRmsNorm
16 : public Operator<AddRmsNorm> {
17 public:
18 AddRmsNorm(const Tensor input, const Tensor residual, const Tensor weight,
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)},
27 dim_{out.size(-1)},
28 ndim_{out.ndim()},
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 "
37 "shape");
38 assert(input.shape() == residual_out.shape() &&
39 "`AddRmsNorm` requires `input` and `residual_out` to have the "
40 "same shape");
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 "
48 "dtype");
49 assert(input.dtype() == residual_out.dtype() &&
50 "`AddRmsNorm` requires `input` and `residual_out` to have the same "
51 "dtype");
52 // The CUDA kernel indexes the normalized dimension with stride 1.
53 assert(input.stride(-1) == 1 &&
54 "`AddRmsNorm` requires the last dimension of `input` to be "
55 "contiguous");
56 assert(residual.stride(-1) == 1 &&
57 "`AddRmsNorm` requires the last dimension of `residual` to be "
58 "contiguous");
59 assert(out.stride(-1) == 1 &&
60 "`AddRmsNorm` requires the last dimension of `out` to be "
61 "contiguous");
62 assert(residual_out.stride(-1) == 1 &&
63 "`AddRmsNorm` requires the last dimension of `residual_out` to be "
64 "contiguous");
65 assert(weight.stride(-1) == 1 &&
66 "`AddRmsNorm` requires the last dimension of `weight` to be "
67 "contiguous");
68 }
69
70 virtual void operator()(const Tensor input, const Tensor residual,
71 const Tensor weight, std::optional<float> eps,
72 Tensor out, Tensor residual_out) const = 0;
73
74 virtual void operator()(const Tensor input, const Tensor residual,
75 const Tensor weight, Tensor out,
76 Tensor residual_out) const {
77 return operator()(input, residual, weight, std::nullopt, out, residual_out);
78 }
79
80 protected:
81 Tensor::Shape input_shape_;
82
83 Tensor::Shape out_shape_;
84
85 Tensor::Strides input_strides_;
86
87 Tensor::Strides residual_strides_;
88
89 Tensor::Strides out_strides_;
90
91 Tensor::Strides residual_out_strides_;
92
93 float eps_{1e-6f};
94
95 Tensor::Size dim_{0};
96
97 Tensor::Size ndim_{0};
98
99 Tensor::Size batch_size_{0};
100
101 Tensor::Size nhead_{1};
102};
103
104} // namespace infini::ops
105
106#endif
Definition add_rms_norm.h:16
Tensor::Strides input_strides_
Definition add_rms_norm.h:85
Tensor::Shape out_shape_
Definition add_rms_norm.h:83
Tensor::Strides residual_out_strides_
Definition add_rms_norm.h:91
Tensor::Shape input_shape_
Definition add_rms_norm.h:81
Tensor::Strides out_strides_
Definition add_rms_norm.h:89
AddRmsNorm(const Tensor input, const Tensor residual, const Tensor weight, std::optional< float > eps, Tensor out, Tensor residual_out)
Definition add_rms_norm.h:18
virtual void operator()(const Tensor input, const Tensor residual, const Tensor weight, Tensor out, Tensor residual_out) const
Definition add_rms_norm.h:74
Tensor::Strides residual_strides_
Definition add_rms_norm.h:87
virtual void operator()(const Tensor input, const Tensor residual, const Tensor weight, std::optional< float > eps, Tensor out, Tensor residual_out) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8