InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
cudnn_convolution.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CUDNN_CONVOLUTION_H_
2#define INFINI_OPS_BASE_CUDNN_CONVOLUTION_H_
3
4#include <vector>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class CudnnConvolution : public Operator<CudnnConvolution> {
11 public:
12 CudnnConvolution(const Tensor input, const Tensor weight,
13 const std::vector<int64_t> padding,
14 const std::vector<int64_t> stride,
15 const std::vector<int64_t> dilation, const int64_t groups,
16 const bool benchmark, const bool deterministic,
17 const bool allow_tf32, Tensor out)
18 : input_shape_{input.shape()},
19 input_strides_{input.strides()},
20 input_type_{input.dtype()},
21 weight_shape_{weight.shape()},
22 weight_strides_{weight.strides()},
23 weight_type_{weight.dtype()},
24 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 out_type_{out.dtype()},
27 padding_{padding},
28 stride_{stride},
29 dilation_{dilation},
30 groups_{groups},
31 benchmark_{benchmark},
32 deterministic_{deterministic},
33 allow_tf32_{allow_tf32},
34 device_index_{out.device().index()} {}
35
36 virtual void operator()(const Tensor input, const Tensor weight,
37 const std::vector<int64_t> padding,
38 const std::vector<int64_t> stride,
39 const std::vector<int64_t> dilation,
40 const int64_t groups, const bool benchmark,
41 const bool deterministic, const bool allow_tf32,
42 Tensor out) const = 0;
43
44 protected:
45 Tensor::Shape input_shape_;
46
47 Tensor::Strides input_strides_;
48
49 DataType input_type_;
50
51 Tensor::Shape weight_shape_;
52
53 Tensor::Strides weight_strides_;
54
55 DataType weight_type_;
56
57 Tensor::Shape out_shape_;
58
59 Tensor::Strides out_strides_;
60
61 DataType out_type_;
62
63 std::vector<int64_t> padding_{};
64
65 std::vector<int64_t> stride_{};
66
67 std::vector<int64_t> dilation_{};
68
69 int64_t groups_{};
70
71 bool benchmark_{};
72
74
76
78};
79
80} // namespace infini::ops
81
82#endif
Definition cudnn_convolution.h:10
bool deterministic_
Definition cudnn_convolution.h:73
virtual void operator()(const Tensor input, const Tensor weight, const std::vector< int64_t > padding, const std::vector< int64_t > stride, const std::vector< int64_t > dilation, const int64_t groups, const bool benchmark, const bool deterministic, const bool allow_tf32, Tensor out) const =0
std::vector< int64_t > padding_
Definition cudnn_convolution.h:63
int64_t groups_
Definition cudnn_convolution.h:69
int device_index_
Definition cudnn_convolution.h:77
Tensor::Shape out_shape_
Definition cudnn_convolution.h:57
std::vector< int64_t > stride_
Definition cudnn_convolution.h:65
bool allow_tf32_
Definition cudnn_convolution.h:75
CudnnConvolution(const Tensor input, const Tensor weight, const std::vector< int64_t > padding, const std::vector< int64_t > stride, const std::vector< int64_t > dilation, const int64_t groups, const bool benchmark, const bool deterministic, const bool allow_tf32, Tensor out)
Definition cudnn_convolution.h:12
DataType weight_type_
Definition cudnn_convolution.h:55
bool benchmark_
Definition cudnn_convolution.h:71
DataType out_type_
Definition cudnn_convolution.h:61
DataType input_type_
Definition cudnn_convolution.h:49
Tensor::Shape input_shape_
Definition cudnn_convolution.h:45
Tensor::Shape weight_shape_
Definition cudnn_convolution.h:51
Tensor::Strides out_strides_
Definition cudnn_convolution.h:59
Tensor::Strides weight_strides_
Definition cudnn_convolution.h:53
Tensor::Strides input_strides_
Definition cudnn_convolution.h:47
std::vector< int64_t > dilation_
Definition cudnn_convolution.h:67
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8