InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
max_unpool3d.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MAX_UNPOOL3D_H_
2#define INFINI_OPS_BASE_MAX_UNPOOL3D_H_
3
4#include <cstdint>
5#include <optional>
6#include <vector>
7
8#include "common/op_utils/max_unpool.h"
9
10namespace infini::ops {
11
12class MaxUnpool3d : public Operator<MaxUnpool3d> {
13 public:
14 MaxUnpool3d(const Tensor input, const Tensor indices,
15 const std::vector<int64_t> kernel_size,
16 const std::optional<std::vector<int64_t>> stride,
17 const std::vector<int64_t> padding,
18 const std::optional<std::vector<int64_t>> output_size, Tensor out)
19 : input_shape_{input.shape()},
20 input_strides_{input.strides()},
21 input_type_{input.dtype()},
22 indices_shape_{indices.shape()},
23 indices_strides_{indices.strides()},
24 indices_type_{indices.dtype()},
25 out_shape_{out.shape()},
26 out_strides_{out.strides()},
27 out_type_{out.dtype()},
29 stride_{},
30 padding_{padding},
31 device_index_{out.device().index()} {
32 auto geometry = max_unpool_detail::ResolveGeometry<3>(
33 input, kernel_size, stride, padding, output_size);
34 output_size_ = std::move(geometry.first);
35 stride_ = std::move(geometry.second);
36 }
37
40 [[deprecated("Use the `kernel_size` overload instead.")]]
41 MaxUnpool3d(const Tensor input, const Tensor indices,
42 const std::vector<int64_t> output_size,
43 const std::vector<int64_t> stride,
44 const std::vector<int64_t> padding, Tensor out)
45 : input_shape_{input.shape()},
46 input_strides_{input.strides()},
47 input_type_{input.dtype()},
48 indices_shape_{indices.shape()},
49 indices_strides_{indices.strides()},
50 indices_type_{indices.dtype()},
51 out_shape_{out.shape()},
52 out_strides_{out.strides()},
53 out_type_{out.dtype()},
54 output_size_{output_size},
55 stride_{stride},
56 padding_{padding},
57 device_index_{out.device().index()} {}
58
59 void operator()(const Tensor input, const Tensor indices,
60 const std::vector<int64_t> kernel_size,
61 const std::optional<std::vector<int64_t>> stride,
62 const std::vector<int64_t> padding,
63 const std::optional<std::vector<int64_t>> output_size,
64 Tensor out) const {
65 const auto geometry = max_unpool_detail::ResolveGeometry<3>(
66 input, kernel_size, stride, padding, output_size);
67 (*this)(input, indices, geometry.first, geometry.second, padding, out);
68 }
69
72 [[deprecated("Use the `kernel_size` overload instead.")]] virtual void
73 operator()(const Tensor input, const Tensor indices,
74 const std::vector<int64_t> output_size,
75 const std::vector<int64_t> stride,
76 const std::vector<int64_t> padding, Tensor out) const = 0;
77
78 protected:
79 Tensor::Shape input_shape_;
80
81 Tensor::Strides input_strides_;
82
83 DataType input_type_;
84
85 Tensor::Shape indices_shape_;
86
87 Tensor::Strides indices_strides_;
88
89 DataType indices_type_;
90
91 Tensor::Shape out_shape_;
92
93 Tensor::Strides out_strides_;
94
95 DataType out_type_;
96
97 std::vector<int64_t> output_size_{};
98
99 std::vector<int64_t> stride_{};
100
101 std::vector<int64_t> padding_{};
102
104};
105
106} // namespace infini::ops
107
108#endif
Definition max_unpool3d.h:12
Tensor::Strides out_strides_
Definition max_unpool3d.h:93
int device_index_
Definition max_unpool3d.h:103
Tensor::Shape out_shape_
Definition max_unpool3d.h:91
std::vector< int64_t > output_size_
Definition max_unpool3d.h:97
Tensor::Shape input_shape_
Definition max_unpool3d.h:79
Tensor::Strides input_strides_
Definition max_unpool3d.h:81
std::vector< int64_t > stride_
Definition max_unpool3d.h:99
DataType input_type_
Definition max_unpool3d.h:83
DataType indices_type_
Definition max_unpool3d.h:89
virtual void operator()(const Tensor input, const Tensor indices, const std::vector< int64_t > output_size, const std::vector< int64_t > stride, const std::vector< int64_t > padding, Tensor out) const =0
Tensor::Shape indices_shape_
Definition max_unpool3d.h:85
MaxUnpool3d(const Tensor input, const Tensor indices, const std::vector< int64_t > output_size, const std::vector< int64_t > stride, const std::vector< int64_t > padding, Tensor out)
Definition max_unpool3d.h:41
MaxUnpool3d(const Tensor input, const Tensor indices, const std::vector< int64_t > kernel_size, const std::optional< std::vector< int64_t > > stride, const std::vector< int64_t > padding, const std::optional< std::vector< int64_t > > output_size, Tensor out)
Definition max_unpool3d.h:14
void operator()(const Tensor input, const Tensor indices, const std::vector< int64_t > kernel_size, const std::optional< std::vector< int64_t > > stride, const std::vector< int64_t > padding, const std::optional< std::vector< int64_t > > output_size, Tensor out) const
Definition max_unpool3d.h:59
Tensor::Strides indices_strides_
Definition max_unpool3d.h:87
std::vector< int64_t > padding_
Definition max_unpool3d.h:101
DataType out_type_
Definition max_unpool3d.h:95
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8