InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
conv2d.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CONV2D_H_
2#define INFINI_OPS_BASE_CONV2D_H_
3
4#include <cstdint>
5#include <optional>
6#include <string>
7#include <vector>
8
9#include "common/op_utils/conv.h"
10
11namespace infini::ops {
12
13class Conv2d : public Operator<Conv2d> {
14 public:
15 Conv2d(const Tensor input, const Tensor weight, std::optional<Tensor> bias,
16 const std::vector<int64_t> stride, const std::string padding,
17 const std::vector<int64_t> dilation, const int64_t groups, Tensor out)
18 : metadata_{conv_detail::MakeMetadata<2>(
19 input, weight, bias, stride,
20 conv_detail::ResolvePadding<2>(weight, stride, padding, dilation),
21 dilation, groups, out)} {}
22
23 Conv2d(const Tensor input, const Tensor weight, std::optional<Tensor> bias,
24 const std::vector<int64_t> stride, const std::vector<int64_t> padding,
25 const std::vector<int64_t> dilation, const int64_t groups, Tensor out)
26 : metadata_{conv_detail::MakeMetadata<2>(
27 input, weight, bias, stride,
28 conv_detail::ResolvePadding<2>(weight, stride, padding, dilation),
29 dilation, groups, out)} {}
30
31 void operator()(const Tensor input, const Tensor weight,
32 std::optional<Tensor> bias, const std::vector<int64_t> stride,
33 const std::string padding,
34 const std::vector<int64_t> dilation, const int64_t groups,
35 Tensor out) const {
36 auto resolved =
37 conv_detail::ResolvePadding<2>(weight, stride, padding, dilation);
38 (*this)(input, weight, bias, stride, resolved.left, dilation, groups, out);
39 }
40
41 virtual void operator()(const Tensor input, const Tensor weight,
42 std::optional<Tensor> bias,
43 const std::vector<int64_t> stride,
44 const std::vector<int64_t> padding,
45 const std::vector<int64_t> dilation,
46 const int64_t groups, Tensor out) const = 0;
47
48 protected:
49 conv_detail::Metadata metadata_;
50};
51
52} // namespace infini::ops
53
54#endif
Definition conv2d.h:13
Conv2d(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > stride, const std::string padding, const std::vector< int64_t > dilation, const int64_t groups, Tensor out)
Definition conv2d.h:15
void operator()(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > stride, const std::string padding, const std::vector< int64_t > dilation, const int64_t groups, Tensor out) const
Definition conv2d.h:31
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 int64_t groups, Tensor out) const =0
Conv2d(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 int64_t groups, Tensor out)
Definition conv2d.h:23
conv_detail::Metadata metadata_
Definition conv2d.h:49
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8