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,
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()}
32 bias_strides_{bias.has_value() ?
Tensor::Strides{bias->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()},
42 spatial_ndim_{input.ndim() - 2},
43 output_size_{out.numel()},
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");
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,)`");
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");
83 const auto expected = (input_shape_[i + 2] + 2 * padding_[i] -
84 dilation_[i] * (weight_shape_[i + 2] - 1) - 1) /
87 assert(out_shape_[i + 2] == expected &&
88 "`ConvInfinilm` output spatial shape mismatch");
89 kernel_size_ *= weight_shape_[i + 2];