InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
conv_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_CONV_INFINILM_H_
2#define INFINI_OPS_BASE_CONV_INFINILM_H_
3
4#include <cassert>
5#include <cstdint>
6#include <optional>
7#include <vector>
8
9#include "operator.h"
10
11namespace infini::ops {
12
15class [[deprecated(
16 "Migrate to an open-source-aligned operator when available.")]] ConvInfinilm
17 : public Operator<ConvInfinilm> {
18 public:
19 ConvInfinilm(const Tensor input, const Tensor weight,
20 std::optional<Tensor> bias, const std::vector<int64_t> padding,
21 const std::vector<int64_t> stride,
22 const std::vector<int64_t> dilation, const int64_t groups,
23 Tensor out)
24 : input_shape_{input.shape()},
25 input_strides_{input.strides()},
26 weight_shape_{weight.shape()},
27 weight_strides_{weight.strides()},
28 out_shape_{out.shape()},
29 out_strides_{out.strides()},
30 bias_shape_{bias.has_value() ? Tensor::Shape{bias->shape()}
31 : Tensor::Shape{}},
32 bias_strides_{bias.has_value() ? Tensor::Strides{bias->strides()}
33 : Tensor::Strides{}},
34 input_type_{input.dtype()},
35 weight_type_{weight.dtype()},
36 out_type_{out.dtype()},
37 bias_type_{bias.has_value() ? bias->dtype() : out.dtype()},
38 padding_{padding},
39 stride_{stride},
40 dilation_{dilation},
41 groups_{groups},
42 spatial_ndim_{input.ndim() - 2},
43 output_size_{out.numel()},
44 kernel_size_{1},
45 device_index_{out.device().index()},
46 has_bias_{bias.has_value()} {
47 assert(input.ndim() >= 3 && input.ndim() <= 5 &&
48 "`ConvInfinilm` supports 1D, 2D, and 3D conv_infinilmolution");
49 assert(input.ndim() == weight.ndim() && input.ndim() == out.ndim() &&
50 "`ConvInfinilm` input, weight, and output ranks must match");
51 assert(padding.size() == spatial_ndim_ && stride.size() == spatial_ndim_ &&
52 dilation.size() == spatial_ndim_ &&
53 "`ConvInfinilm` padding, stride, and dilation rank mismatch");
54 assert(groups > 0 && "`ConvInfinilm` groups must be positive");
55 assert(input_type_ == weight_type_ && input_type_ == out_type_ &&
56 "`ConvInfinilm` input, weight, and output dtypes must match");
57 assert(input_shape_[1] % groups == 0 &&
58 "`ConvInfinilm` input channels must be divisible by groups");
59 assert(weight_shape_[0] % groups == 0 &&
60 "`ConvInfinilm` output channels must be divisible by groups");
61 assert(weight_shape_[1] == input_shape_[1] / groups &&
62 "`ConvInfinilm` weight input channels mismatch");
63 assert(out_shape_[0] == input_shape_[0] &&
64 "`ConvInfinilm` output batch size mismatch");
65 assert(out_shape_[1] == weight_shape_[0] &&
66 "`ConvInfinilm` output channels mismatch");
67 assert(!out.HasBroadcastDim() &&
68 "`ConvInfinilm` output must not have broadcasted dimensions");
69
70 if (has_bias_) {
71 assert(bias_type_ == out_type_ && "`ConvInfinilm` bias dtype mismatch");
72 assert(bias_shape_.size() == 1 && bias_shape_[0] == out_shape_[1] &&
73 "`ConvInfinilm` bias shape must be `(out_channels,)`");
74 }
75
76 for (std::size_t i = 0; i < spatial_ndim_; ++i) {
77 assert(stride_[i] > 0 && "`ConvInfinilm` stride values must be positive");
78 assert(dilation_[i] > 0 &&
79 "`ConvInfinilm` dilation values must be positive");
80 assert(padding_[i] >= 0 &&
81 "`ConvInfinilm` padding values must be non-negative");
82
83 const auto expected = (input_shape_[i + 2] + 2 * padding_[i] -
84 dilation_[i] * (weight_shape_[i + 2] - 1) - 1) /
85 stride_[i] +
86 1;
87 assert(out_shape_[i + 2] == expected &&
88 "`ConvInfinilm` output spatial shape mismatch");
89 kernel_size_ *= weight_shape_[i + 2];
90 }
91 }
92
93 virtual void operator()(const Tensor input, const Tensor weight,
94 std::optional<Tensor> bias,
95 const std::vector<int64_t> padding,
96 const std::vector<int64_t> stride,
97 const std::vector<int64_t> dilation,
98 const int64_t groups, Tensor out) const = 0;
99
100 protected:
101 Tensor::Shape input_shape_;
102
103 Tensor::Strides input_strides_;
104
105 Tensor::Shape weight_shape_;
106
107 Tensor::Strides weight_strides_;
108
109 Tensor::Shape out_shape_;
110
111 Tensor::Strides out_strides_;
112
113 Tensor::Shape bias_shape_;
114
115 Tensor::Strides bias_strides_;
116
117 DataType input_type_;
118
119 DataType weight_type_;
120
121 DataType out_type_;
122
123 DataType bias_type_;
124
125 std::vector<int64_t> padding_;
126
127 std::vector<int64_t> stride_;
128
129 std::vector<int64_t> dilation_;
130
131 int64_t groups_{1};
132
133 Tensor::Size spatial_ndim_{0};
134
135 Tensor::Size output_size_{0};
136
137 Tensor::Size kernel_size_{1};
138
139 int device_index_{0};
140
141 bool has_bias_{false};
142};
143
144} // namespace infini::ops
145
146#endif
Definition conv_infinilm.h:17
DataType bias_type_
Definition conv_infinilm.h:123
ConvInfinilm(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > padding, const std::vector< int64_t > stride, const std::vector< int64_t > dilation, const int64_t groups, Tensor out)
Definition conv_infinilm.h:19
DataType weight_type_
Definition conv_infinilm.h:119
Tensor::Shape out_shape_
Definition conv_infinilm.h:109
Tensor::Shape input_shape_
Definition conv_infinilm.h:101
DataType input_type_
Definition conv_infinilm.h:117
std::vector< int64_t > stride_
Definition conv_infinilm.h:127
Tensor::Shape bias_shape_
Definition conv_infinilm.h:113
Tensor::Shape weight_shape_
Definition conv_infinilm.h:105
virtual void operator()(const Tensor input, const Tensor weight, std::optional< Tensor > bias, const std::vector< int64_t > padding, const std::vector< int64_t > stride, const std::vector< int64_t > dilation, const int64_t groups, Tensor out) const =0
Tensor::Strides bias_strides_
Definition conv_infinilm.h:115
Tensor::Strides input_strides_
Definition conv_infinilm.h:103
Tensor::Strides weight_strides_
Definition conv_infinilm.h:107
Tensor::Strides out_strides_
Definition conv_infinilm.h:111
std::vector< int64_t > padding_
Definition conv_infinilm.h:125
std::vector< int64_t > dilation_
Definition conv_infinilm.h:129
DataType out_type_
Definition conv_infinilm.h:121
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8