InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
max_unpool2d.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_MAX_UNPOOL2D_H_
2#define INFINI_OPS_BASE_MAX_UNPOOL2D_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 MaxUnpool2d : public Operator<MaxUnpool2d> {
13 public:
14 MaxUnpool2d(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 device_index_{out.device().index()} {
30 auto geometry = max_unpool_detail::ResolveGeometry<2>(
31 input, kernel_size, stride, padding, output_size);
32 output_size_ = std::move(geometry.first);
33 }
34
37 [[deprecated("Use the `kernel_size` overload instead.")]]
38 MaxUnpool2d(const Tensor input, const Tensor indices,
39 const std::vector<int64_t> output_size, Tensor out)
40 : input_shape_{input.shape()},
41 input_strides_{input.strides()},
42 input_type_{input.dtype()},
43 indices_shape_{indices.shape()},
44 indices_strides_{indices.strides()},
45 indices_type_{indices.dtype()},
46 out_shape_{out.shape()},
47 out_strides_{out.strides()},
48 out_type_{out.dtype()},
49 output_size_{output_size},
50 device_index_{out.device().index()} {}
51
52 void operator()(const Tensor input, const Tensor indices,
53 const std::vector<int64_t> kernel_size,
54 const std::optional<std::vector<int64_t>> stride,
55 const std::vector<int64_t> padding,
56 const std::optional<std::vector<int64_t>> output_size,
57 Tensor out) const {
58 auto geometry = max_unpool_detail::ResolveGeometry<2>(
59 input, kernel_size, stride, padding, output_size);
60 (*this)(input, indices, std::move(geometry.first), out);
61 }
62
65 [[deprecated("Use the `kernel_size` overload instead.")]] virtual void
66 operator()(const Tensor input, const Tensor indices,
67 const std::vector<int64_t> output_size, Tensor out) const = 0;
68
69 protected:
70 Tensor::Shape input_shape_;
71
72 Tensor::Strides input_strides_;
73
74 DataType input_type_;
75
76 Tensor::Shape indices_shape_;
77
78 Tensor::Strides indices_strides_;
79
80 DataType indices_type_;
81
82 Tensor::Shape out_shape_;
83
84 Tensor::Strides out_strides_;
85
86 DataType out_type_;
87
88 std::vector<int64_t> output_size_{};
89
91};
92
93} // namespace infini::ops
94
95#endif
Definition max_unpool2d.h:12
std::vector< int64_t > output_size_
Definition max_unpool2d.h:88
DataType out_type_
Definition max_unpool2d.h:86
Tensor::Shape out_shape_
Definition max_unpool2d.h:82
Tensor::Strides input_strides_
Definition max_unpool2d.h:72
Tensor::Strides indices_strides_
Definition max_unpool2d.h:78
virtual void operator()(const Tensor input, const Tensor indices, const std::vector< int64_t > output_size, Tensor out) const =0
MaxUnpool2d(const Tensor input, const Tensor indices, const std::vector< int64_t > output_size, Tensor out)
Definition max_unpool2d.h:38
DataType indices_type_
Definition max_unpool2d.h:80
int device_index_
Definition max_unpool2d.h:90
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_unpool2d.h:52
Tensor::Shape input_shape_
Definition max_unpool2d.h:70
Tensor::Strides out_strides_
Definition max_unpool2d.h:84
MaxUnpool2d(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_unpool2d.h:14
Tensor::Shape indices_shape_
Definition max_unpool2d.h:76
DataType input_type_
Definition max_unpool2d.h:74
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8