1#ifndef INFINI_OPS_BASE_MAX_UNPOOL2D_H_
2#define INFINI_OPS_BASE_MAX_UNPOOL2D_H_
8#include "common/op_utils/max_unpool.h"
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)
30 auto geometry = max_unpool_detail::ResolveGeometry<2>(
31 input, kernel_size, stride, padding, output_size);
37 [[deprecated(
"Use the `kernel_size` overload instead.")]]
39 const std::vector<int64_t> output_size,
Tensor out)
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,
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);
65 [[deprecated(
"Use the `kernel_size` overload instead.")]]
virtual void
67 const std::vector<int64_t> output_size,
Tensor out)
const = 0;
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