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,
22 :
metadata_{MakeMetadata(input, weight, bias, stride, padding, dilation,
23 transposed, output_padding, groups, out)} {}
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;
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,
44 assert((input.ndim() >= 3 && input.ndim() <= 5) &&
45 "operator `Convolution` currently supports only 1D, 2D, and 3D "
48 "operator `Convolution` does not currently support transposed "
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");
57 switch (input.ndim()) {
59 return conv_detail::MakeMetadata<1>(
60 input, weight, bias, stride,
61 conv_detail::ResolvePadding<1>(weight, stride, padding, dilation),
62 dilation, groups, out);
64 return conv_detail::MakeMetadata<2>(
65 input, weight, bias, stride,
66 conv_detail::ResolvePadding<2>(weight, stride, padding, dilation),
67 dilation, groups, out);
69 return conv_detail::MakeMetadata<3>(
70 input, weight, bias, stride,
71 conv_detail::ResolvePadding<3>(weight, stride, padding, dilation),
72 dilation, groups, out);
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