Skip to main content

cubek_std/
size.rs

1use cubecl::prelude::*;
2
3#[derive(Debug, Clone, Copy)]
4/// Matrix dimension specifier for matmul operations.
5pub enum MatmulDim {
6    /// Rows of the output matrix.
7    M,
8    /// Columns of the output matrix.
9    N,
10    /// Reduction dimension.
11    K,
12}
13
14#[macro_export]
15macro_rules! define_3d_size_base {
16    ($name:ident, $ty:ty) => {
17        #[derive(CubeType, Copy, Clone, Debug, Hash, PartialEq, Eq)]
18        pub struct $name {
19            pub m: $ty,
20            pub n: $ty,
21            pub k: $ty,
22        }
23
24        impl $name {
25            pub fn new(m: u32, n: u32, k: u32) -> Self {
26                $name {
27                    m: <$ty>::try_from(m).unwrap(),
28                    n: <$ty>::try_from(n).unwrap(),
29                    k: <$ty>::try_from(k).unwrap(),
30                }
31            }
32
33            pub fn get(&self, dim: $crate::MatmulDim) -> u32 {
34                (match dim {
35                    $crate::MatmulDim::M => self.m,
36                    $crate::MatmulDim::N => self.n,
37                    $crate::MatmulDim::K => self.k,
38                }) as u32
39            }
40
41            pub fn m(&self) -> u32 {
42                self.get($crate::MatmulDim::M)
43            }
44
45            pub fn n(&self) -> u32 {
46                self.get($crate::MatmulDim::N)
47            }
48
49            pub fn k(&self) -> u32 {
50                self.get($crate::MatmulDim::K)
51            }
52
53            pub fn mn(&self) -> u32 {
54                self.get($crate::MatmulDim::M) * self.get($crate::MatmulDim::N)
55            }
56
57            pub fn mk(&self) -> u32 {
58                self.get($crate::MatmulDim::M) * self.get($crate::MatmulDim::K)
59            }
60
61            pub fn nk(&self) -> u32 {
62                self.get($crate::MatmulDim::N) * self.get($crate::MatmulDim::K)
63            }
64
65            pub fn mnk(&self) -> u32 {
66                self.get($crate::MatmulDim::M)
67                    * self.get($crate::MatmulDim::N)
68                    * self.get($crate::MatmulDim::K)
69            }
70        }
71    };
72}
73
74#[macro_export]
75macro_rules! impl_3d_size_from_tuple {
76    ($name:ident, $ty_struct:ty, $ty_tuple:ty) => {
77        impl From<($ty_tuple, $ty_tuple, $ty_tuple)> for $name {
78            fn from(value: ($ty_tuple, $ty_tuple, $ty_tuple)) -> Self {
79                Self {
80                    m: value.0 as $ty_struct,
81                    n: value.1 as $ty_struct,
82                    k: value.2 as $ty_struct,
83                }
84            }
85        }
86
87        impl From<$name> for ($ty_tuple, $ty_tuple, $ty_tuple) {
88            fn from(value: $name) -> Self {
89                (
90                    value.m as $ty_tuple,
91                    value.n as $ty_tuple,
92                    value.k as $ty_tuple,
93                )
94            }
95        }
96    };
97}
98
99// Shapes m,n,k of the problem
100define_3d_size_base!(MatmulProblemSize, u32);
101impl_3d_size_from_tuple!(MatmulProblemSize, u32, u8);
102impl_3d_size_from_tuple!(MatmulProblemSize, u32, u32);
103impl_3d_size_from_tuple!(MatmulProblemSize, u32, i32);
104impl_3d_size_from_tuple!(MatmulProblemSize, u32, u16);
105impl_3d_size_from_tuple!(MatmulProblemSize, u32, usize);