cubek_convolution/components/
problem.rs1use 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)]
16pub 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 pub padded_channels: usize,
41 pub operation: ConvolutionOperation,
42
43 pub dimensionality: Dimensionality,
44
45 pub global_dtypes: MatmulGlobalElems,
46 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 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#[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}