InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
elu_backward.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ELU_BACKWARD_H_
2#define INFINI_OPS_BASE_ELU_BACKWARD_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class EluBackward : public Operator<EluBackward> {
9 public:
10 EluBackward(const Tensor grad_output, const Tensor self_or_result,
11 const double alpha, const double scale, const double input_scale,
12 const bool is_result, Tensor grad_input)
13 : grad_output_shape_{grad_output.shape()},
14 grad_output_strides_{grad_output.strides()},
15 grad_output_type_{grad_output.dtype()},
16 self_or_result_shape_{self_or_result.shape()},
17 self_or_result_strides_{self_or_result.strides()},
18 self_or_result_type_{self_or_result.dtype()},
19 grad_input_shape_{grad_input.shape()},
20 grad_input_strides_{grad_input.strides()},
21 grad_input_type_{grad_input.dtype()},
22 alpha_{alpha},
23 scale_{scale},
24 input_scale_{input_scale},
25 is_result_{is_result},
26 device_index_{grad_input.device().index()} {}
27
28 virtual void operator()(const Tensor grad_output, const Tensor self_or_result,
29 const double alpha, const double scale,
30 const double input_scale, const bool is_result,
31 Tensor grad_input) const = 0;
32
33 protected:
34 Tensor::Shape grad_output_shape_;
35
36 Tensor::Strides grad_output_strides_;
37
39
40 Tensor::Shape self_or_result_shape_;
41
42 Tensor::Strides self_or_result_strides_;
43
45
46 Tensor::Shape grad_input_shape_;
47
48 Tensor::Strides grad_input_strides_;
49
51
52 double alpha_{};
53
54 double scale_{};
55
56 double input_scale_{};
57
58 bool is_result_{};
59
61};
62
63} // namespace infini::ops
64
65#endif
Definition elu_backward.h:8
virtual void operator()(const Tensor grad_output, const Tensor self_or_result, const double alpha, const double scale, const double input_scale, const bool is_result, Tensor grad_input) const =0
double scale_
Definition elu_backward.h:54
EluBackward(const Tensor grad_output, const Tensor self_or_result, const double alpha, const double scale, const double input_scale, const bool is_result, Tensor grad_input)
Definition elu_backward.h:10
Tensor::Shape self_or_result_shape_
Definition elu_backward.h:40
int device_index_
Definition elu_backward.h:60
bool is_result_
Definition elu_backward.h:58
Tensor::Strides self_or_result_strides_
Definition elu_backward.h:42
DataType self_or_result_type_
Definition elu_backward.h:44
Tensor::Shape grad_output_shape_
Definition elu_backward.h:34
double alpha_
Definition elu_backward.h:52
DataType grad_input_type_
Definition elu_backward.h:50
Tensor::Strides grad_output_strides_
Definition elu_backward.h:36
Tensor::Shape grad_input_shape_
Definition elu_backward.h:46
Tensor::Strides grad_input_strides_
Definition elu_backward.h:48
DataType grad_output_type_
Definition elu_backward.h:38
double input_scale_
Definition elu_backward.h:56
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8