Skip to main content

cubek_convolution/components/global/layout/
im2col.rs

1use 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/// Maps a 4D NHWC tensor to a 2D column matrix using the im2col transformation
17/// It first decomposes the `(m, k)` matrix into `((n, out_h, out_w), (k_h, k_w, c))`, then applies
18/// the convolution parameters to calculate the position in the input tensor for that kernel element.
19#[derive(CubeType, CubeLaunch, Clone)]
20pub struct Im2colLayout {
21    /// Shape of output DHW
22    pub shape_out: Sequence<FastDivmod<u32>>,
23    /// Shape of channel, for decomposing k
24    pub padded_channels: FastDivmod<u32>,
25
26    /// Shape of the combined `m` dimension, including padding
27    pub rows: u32,
28    /// Shape of the combined `k` dimension, including padding
29    pub cols: u32,
30
31    /// Comptime parameters for the convolution
32    #[cube(comptime)]
33    pub params: ConvolutionParams,
34    /// Global memory config for the backing tensor
35    #[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    /// Whether a transposed-convolution coordinate maps to an actual source pixel.
60    ///
61    /// Solving `out * stride - padding + kernel = input` for `out` introduces a
62    /// division by `stride`. Integer division alone would incorrectly map numerators that are not
63    /// divisible by the stride to a neighboring source pixel.
64    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        // Shouldn't be relied on because it doesn't check spatial
158        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}