InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
fractional_max_pool3d_backward.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_FRACTIONAL_MAX_POOL3D_BACKWARD_H_
2#define INFINI_OPS_BASE_FRACTIONAL_MAX_POOL3D_BACKWARD_H_
3
4#include <vector>
5
6#include "operator.h"
7
8namespace infini::ops {
9
11 : public Operator<FractionalMaxPool3dBackward> {
12 public:
13 FractionalMaxPool3dBackward(const Tensor grad_output, const Tensor input,
14 const Tensor indices,
15 const std::vector<int64_t> kernel_size,
16 const std::vector<int64_t> output_size,
17 Tensor grad_input)
18 : grad_output_shape_{grad_output.shape()},
19 grad_output_strides_{grad_output.strides()},
20 grad_output_type_{grad_output.dtype()},
21 input_shape_{input.shape()},
22 input_strides_{input.strides()},
23 input_type_{input.dtype()},
24 indices_shape_{indices.shape()},
25 indices_strides_{indices.strides()},
26 indices_type_{indices.dtype()},
27 grad_input_shape_{grad_input.shape()},
28 grad_input_strides_{grad_input.strides()},
29 grad_input_type_{grad_input.dtype()},
30 kernel_size_{kernel_size},
31 output_size_{output_size},
32 device_index_{grad_input.device().index()} {}
33
34 virtual void operator()(const Tensor grad_output, const Tensor input,
35 const Tensor indices,
36 const std::vector<int64_t> kernel_size,
37 const std::vector<int64_t> output_size,
38 Tensor grad_input) const = 0;
39
40 protected:
41 Tensor::Shape grad_output_shape_;
42
43 Tensor::Strides grad_output_strides_;
44
46
47 Tensor::Shape input_shape_;
48
49 Tensor::Strides input_strides_;
50
51 DataType input_type_;
52
53 Tensor::Shape indices_shape_;
54
55 Tensor::Strides indices_strides_;
56
57 DataType indices_type_;
58
59 Tensor::Shape grad_input_shape_;
60
61 Tensor::Strides grad_input_strides_;
62
64
65 std::vector<int64_t> kernel_size_{};
66
67 std::vector<int64_t> output_size_{};
68
70};
71
72} // namespace infini::ops
73
74#endif
Definition fractional_max_pool3d_backward.h:11
virtual void operator()(const Tensor grad_output, const Tensor input, const Tensor indices, const std::vector< int64_t > kernel_size, const std::vector< int64_t > output_size, Tensor grad_input) const =0
FractionalMaxPool3dBackward(const Tensor grad_output, const Tensor input, const Tensor indices, const std::vector< int64_t > kernel_size, const std::vector< int64_t > output_size, Tensor grad_input)
Definition fractional_max_pool3d_backward.h:13
Tensor::Strides grad_input_strides_
Definition fractional_max_pool3d_backward.h:61
Tensor::Shape input_shape_
Definition fractional_max_pool3d_backward.h:47
DataType grad_input_type_
Definition fractional_max_pool3d_backward.h:63
DataType indices_type_
Definition fractional_max_pool3d_backward.h:57
Tensor::Shape indices_shape_
Definition fractional_max_pool3d_backward.h:53
Tensor::Shape grad_input_shape_
Definition fractional_max_pool3d_backward.h:59
Tensor::Strides indices_strides_
Definition fractional_max_pool3d_backward.h:55
int device_index_
Definition fractional_max_pool3d_backward.h:69
std::vector< int64_t > kernel_size_
Definition fractional_max_pool3d_backward.h:65
DataType grad_output_type_
Definition fractional_max_pool3d_backward.h:45
DataType input_type_
Definition fractional_max_pool3d_backward.h:51
std::vector< int64_t > output_size_
Definition fractional_max_pool3d_backward.h:67
Tensor::Strides grad_output_strides_
Definition fractional_max_pool3d_backward.h:43
Tensor::Strides input_strides_
Definition fractional_max_pool3d_backward.h:49
Tensor::Shape grad_output_shape_
Definition fractional_max_pool3d_backward.h:41
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8