cubek_convolution/components/
config.rs1use 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
12pub 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 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}