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