Skip to main content

cubek_convolution/components/
config.rs

1use std::ops::Deref;
2
3use cubecl::CubeDim;
4use cubek_matmul::{
5    definition::{MatmulSetupError, MatmulVectorSizes},
6    multi_level::components::global::{GlobalConfig, memory::GlobalMemoryConfig},
7};
8use std::{fmt::Debug, hash::Hash};
9
10use super::*;
11
12/// Convolution specific config, extends regular matmul `Config`.
13pub trait ConvGemmConfig:
14    Deref<Target: GlobalConfig> + Copy + Clone + Eq + PartialEq + Hash + Debug + Send + Sync + 'static
15{
16    type GlobalMatmulConfig: GlobalConfig;
17
18    fn matmul_config(&self) -> Self::GlobalMatmulConfig;
19
20    /// The size of the convolution kernel at `dim`
21    fn params(&self) -> ConvolutionParams;
22    fn operation(&self) -> ConvolutionOperation;
23    fn vector_sizes(&self) -> MatmulVectorSizes;
24    fn check_spatial_bounds(&self) -> bool;
25    fn cube_dim(&self) -> CubeDim;
26    fn lhs_global_memory_config(&self) -> GlobalMemoryConfig;
27    fn rhs_global_memory_config(&self) -> GlobalMemoryConfig;
28    fn out_global_memory_config(&self) -> GlobalMemoryConfig;
29}
30
31#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
32pub struct ConvolutionConfig<M: GlobalConfig> {
33    pub matmul: M,
34    pub params: ConvolutionParams,
35}
36
37#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
38pub struct ConvolutionParams {
39    pub kernel_size: [u32; 3],
40    pub stride: [u32; 3],
41    pub dilation: [u32; 3],
42    pub padding: [i32; 3],
43    pub dimensionality: Dimensionality,
44    pub operation: ConvolutionOperation,
45}
46
47impl ConvolutionParams {
48    pub fn from_problem(problem: &ConvolutionProblem) -> Self {
49        let dims = problem.dimensionality.num_dims();
50
51        let mut params = ConvolutionParams {
52            kernel_size: [0; 3],
53            stride: [0; 3],
54            dilation: [0; 3],
55            padding: [0; 3],
56            dimensionality: problem.dimensionality,
57            operation: problem.operation,
58        };
59        params.kernel_size[0..dims].copy_from_slice(&problem.kernel_size);
60        params.stride[0..dims].copy_from_slice(&problem.stride);
61        params.dilation[0..dims].copy_from_slice(&problem.dilation);
62        params.padding[0..dims].copy_from_slice(&problem.padding);
63        params
64    }
65
66    pub(crate) fn has_non_unit_stride(&self) -> bool {
67        let dims = self.dimensionality.num_dims();
68        self.stride[..dims].iter().any(|&stride| stride != 1)
69    }
70}
71
72impl<M: GlobalConfig> Deref for ConvolutionConfig<M> {
73    type Target = M;
74
75    fn deref(&self) -> &Self::Target {
76        &self.matmul
77    }
78}
79
80impl<M: GlobalConfig> ConvGemmConfig for ConvolutionConfig<M> {
81    type GlobalMatmulConfig = M;
82
83    fn matmul_config(&self) -> Self::GlobalMatmulConfig {
84        self.matmul
85    }
86
87    fn vector_sizes(&self) -> MatmulVectorSizes {
88        self.matmul.global_vector_sizes()
89    }
90
91    fn cube_dim(&self) -> CubeDim {
92        self.matmul.cube_dim()
93    }
94
95    fn check_spatial_bounds(&self) -> bool {
96        let spatial_dims = self.params.dimensionality.num_dims();
97        let mut has_padding = false;
98        for i in 0..spatial_dims {
99            has_padding |= self.params.padding[i] != 0;
100        }
101        has_padding
102    }
103
104    fn params(&self) -> ConvolutionParams {
105        self.params
106    }
107
108    fn operation(&self) -> ConvolutionOperation {
109        self.params.operation
110    }
111
112    fn lhs_global_memory_config(&self) -> GlobalMemoryConfig {
113        self.matmul.lhs_reader_config().gmem_config
114    }
115
116    fn rhs_global_memory_config(&self) -> GlobalMemoryConfig {
117        self.matmul.rhs_reader_config().gmem_config
118    }
119
120    fn out_global_memory_config(&self) -> GlobalMemoryConfig {
121        self.matmul.writer_config().gmem_config
122    }
123}
124
125impl<M: GlobalConfig> ConvolutionConfig<M> {
126    #[allow(clippy::too_many_arguments)]
127    pub fn new(
128        matmul: M,
129        kernel_size: &[u32],
130        stride: &[u32],
131        dilation: &[u32],
132        padding: &[i32],
133        dim: Dimensionality,
134        operation: ConvolutionOperation,
135    ) -> Result<Self, MatmulSetupError> {
136        let dims = kernel_size.len();
137
138        let mut params = ConvolutionParams {
139            kernel_size: [0; 3],
140            stride: [0; 3],
141            dilation: [0; 3],
142            padding: [0; 3],
143            dimensionality: dim,
144            operation,
145        };
146        params.kernel_size[0..dims].copy_from_slice(kernel_size);
147        params.stride[0..dims].copy_from_slice(stride);
148        params.dilation[0..dims].copy_from_slice(dilation);
149        params.padding[0..dims].copy_from_slice(padding);
150        Ok(Self { matmul, params })
151    }
152}