cubek_convolution/components/global/layout/
out.rs1use 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#[derive(CubeType, CubeLaunch, Clone)]
18pub struct OutLayout {
19 pub shape_out: Sequence<FastDivmod<u32>>,
21
22 pub rows: u32,
24 pub cols: u32,
26
27 #[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 OutLayoutLaunch {
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}