1use cubecl::prelude::*;
2
3#[derive(Debug, Clone, Copy)]
4pub enum MatmulDim {
6 M,
8 N,
10 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
99define_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);