InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
fused_add_rms_norm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FUSED_ADD_RMS_NORM_H_
2#define INFINI_OPS_BASE_FUSED_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
12class FusedAddRmsNorm : public Operator<FusedAddRmsNorm> {
13 public:
14 FusedAddRmsNorm(Tensor input, Tensor residual,
15 const std::optional<Tensor> weight, float epsilon)
16 : input_strides_{input.strides()},
17 residual_strides_{residual.strides()},
18 epsilon_{epsilon},
19 dim_{input.size(-1)},
20 num_tokens_{dim_ == 0 ? 0 : input.numel() / dim_} {
21 assert(input.ndim() >= 2 &&
22 "`FusedAddRmsNorm` requires `input` to have at least 2 dimensions");
23 assert(dim_ > 0 &&
24 "`FusedAddRmsNorm` requires a non-empty normalized dimension");
25 assert(input.shape() == residual.shape() &&
26 "`FusedAddRmsNorm` requires `input` and `residual` to have the same "
27 "shape");
28 assert(input.dtype() == residual.dtype() &&
29 "`FusedAddRmsNorm` requires `input` and `residual` to have the same "
30 "dtype");
31 assert(input.stride(-1) == 1 &&
32 "`FusedAddRmsNorm` requires the last dimension of `input` to be "
33 "contiguous");
34 assert(residual.stride(-1) == 1 &&
35 "`FusedAddRmsNorm` requires the last dimension of `residual` to be "
36 "contiguous");
37
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 "
41 "stride");
42 assert(residual.stride(i) ==
43 residual.size(i + 1) * residual.stride(i + 1) &&
44 "`FusedAddRmsNorm` requires `residual` rows to have a uniform "
45 "stride");
46 }
47
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 "
54 "same dtype");
55 assert(weight->stride(0) == 1 &&
56 "`FusedAddRmsNorm` requires `weight` to be contiguous");
57 }
58 }
59
60 virtual void operator()(Tensor input, Tensor residual,
61 const std::optional<Tensor> weight,
62 float epsilon) const = 0;
63
64 protected:
65 Tensor::Strides input_strides_;
66
67 Tensor::Strides residual_strides_;
68
69 float epsilon_{};
70
71 Tensor::Size dim_{0};
72
73 Tensor::Size num_tokens_{0};
74};
75
76} // namespace infini::ops
77
78#endif
Definition fused_add_rms_norm.h:12
float epsilon_
Definition fused_add_rms_norm.h:69
Tensor::Size num_tokens_
Definition fused_add_rms_norm.h:73
FusedAddRmsNorm(Tensor input, Tensor residual, const std::optional< Tensor > weight, float epsilon)
Definition fused_add_rms_norm.h:14
Tensor::Strides input_strides_
Definition fused_add_rms_norm.h:65
Tensor::Strides residual_strides_
Definition fused_add_rms_norm.h:67
Tensor::Size dim_
Definition fused_add_rms_norm.h:71
virtual void operator()(Tensor input, Tensor residual, const std::optional< Tensor > weight, float epsilon) const =0
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8