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