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