cubek_convolution/components/
problem.rs1use cubecl::{
2 ir::AddressType,
3 zspace::{Shape, Strides, shape},
4};
5use cubek_matmul::definition::{AccumulatorOperand, MatmulGlobalElems, MatmulProblem};
6use cubek_std::MatrixLayout;
7
8#[derive(Clone, Debug, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
9pub enum ConvolutionOperation {
10 Forward,
11 BackwardData,
12 BackwardWeight,
13 ForwardTransposed,
14}
15
16#[derive(Clone, Debug)]
17pub struct ConvolutionProblem {
19 pub m: usize,
20 pub n: usize,
21 pub k: usize,
22
23 pub lhs_strides: Strides,
24 pub rhs_strides: Strides,
25
26 pub lhs_layout: MatrixLayout,
27 pub rhs_layout: MatrixLayout,
28
29 pub kernel_size: Vec<u32>,
30 pub stride: Vec<u32>,
31 pub padding: Vec<i32>,
33 pub dilation: Vec<u32>,
34
35 pub batches: usize,
36 pub channels: usize,
37 pub out_channels: usize,
38 pub in_shape: Shape,
39 pub out_shape: Shape,
40
41 pub padded_channels: usize,
43 pub operation: ConvolutionOperation,
44
45 pub dimensionality: Dimensionality,
46
47 pub global_dtypes: MatmulGlobalElems,
48 pub address_type: AddressType,
50}
51
52impl ConvolutionProblem {
53 pub fn as_matmul_problem(&self, accumulator: AccumulatorOperand) -> MatmulProblem {
57 let rank = self.lhs_strides.len();
58
59 let lhs_strides = match self.lhs_layout {
65 MatrixLayout::RowMajor => self.lhs_strides.clone(),
66 MatrixLayout::ColMajor => {
67 let mut lhs_strides: Strides = self.lhs_strides[1..rank].into();
68 lhs_strides.push(self.lhs_strides[0]);
69 lhs_strides
70 }
71 };
72 let rhs_strides = match self.rhs_layout {
73 MatrixLayout::RowMajor => self.rhs_strides.clone(),
74 MatrixLayout::ColMajor => {
75 let mut rhs_strides: Strides = self.rhs_strides[1..rank].into();
76 rhs_strides.push(self.rhs_strides[0]);
77 rhs_strides
78 }
79 };
80
81 MatmulProblem {
82 m: self.m,
83 n: self.n,
84 k: self.k,
85 lhs_batches: shape![],
86 rhs_batches: shape![],
87 out_batches: shape![],
88 lhs_strides,
89 rhs_strides,
90 lhs_layout: self.lhs_layout,
91 rhs_layout: self.rhs_layout,
92 lhs_shape: shape![self.m, self.k],
93 rhs_shape: shape![self.k, self.n],
94 out_shape: shape![self.m, self.n],
95 out_strides: MatrixLayout::RowMajor.to_strides(&[self.m, self.n]),
96 out_layout: MatrixLayout::RowMajor,
97 lhs_scheme: None,
98 rhs_scheme: None,
99 global_dtypes: self.global_dtypes.clone(),
100 address_type: self.address_type,
101 accumulator,
102 }
103 }
104
105 pub fn should_check_channel(&self) -> bool {
106 self.channels != self.padded_channels
107 }
108
109 pub fn should_check_spatial_bounds(&self) -> bool {
110 spatial_bounds_required(
111 self.operation,
112 &self.kernel_size,
113 &self.stride,
114 &self.padding,
115 &self.dilation,
116 &self.in_shape,
117 &self.out_shape,
118 )
119 }
120}
121
122fn spatial_bounds_required(
123 operation: ConvolutionOperation,
124 kernel_size: &[u32],
125 stride: &[u32],
126 padding: &[i32],
127 dilation: &[u32],
128 in_shape: &[usize],
129 out_shape: &[usize],
130) -> bool {
131 (0..kernel_size.len()).any(|dim| {
132 let kernel_extent = (kernel_size[dim] as i64 - 1) * dilation[dim] as i64;
133 let padding = padding[dim] as i64;
134
135 match operation {
136 ConvolutionOperation::Forward | ConvolutionOperation::BackwardWeight => {
137 let first = -padding;
138 let last =
139 (out_shape[dim] as i64 - 1) * stride[dim] as i64 + kernel_extent - padding;
140 first < 0 || last >= in_shape[dim] as i64
141 }
142 ConvolutionOperation::ForwardTransposed | ConvolutionOperation::BackwardData => {
143 let first_numerator = padding - kernel_extent;
144 let last_numerator = in_shape[dim] as i64 - 1 + padding;
145 stride[dim] != 1
149 || first_numerator < 0
150 || last_numerator >= out_shape[dim] as i64 * stride[dim] as i64
151 }
152 }
153 })
154}
155
156#[cfg(test)]
157mod tests {
158 use super::{ConvolutionOperation, spatial_bounds_required};
159
160 #[test]
161 fn forward_checks_bounds_for_end_only_padding() {
162 assert!(spatial_bounds_required(
163 ConvolutionOperation::Forward,
164 &[3],
165 &[1],
166 &[0],
167 &[1],
168 &[5],
169 &[5],
170 ));
171 }
172
173 #[test]
174 fn forward_skips_bounds_for_exact_unpadded_geometry() {
175 assert!(!spatial_bounds_required(
176 ConvolutionOperation::Forward,
177 &[3],
178 &[1],
179 &[0],
180 &[1],
181 &[5],
182 &[3],
183 ));
184 }
185
186 #[test]
187 fn backward_data_checks_kernel_overhang_without_begin_padding() {
188 assert!(spatial_bounds_required(
189 ConvolutionOperation::BackwardData,
190 &[3],
191 &[1],
192 &[0],
193 &[1],
194 &[5],
195 &[3],
196 ));
197 }
198
199 #[test]
200 fn backward_data_skips_bounds_for_pointwise_geometry() {
201 assert!(!spatial_bounds_required(
202 ConvolutionOperation::BackwardData,
203 &[1],
204 &[1],
205 &[0],
206 &[1],
207 &[5],
208 &[5],
209 ));
210 }
211
212 #[test]
213 fn backward_data_checks_stride_divisibility() {
214 assert!(spatial_bounds_required(
215 ConvolutionOperation::BackwardData,
216 &[1],
217 &[2],
218 &[0],
219 &[1],
220 &[5],
221 &[3],
222 ));
223 }
224}
225
226#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug)]
228pub enum Dimensionality {
229 Dim1,
230 Dim2,
231 Dim3,
232}
233
234impl Dimensionality {
235 pub fn num_dims(&self) -> usize {
236 match self {
237 Dimensionality::Dim1 => 1,
238 Dimensionality::Dim2 => 2,
239 Dimensionality::Dim3 => 3,
240 }
241 }
242}