1#ifndef INFINI_OPS_BASE_NATIVE_BATCH_NORM_H_
2#define INFINI_OPS_BASE_NATIVE_BATCH_NORM_H_
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,
35 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
40 bias_type_{bias ? bias->dtype() : DataType::kFloat32},
45 ?
Tensor::Strides{running_mean->strides()}
48 : DataType::kFloat32},
53 ?
Tensor::Strides{running_var->strides()}
56 : DataType::kFloat32},
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,
69 Tensor save_invstd)
const = 0;
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