InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
convolution.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CONVOLUTION_H_
2#define INFINI_OPS_BASE_CONVOLUTION_H_
3
4#include <algorithm>
5#include <cassert>
6#include <cstdint>
7#include <optional>
8#include <vector>
9
10#include "common/op_utils/conv.h"
11
12namespace infini::ops {
13
14class Convolution : public Operator<Convolution> {
15 public:
16 Convolution(const Tensor input, const Tensor weight,
17 std::optional<Tensor> bias, const std::vector<int64_t> stride,
18 const std::vector<int64_t> padding,
19 const std::vector<int64_t> dilation, const bool transposed,
20 const std::vector<int64_t> output_padding, const int64_t groups,
21 Tensor out)
22 : metadata_{MakeMetadata(input, weight, bias, stride, padding, dilation,
23 transposed, output_padding, groups, out)} {}
24
25 virtual void operator()(const Tensor input, const Tensor weight,
26 std::optional<Tensor> bias,
27 const std::vector<int64_t> stride,
28 const std::vector<int64_t> padding,
29 const std::vector<int64_t> dilation,
30 const bool transposed,
31 const std::vector<int64_t> output_padding,
32 const int64_t groups, Tensor out) const = 0;
33
34 protected:
35 conv_detail::Metadata metadata_;
36
37 private:
38 static conv_detail::Metadata MakeMetadata(
39 const Tensor input, const Tensor weight, std::optional<Tensor> bias,
40 const std::vector<int64_t>& stride, const std::vector<int64_t>& padding,
41 const std::vector<int64_t>& dilation, const bool transposed,
42 const std::vector<int64_t>& output_padding, const int64_t groups,
43 Tensor out) {
44 assert((input.ndim() >= 3 && input.ndim() <= 5) &&
45 "operator `Convolution` currently supports only 1D, 2D, and 3D "
46 "inputs");
47 assert(!transposed &&
48 "operator `Convolution` does not currently support transposed "
49 "convolution");
50 assert(output_padding.size() + 2 == input.ndim() &&
51 "operator `Convolution` `output_padding` has the wrong length");
52 assert(std::all_of(output_padding.begin(), output_padding.end(),
53 [](int64_t value) { return value == 0; }) &&
54 "operator `Convolution` does not currently support nonzero "
55 "`output_padding` values");
56
57 switch (input.ndim()) {
58 case 3:
59 return conv_detail::MakeMetadata<1>(
60 input, weight, bias, stride,
61 conv_detail::ResolvePadding<1>(weight, stride, padding, dilation),
62 dilation, groups, out);
63 case 4:
64 return conv_detail::MakeMetadata<2>(
65 input, weight, bias, stride,
66 conv_detail::ResolvePadding<2>(weight, stride, padding, dilation),
67 dilation, groups, out);
68 case 5:
69 return conv_detail::MakeMetadata<3>(
70 input, weight, bias, stride,
71 conv_detail::ResolvePadding<3>(weight, stride, padding, dilation),
72 dilation, groups, out);
73 default:
74 return {};
75 }
76 }
77};
78
79} // namespace infini::ops
80
81#endif
Definition convolution.h:14
conv_detail::Metadata metadata_
Definition convolution.h:35
virtual void operator()(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > stride, const std::vector< int64_t > padding, const std::vector< int64_t > dilation, const bool transposed, const std::vector< int64_t > output_padding, const int64_t groups, Tensor out) const =0
Convolution(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > stride, const std::vector< int64_t > padding, const std::vector< int64_t > dilation, const bool transposed, const std::vector< int64_t > output_padding, const int64_t groups, Tensor out)
Definition convolution.h:16
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8