InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
batch_norm_elemt.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_BATCH_NORM_ELEMT_H_
2#define INFINI_OPS_BASE_BATCH_NORM_ELEMT_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class BatchNormElemt : public Operator<BatchNormElemt> {
11 public:
12 BatchNormElemt(const Tensor input, const std::optional<Tensor> weight,
13 const std::optional<Tensor> bias, const Tensor mean,
14 const Tensor invstd, const double eps, Tensor out)
15 : input_shape_{input.shape()},
16 input_strides_{input.strides()},
17 input_type_{input.dtype()},
18 mean_shape_{mean.shape()},
19 mean_strides_{mean.strides()},
20 mean_type_{mean.dtype()},
21 invstd_shape_{invstd.shape()},
22 invstd_strides_{invstd.strides()},
23 invstd_type_{invstd.dtype()},
24 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 out_type_{out.dtype()},
27 has_weight_{weight.has_value()},
28 weight_shape_{weight ? Tensor::Shape{weight->shape()}
29 : Tensor::Shape{}},
30 weight_strides_{weight ? Tensor::Strides{weight->strides()}
31 : Tensor::Strides{}},
32 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
33 has_bias_{bias.has_value()},
34 bias_shape_{bias ? Tensor::Shape{bias->shape()} : Tensor::Shape{}},
35 bias_strides_{bias ? Tensor::Strides{bias->strides()}
36 : Tensor::Strides{}},
37 bias_type_{bias ? bias->dtype() : DataType::kFloat32},
38 eps_{eps},
39 device_index_{out.device().index()} {}
40
41 virtual void operator()(const Tensor input,
42 const std::optional<Tensor> weight,
43 const std::optional<Tensor> bias, const Tensor mean,
44 const Tensor invstd, const double eps,
45 Tensor out) const = 0;
46
47 protected:
48 Tensor::Shape input_shape_;
49
50 Tensor::Strides input_strides_;
51
52 DataType input_type_;
53
54 Tensor::Shape mean_shape_;
55
56 Tensor::Strides mean_strides_;
57
58 DataType mean_type_;
59
60 Tensor::Shape invstd_shape_;
61
62 Tensor::Strides invstd_strides_;
63
64 DataType invstd_type_;
65
66 Tensor::Shape out_shape_;
67
68 Tensor::Strides out_strides_;
69
70 DataType out_type_;
71
72 bool has_weight_{false};
73
74 Tensor::Shape weight_shape_;
75
76 Tensor::Strides weight_strides_;
77
78 DataType weight_type_{DataType::kFloat32};
79
80 bool has_bias_{false};
81
82 Tensor::Shape bias_shape_;
83
84 Tensor::Strides bias_strides_;
85
86 DataType bias_type_{DataType::kFloat32};
87
88 double eps_{};
89
91};
92
93} // namespace infini::ops
94
95#endif
Definition batch_norm_elemt.h:10
virtual void operator()(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, const Tensor mean, const Tensor invstd, const double eps, Tensor out) const =0
int device_index_
Definition batch_norm_elemt.h:90
Tensor::Shape weight_shape_
Definition batch_norm_elemt.h:74
bool has_weight_
Definition batch_norm_elemt.h:72
DataType out_type_
Definition batch_norm_elemt.h:70
bool has_bias_
Definition batch_norm_elemt.h:80
Tensor::Shape bias_shape_
Definition batch_norm_elemt.h:82
DataType bias_type_
Definition batch_norm_elemt.h:86
Tensor::Strides mean_strides_
Definition batch_norm_elemt.h:56
Tensor::Strides input_strides_
Definition batch_norm_elemt.h:50
Tensor::Shape out_shape_
Definition batch_norm_elemt.h:66
Tensor::Strides bias_strides_
Definition batch_norm_elemt.h:84
Tensor::Shape invstd_shape_
Definition batch_norm_elemt.h:60
DataType invstd_type_
Definition batch_norm_elemt.h:64
Tensor::Shape mean_shape_
Definition batch_norm_elemt.h:54
Tensor::Strides weight_strides_
Definition batch_norm_elemt.h:76
Tensor::Strides invstd_strides_
Definition batch_norm_elemt.h:62
Tensor::Shape input_shape_
Definition batch_norm_elemt.h:48
DataType weight_type_
Definition batch_norm_elemt.h:78
Tensor::Strides out_strides_
Definition batch_norm_elemt.h:68
double eps_
Definition batch_norm_elemt.h:88
BatchNormElemt(const Tensor input, const std::optional< Tensor > weight, const std::optional< Tensor > bias, const Tensor mean, const Tensor invstd, const double eps, Tensor out)
Definition batch_norm_elemt.h:12
DataType input_type_
Definition batch_norm_elemt.h:52
DataType mean_type_
Definition batch_norm_elemt.h:58
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8