InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
internal_native_batch_norm_legit.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_INTERNAL_NATIVE_BATCH_NORM_LEGIT_H_
2#define INFINI_OPS_BASE_INTERNAL_NATIVE_BATCH_NORM_LEGIT_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops::internal {
9
10class NativeBatchNormLegit : public Operator<NativeBatchNormLegit> {
11 public:
12 NativeBatchNormLegit(const Tensor input, const std::optional<Tensor> weight,
13 const std::optional<Tensor> bias, Tensor running_mean,
14 Tensor running_var, const bool training,
15 const double momentum, const double eps, Tensor out,
16 Tensor save_mean, Tensor save_invstd)
17 : input_shape_{input.shape()},
18 input_strides_{input.strides()},
19 input_type_{input.dtype()},
20 running_mean_shape_{running_mean.shape()},
21 running_mean_strides_{running_mean.strides()},
22 running_mean_type_{running_mean.dtype()},
23 running_var_shape_{running_var.shape()},
24 running_var_strides_{running_var.strides()},
25 running_var_type_{running_var.dtype()},
26 out_shape_{out.shape()},
27 out_strides_{out.strides()},
28 out_type_{out.dtype()},
29 save_mean_shape_{save_mean.shape()},
30 save_mean_strides_{save_mean.strides()},
31 save_mean_type_{save_mean.dtype()},
32 save_invstd_shape_{save_invstd.shape()},
33 save_invstd_strides_{save_invstd.strides()},
34 save_invstd_type_{save_invstd.dtype()},
35 has_weight_{weight.has_value()},
36 weight_shape_{weight ? Tensor::Shape{weight->shape()}
37 : Tensor::Shape{}},
38 weight_strides_{weight ? Tensor::Strides{weight->strides()}
39 : Tensor::Strides{}},
40 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
41 has_bias_{bias.has_value()},
42 bias_shape_{bias ? Tensor::Shape{bias->shape()} : Tensor::Shape{}},
43 bias_strides_{bias ? Tensor::Strides{bias->strides()}
44 : Tensor::Strides{}},
45 bias_type_{bias ? bias->dtype() : DataType::kFloat32},
46 training_{training},
47 momentum_{momentum},
48 eps_{eps},
49 device_index_{out.device().index()} {}
50
51 NativeBatchNormLegit(const Tensor input, const std::optional<Tensor> weight,
52 const std::optional<Tensor> bias, const bool training,
53 const double momentum, const double eps, Tensor out,
54 Tensor save_mean, Tensor save_invstd)
55 : input_shape_{input.shape()},
56 input_strides_{input.strides()},
57 input_type_{input.dtype()},
58 out_shape_{out.shape()},
59 out_strides_{out.strides()},
60 out_type_{out.dtype()},
61 save_mean_shape_{save_mean.shape()},
62 save_mean_strides_{save_mean.strides()},
63 save_mean_type_{save_mean.dtype()},
64 save_invstd_shape_{save_invstd.shape()},
65 save_invstd_strides_{save_invstd.strides()},
66 save_invstd_type_{save_invstd.dtype()},
67 has_weight_{weight.has_value()},
68 weight_shape_{weight ? Tensor::Shape{weight->shape()}
69 : Tensor::Shape{}},
70 weight_strides_{weight ? Tensor::Strides{weight->strides()}
71 : Tensor::Strides{}},
72 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
73 has_bias_{bias.has_value()},
74 bias_shape_{bias ? Tensor::Shape{bias->shape()} : Tensor::Shape{}},
75 bias_strides_{bias ? Tensor::Strides{bias->strides()}
76 : Tensor::Strides{}},
77 bias_type_{bias ? bias->dtype() : DataType::kFloat32},
78 training_{training},
79 momentum_{momentum},
80 eps_{eps},
81 device_index_{out.device().index()} {}
82
83 virtual void operator()(const Tensor input,
84 const std::optional<Tensor> weight,
85 const std::optional<Tensor> bias, Tensor running_mean,
86 Tensor running_var, const bool training,
87 const double momentum, const double eps, Tensor out,
88 Tensor save_mean, Tensor save_invstd) const = 0;
89
90 virtual void operator()(const Tensor input,
91 const std::optional<Tensor> weight,
92 const std::optional<Tensor> bias, const bool training,
93 const double momentum, const double eps, Tensor out,
94 Tensor save_mean, Tensor save_invstd) const = 0;
95
96 protected:
97 Tensor::Shape input_shape_;
98
99 Tensor::Strides input_strides_;
100
101 DataType input_type_;
102
103 Tensor::Shape running_mean_shape_;
104
105 Tensor::Strides running_mean_strides_;
106
108
109 Tensor::Shape running_var_shape_;
110
111 Tensor::Strides running_var_strides_;
112
114
115 Tensor::Shape out_shape_;
116
117 Tensor::Strides out_strides_;
118
119 DataType out_type_;
120
121 Tensor::Shape save_mean_shape_;
122
123 Tensor::Strides save_mean_strides_;
124
126
127 Tensor::Shape save_invstd_shape_;
128
129 Tensor::Strides save_invstd_strides_;
130
132
133 bool has_weight_{false};
134
135 Tensor::Shape weight_shape_;
136
137 Tensor::Strides weight_strides_;
138
139 DataType weight_type_{DataType::kFloat32};
140
141 bool has_bias_{false};
142
143 Tensor::Shape bias_shape_;
144
145 Tensor::Strides bias_strides_;
146
147 DataType bias_type_{DataType::kFloat32};
148
149 bool training_{};
150
151 double momentum_{};
152
153 double eps_{};
154
156};
157
158} // namespace infini::ops::internal
159
160#endif
Definition generated/include/operator.h:282
Definition internal_native_batch_norm_legit.h:10
DataType weight_type_
Definition internal_native_batch_norm_legit.h:139
int device_index_
Definition internal_native_batch_norm_legit.h:155
Tensor::Strides out_strides_
Definition internal_native_batch_norm_legit.h:117
Tensor::Shape save_mean_shape_
Definition internal_native_batch_norm_legit.h:121
NativeBatchNormLegit(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, const bool training, const double momentum, const double eps, Tensor out, Tensor save_mean, Tensor save_invstd)
Definition internal_native_batch_norm_legit.h:51
DataType out_type_
Definition internal_native_batch_norm_legit.h:119
DataType running_var_type_
Definition internal_native_batch_norm_legit.h:113
Tensor::Shape input_shape_
Definition internal_native_batch_norm_legit.h:97
Tensor::Strides save_mean_strides_
Definition internal_native_batch_norm_legit.h:123
Tensor::Shape running_mean_shape_
Definition internal_native_batch_norm_legit.h:103
Tensor::Strides running_mean_strides_
Definition internal_native_batch_norm_legit.h:105
Tensor::Shape out_shape_
Definition internal_native_batch_norm_legit.h:115
DataType bias_type_
Definition internal_native_batch_norm_legit.h:147
Tensor::Strides save_invstd_strides_
Definition internal_native_batch_norm_legit.h:129
virtual void operator()(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, const bool training, const double momentum, const double eps, Tensor out, Tensor save_mean, Tensor save_invstd) const =0
Tensor::Shape bias_shape_
Definition internal_native_batch_norm_legit.h:143
NativeBatchNormLegit(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, Tensor running_mean, Tensor running_var, const bool training, const double momentum, const double eps, Tensor out, Tensor save_mean, Tensor save_invstd)
Definition internal_native_batch_norm_legit.h:12
bool has_bias_
Definition internal_native_batch_norm_legit.h:141
Tensor::Strides running_var_strides_
Definition internal_native_batch_norm_legit.h:111
double momentum_
Definition internal_native_batch_norm_legit.h:151
DataType running_mean_type_
Definition internal_native_batch_norm_legit.h:107
Tensor::Strides bias_strides_
Definition internal_native_batch_norm_legit.h:145
Tensor::Strides input_strides_
Definition internal_native_batch_norm_legit.h:99
bool training_
Definition internal_native_batch_norm_legit.h:149
Tensor::Shape save_invstd_shape_
Definition internal_native_batch_norm_legit.h:127
bool has_weight_
Definition internal_native_batch_norm_legit.h:133
Tensor::Shape weight_shape_
Definition internal_native_batch_norm_legit.h:135
DataType save_mean_type_
Definition internal_native_batch_norm_legit.h:125
double eps_
Definition internal_native_batch_norm_legit.h:153
Tensor::Strides weight_strides_
Definition internal_native_batch_norm_legit.h:137
Tensor::Shape running_var_shape_
Definition internal_native_batch_norm_legit.h:109
DataType input_type_
Definition internal_native_batch_norm_legit.h:101
DataType save_invstd_type_
Definition internal_native_batch_norm_legit.h:131
virtual void operator()(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, Tensor running_mean, Tensor running_var, const bool training, const double momentum, const double eps, Tensor out, Tensor save_mean, Tensor save_invstd) const =0
Definition internal_add_relu.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8