Skip to main content

cubek_convolution/components/
problem.rs

1use 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)]
17/// Description of a matmul problem to solve, regardless of actual data
18pub 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    /// Padding at the beginning of each spatial dimension.
32    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    /// Channels after applying loader-specific padding
42    pub padded_channels: usize,
43    pub operation: ConvolutionOperation,
44
45    pub dimensionality: Dimensionality,
46
47    pub global_dtypes: MatmulGlobalElems,
48    /// Address type, defined as the max of each handle's `required_address_type`
49    pub address_type: AddressType,
50}
51
52impl ConvolutionProblem {
53    /// The bias is not part of the convolution problem: it arrives at launch, so
54    /// the caller states whether one is there. The matmul selector charges its
55    /// stage against the shared-memory budget.
56    pub fn as_matmul_problem(&self, accumulator: AccumulatorOperand) -> MatmulProblem {
57        let rank = self.lhs_strides.len();
58
59        // Strides are expected to be in row major (m, n) format so for matmul checks we need to
60        // convert them to that format, with all other dims treated as batch dims so they're still
61        // checked.
62        // lhs already has the right format, but rhs needs special handling.
63        // (h, w, c, n)
64        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                // A transposed convolution only has a source coordinate when the numerator is
146                // divisible by the stride. Non-unit strides therefore require the checked path
147                // even when every quotient falls within the source tensor.
148                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/// Spatial dimensionality of an operation
227#[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}