InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
rms_norm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_RMS_NORM_H_
2#define INFINI_OPS_BASE_RMS_NORM_H_
3
4#include <cstddef>
5#include <vector>
6
7#include "operator.h"
8#include "tensor.h"
9
10namespace infini::ops {
11
12class RmsNorm : public Operator<RmsNorm> {
13 public:
14 RmsNorm(const Tensor input, const Tensor weight, float eps, Tensor out)
15 : input_shape_{input.shape()},
16 out_shape_{out.shape()},
17 input_strides_{input.strides()},
18 out_strides_{out.strides()},
19 eps_{eps},
20 dim_{out.size(-1)},
21 ndim_{out.ndim()},
22 batch_size_{ndim_ == 2 ? out.size(-2) : out.size(-3)},
23 nhead_{ndim_ == 2 ? 1 : out.size(-2)} {
24 assert(input.dtype() == out.dtype());
25 }
26
27 RmsNorm(const Tensor input, const Tensor weight, Tensor out)
28 : RmsNorm{input, weight, 1e-6f, out} {}
29
30 // TODO: Type of `eps` should be `std::optional<float>` instead of `float`.
31 virtual void operator()(const Tensor input, const Tensor weight, float eps,
32 Tensor out) const = 0;
33
34 virtual void operator()(const Tensor input, const Tensor weight,
35 Tensor out) const {
36 return operator()(input, weight, eps_, out);
37 }
38
39 protected:
40 Tensor::Shape input_shape_;
41
42 Tensor::Shape out_shape_;
43
44 Tensor::Strides input_strides_;
45
46 Tensor::Strides out_strides_;
47
48 float eps_{1e-6f};
49
50 Tensor::Size dim_{0};
51
52 Tensor::Size ndim_{0};
53
54 Tensor::Size batch_size_{0};
55
56 Tensor::Size nhead_{1};
57};
58
59} // namespace infini::ops
60
61#endif
Definition generated/include/operator.h:282
Definition rms_norm.h:12
Tensor::Size batch_size_
Definition rms_norm.h:54
Tensor::Size ndim_
Definition rms_norm.h:52
virtual void operator()(const Tensor input, const Tensor weight, Tensor out) const
Definition rms_norm.h:34
RmsNorm(const Tensor input, const Tensor weight, Tensor out)
Definition rms_norm.h:27
Tensor::Shape input_shape_
Definition rms_norm.h:40
RmsNorm(const Tensor input, const Tensor weight, float eps, Tensor out)
Definition rms_norm.h:14
Tensor::Strides out_strides_
Definition rms_norm.h:46
virtual void operator()(const Tensor input, const Tensor weight, float eps, Tensor out) const =0
Tensor::Strides input_strides_
Definition rms_norm.h:44
float eps_
Definition rms_norm.h:48
Tensor::Shape out_shape_
Definition rms_norm.h:42
Tensor::Size nhead_
Definition rms_norm.h:56
Tensor::Size dim_
Definition rms_norm.h:50
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8