Skip to main content

cubek_convolution/components/
problem.rs

1use cubecl::{
2    ir::AddressType,
3    zspace::{Shape, Strides, shape},
4};
5use cubek_matmul::definition::{MatmulGlobalElems, MatmulProblem, MatrixLayout};
6
7#[derive(Clone, Debug, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
8pub enum ConvolutionOperation {
9    Forward,
10    BackwardData,
11    BackwardWeight,
12    ForwardTransposed,
13}
14
15#[derive(Clone, Debug)]
16/// Description of a matmul problem to solve, regardless of actual data
17pub struct ConvolutionProblem {
18    pub m: usize,
19    pub n: usize,
20    pub k: usize,
21
22    pub lhs_strides: Strides,
23    pub rhs_strides: Strides,
24
25    pub lhs_layout: MatrixLayout,
26    pub rhs_layout: MatrixLayout,
27
28    pub kernel_size: Vec<u32>,
29    pub stride: Vec<u32>,
30    pub padding: Vec<i32>,
31    pub dilation: Vec<u32>,
32
33    pub batches: usize,
34    pub channels: usize,
35    pub out_channels: usize,
36    pub in_shape: Shape,
37    pub out_shape: Shape,
38
39    /// Channels after applying loader-specific padding
40    pub padded_channels: usize,
41    pub operation: ConvolutionOperation,
42
43    pub dimensionality: Dimensionality,
44
45    pub global_dtypes: MatmulGlobalElems,
46    /// Address type, defined as the max of each handle's `required_address_type`
47    pub address_type: AddressType,
48}
49
50impl ConvolutionProblem {
51    pub fn as_matmul_problem(&self) -> MatmulProblem {
52        let rank = self.lhs_strides.len();
53
54        // Strides are expected to be in row major (m, n) format so for matmul checks we need to
55        // convert them to that format, with all other dims treated as batch dims so they're still
56        // checked.
57        // lhs already has the right format, but rhs needs special handling.
58        // (h, w, c, n)
59        let lhs_strides = match self.lhs_layout {
60            MatrixLayout::RowMajor => self.lhs_strides.clone(),
61            MatrixLayout::ColMajor => {
62                let mut lhs_strides: Strides = self.lhs_strides[1..rank].into();
63                lhs_strides.push(self.lhs_strides[0]);
64                lhs_strides
65            }
66        };
67        let rhs_strides = match self.rhs_layout {
68            MatrixLayout::RowMajor => self.rhs_strides.clone(),
69            MatrixLayout::ColMajor => {
70                let mut rhs_strides: Strides = self.rhs_strides[1..rank].into();
71                rhs_strides.push(self.rhs_strides[0]);
72                rhs_strides
73            }
74        };
75
76        MatmulProblem {
77            m: self.m,
78            n: self.n,
79            k: self.k,
80            lhs_batches: shape![],
81            rhs_batches: shape![],
82            out_batches: shape![],
83            lhs_strides,
84            rhs_strides,
85            lhs_layout: self.lhs_layout,
86            rhs_layout: self.rhs_layout,
87            lhs_shape: shape![self.m, self.k],
88            rhs_shape: shape![self.k, self.n],
89            out_shape: shape![self.m, self.n],
90            out_strides: MatrixLayout::RowMajor.to_strides(&[self.m, self.n]),
91            out_layout: MatrixLayout::RowMajor,
92            lhs_scheme: None,
93            rhs_scheme: None,
94            global_dtypes: self.global_dtypes.clone(),
95            address_type: self.address_type,
96        }
97    }
98
99    pub fn should_check_channel(&self) -> bool {
100        self.channels != self.padded_channels
101    }
102
103    pub fn should_check_spatial_bounds(&self) -> bool {
104        self.padding.iter().any(|&pad| pad != 0)
105    }
106}
107
108/// Spatial dimensionality of an operation
109#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug)]
110pub enum Dimensionality {
111    Dim1,
112    Dim2,
113    Dim3,
114}
115
116impl Dimensionality {
117    pub fn num_dims(&self) -> usize {
118        match self {
119            Dimensionality::Dim1 => 1,
120            Dimensionality::Dim2 => 2,
121            Dimensionality::Dim3 => 3,
122        }
123    }
124}