cubek_convolution/components/global/layout/
im2col.rs1use cubecl::prelude::*;
2use cubecl::std::{
3 FastDivmod,
4 tensor::layout::{Layout, LayoutExpand},
5};
6use cubek_matmul::multi_level::{
7 args::BatchedCoords,
8 components::global::{GlobalConfig, memory::GlobalLayoutConfig},
9};
10
11use crate::components::{
12 ConvolutionOperation, ConvolutionParams, ConvolutionProblem,
13 global::layout::{NhwcCoords, div_mod_seq},
14};
15
16#[derive(CubeType, CubeLaunch, Clone)]
20pub struct Im2colLayout {
21 pub shape_out: Sequence<FastDivmod<u32>>,
23 pub padded_channels: FastDivmod<u32>,
25
26 pub rows: u32,
28 pub cols: u32,
30
31 #[cube(comptime)]
33 pub params: ConvolutionParams,
34 #[cube(comptime)]
36 pub config: GlobalLayoutConfig,
37}
38
39#[cube]
40impl Im2colLayout {
41 pub fn new<G: GlobalConfig>(
42 rows: u32,
43 cols: u32,
44 padded_channels: FastDivmod<u32>,
45 shape_out: Sequence<FastDivmod<u32>>,
46 #[comptime] config: GlobalLayoutConfig,
47 #[comptime] params: ConvolutionParams,
48 ) -> Im2colLayout {
49 Im2colLayout {
50 shape_out,
51 padded_channels,
52 rows,
53 cols,
54 params,
55 config,
56 }
57 }
58
59 fn stride_is_valid(&self, pos: BatchedCoords) -> bool {
65 let params = self.params.comptime();
66
67 match params.operation {
68 ConvolutionOperation::Forward | ConvolutionOperation::BackwardWeight => true.runtime(),
69 ConvolutionOperation::ForwardTransposed | ConvolutionOperation::BackwardData => {
70 if params.has_non_unit_stride() {
71 let (_, view_m, view_k) = pos;
72 let (_, out_offs) = div_mod_seq(view_m, &self.shape_out);
73 let (mut rem, _) = self.padded_channels.div_mod(view_k);
74
75 let spatial_dims = params.dimensionality.num_dims();
76 let mut valid = true.runtime();
77
78 #[unroll]
79 for i in 0..spatial_dims {
80 let dim = spatial_dims - i - 1;
81 let ksize = params.kernel_size[dim];
82 let k_pos = (rem % ksize) as i32;
83 rem /= ksize;
84
85 let numerator = out_offs[dim] as i32 + params.padding[dim]
86 - k_pos * params.dilation[dim] as i32;
87 valid &= numerator % params.stride[dim] as i32 == 0;
88 }
89
90 valid
91 } else {
92 true.runtime()
93 }
94 }
95 }
96 }
97}
98
99#[cube]
100impl Layout for Im2colLayout {
101 type Coordinates = BatchedCoords;
102 type SourceCoordinates = NhwcCoords;
103
104 fn to_source_pos(&self, pos: Self::Coordinates) -> NhwcCoords {
105 let params = self.params.comptime();
106 let (_, view_m, view_k) = pos;
107
108 let (batch, out_offs) = div_mod_seq(view_m, &self.shape_out);
109
110 let (mut rem, channel) = self.padded_channels.div_mod(view_k);
111
112 let spatial_dims = params.dimensionality.num_dims();
113 let mut in_pos = Sequence::<i32>::new();
114
115 #[unroll]
116 for i in 0..spatial_dims {
117 let dim = spatial_dims - i - 1;
118 let ksize = params.kernel_size[dim];
119 let k_pos = (rem % ksize) as i32;
120 rem /= ksize;
121
122 let out_pos = out_offs[dim];
123 let stride = params.stride[dim] as i32;
124 let dilate = params.dilation[dim] as i32;
125 let pad = params.padding[dim];
126
127 let pos = match params.operation {
128 ConvolutionOperation::Forward | ConvolutionOperation::BackwardWeight => {
129 (out_pos as i32 * stride + k_pos * dilate) - pad
130 }
131 ConvolutionOperation::ForwardTransposed | ConvolutionOperation::BackwardData => {
132 (out_pos as i32 + pad - k_pos * dilate) / stride
133 }
134 };
135 in_pos.push(pos);
136 }
137
138 let in_pos = in_pos.reversed();
139
140 NhwcCoords {
141 batch,
142 spatial: in_pos,
143 channel,
144 }
145 }
146
147 fn shape(&self) -> Self::Coordinates {
148 (1, self.rows, self.cols)
149 }
150
151 fn to_source_pos_checked(&self, pos: Self::Coordinates) -> (NhwcCoords, bool) {
152 (self.to_source_pos(pos), self.is_in_bounds(pos))
153 }
154
155 fn is_in_bounds(&self, pos: Self::Coordinates) -> bool {
156 let (_, view_m, view_k) = pos;
157 let m_in_bounds = !self.config.check_row_bounds || view_m < self.rows;
159 let k_in_bounds = !self.config.check_col_bounds || view_k < self.cols;
160 m_in_bounds && k_in_bounds && self.stride_is_valid(pos)
161 }
162}
163
164impl Im2colLayoutLaunch {
165 pub fn from_args(
166 problem: &ConvolutionProblem,
167 params: ConvolutionParams,
168 config: GlobalLayoutConfig,
169 ) -> Self {
170 match problem.operation {
171 ConvolutionOperation::Forward => Self::from_args_fprop(problem, params, config),
172 ConvolutionOperation::ForwardTransposed | ConvolutionOperation::BackwardData => {
173 Self::from_args_dgrad(problem, params, config)
174 }
175 ConvolutionOperation::BackwardWeight => Self::from_args_wgrad(problem, params, config),
176 }
177 }
178
179 fn from_args_fprop(
180 problem: &ConvolutionProblem,
181 params: ConvolutionParams,
182 config: GlobalLayoutConfig,
183 ) -> Self {
184 let shape_out = problem.out_shape.iter().map(|s| *s as u32).collect();
185
186 let padded_channels = problem.padded_channels as u32;
187
188 let shape_m = problem.m as u32;
189 let shape_k = problem.k as u32;
190
191 Im2colLayoutLaunch::new(shape_out, padded_channels, shape_m, shape_k, params, config)
192 }
193
194 fn from_args_dgrad(
195 problem: &ConvolutionProblem,
196 params: ConvolutionParams,
197 config: GlobalLayoutConfig,
198 ) -> Self {
199 let shape = problem.in_shape.iter().map(|s| *s as u32).collect();
200
201 let padded_channels = problem.padded_channels as u32;
202
203 let shape_m = problem.m as u32;
204 let shape_k = problem.k as u32;
205
206 Im2colLayoutLaunch::new(shape, padded_channels, shape_m, shape_k, params, config)
207 }
208
209 fn from_args_wgrad(
210 problem: &ConvolutionProblem,
211 params: ConvolutionParams,
212 config: GlobalLayoutConfig,
213 ) -> Self {
214 let shape_out = problem.out_shape.iter().map(|s| *s as u32).collect();
215
216 let padded_channels = problem.padded_channels as u32;
217
218 let shape_k = problem.k as u32;
219 let shape_n = problem.n as u32;
220
221 Im2colLayoutLaunch::new(shape_out, padded_channels, shape_k, shape_n, params, config)
222 }
223}