Skip to main content

cubek_convolution/components/global/layout/
out.rs

1use cubecl::prelude::*;
2use cubecl::std::{
3    FastDivmod,
4    tensor::layout::{Layout, LayoutExpand},
5};
6use cubek_matmul::multi_level::{
7    args::BatchedCoords, components::global::memory::GlobalLayoutConfig,
8};
9
10use crate::components::{
11    ConvolutionOperation, ConvolutionProblem,
12    global::layout::{NhwcCoords, cast_seq, div_mod_seq},
13};
14
15/// Maps a 4D NHWC out tensor of shape `((n, h, w), c)` to a col-major 2D matmul tile with
16/// shape `(m, n)`
17#[derive(CubeType, CubeLaunch, Clone)]
18pub struct OutLayout {
19    /// Shape of DHW
20    pub shape_out: Sequence<FastDivmod<u32>>,
21
22    /// Shape of the conceptual `m` size
23    pub rows: u32,
24    /// Shape of the conceptual `n`size, or channels
25    pub cols: u32,
26
27    /// Global memory config for the backing tensor
28    #[cube(comptime)]
29    pub config: GlobalLayoutConfig,
30}
31
32#[cube]
33impl OutLayout {
34    pub fn new(
35        rows: u32,
36        cols: u32,
37        shape_out: Sequence<FastDivmod<u32>>,
38        #[comptime] config: GlobalLayoutConfig,
39    ) -> OutLayout {
40        OutLayout {
41            shape_out,
42            rows,
43            cols,
44            config,
45        }
46    }
47}
48
49#[cube]
50impl Layout for OutLayout {
51    type Coordinates = BatchedCoords;
52    type SourceCoordinates = NhwcCoords;
53
54    fn to_source_pos(&self, coords: Self::Coordinates) -> NhwcCoords {
55        let (_, view_m, view_n) = coords;
56        let (batch, spatial) = div_mod_seq(view_m, &self.shape_out);
57
58        NhwcCoords {
59            batch,
60            spatial: cast_seq(spatial),
61            channel: view_n,
62        }
63    }
64
65    fn to_source_pos_checked(&self, coords: Self::Coordinates) -> (NhwcCoords, bool) {
66        (self.to_source_pos(coords), self.is_in_bounds(coords))
67    }
68
69    fn shape(&self) -> Self::Coordinates {
70        (1, self.rows, self.cols)
71    }
72
73    fn is_in_bounds(&self, pos: Self::Coordinates) -> bool {
74        let (_, row, col) = pos;
75        (!self.config.check_row_bounds || row < self.rows)
76            && (!self.config.check_col_bounds || col < self.cols)
77    }
78}
79
80impl<R: Runtime> OutLayoutLaunch<R> {
81    pub fn from_args(problem: &ConvolutionProblem, config: GlobalLayoutConfig) -> Self {
82        match problem.operation {
83            ConvolutionOperation::Forward => Self::from_args_fprop(problem, config),
84            ConvolutionOperation::ForwardTransposed | ConvolutionOperation::BackwardData => {
85                Self::from_args_dgrad(problem, config)
86            }
87            ConvolutionOperation::BackwardWeight => Self::from_args_wgrad(problem, config),
88        }
89    }
90
91    fn from_args_fprop(problem: &ConvolutionProblem, config: GlobalLayoutConfig) -> Self {
92        let shape_out = problem.out_shape.iter().map(|s| *s as u32).collect();
93        let shape_m = problem.m as u32;
94        let shape_n = problem.n as u32;
95
96        Self::new(shape_out, shape_m, shape_n, config)
97    }
98
99    fn from_args_dgrad(problem: &ConvolutionProblem, config: GlobalLayoutConfig) -> Self {
100        let shape = problem.in_shape.iter().map(|s| *s as u32).collect();
101        let shape_m = problem.m as u32;
102        let shape_n = problem.n as u32;
103
104        Self::new(shape, shape_m, shape_n, config)
105    }
106
107    fn from_args_wgrad(problem: &ConvolutionProblem, config: GlobalLayoutConfig) -> Self {
108        let shape_out = problem.out_shape.iter().map(|s| *s as u32).collect();
109        let shape_m = problem.m as u32;
110        let shape_k = problem.k as u32;
111
112        Self::new(shape_out, shape_k, shape_m, config)
113    }
114}