Skip to main content

cubecl_cpp/shared/
element.rs

1use cubecl_common::{e2m1, e2m1x2, e3m2, e5m2};
2use cubecl_core::tf32;
3use half::{bf16, f16};
4use std::fmt::Display;
5
6use super::Dialect;
7
8#[derive(Debug, Clone, PartialEq, Eq, Copy, Hash)]
9pub enum Elem<D: Dialect> {
10    TF32,
11    F32,
12    F64,
13    F16,
14    F16x2,
15    BF16,
16    BF16x2,
17    FP4(FP4Kind),
18    FP4x2(FP4Kind),
19    FP6(FP6Kind),
20    FP6x2(FP6Kind),
21    FP8(FP8Kind),
22    FP8x2(FP8Kind),
23    I8,
24    I16,
25    I32,
26    I64,
27    U8,
28    U16,
29    U32,
30    U64,
31    Bool,
32    None,
33    _Dialect(std::marker::PhantomData<D>),
34}
35
36#[derive(Debug, Clone, PartialEq, Eq, Copy, Hash)]
37pub enum FP4Kind {
38    E2M1,
39}
40
41#[derive(Debug, Clone, PartialEq, Eq, Copy, Hash)]
42pub enum FP6Kind {
43    E2M3,
44    E3M2,
45}
46
47#[derive(Debug, Clone, PartialEq, Eq, Copy, Hash)]
48pub enum FP8Kind {
49    E4M3,
50    E5M2,
51    UE8M0,
52}
53
54impl Display for FP4Kind {
55    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
56        let name = match self {
57            FP4Kind::E2M1 => "e2m1",
58        };
59        f.write_str(name)
60    }
61}
62
63impl Display for FP6Kind {
64    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
65        let name = match self {
66            FP6Kind::E2M3 => "e2m3",
67            FP6Kind::E3M2 => "e3m2",
68        };
69        f.write_str(name)
70    }
71}
72
73impl Display for FP8Kind {
74    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
75        let name = match self {
76            FP8Kind::E4M3 => "e4m3",
77            FP8Kind::E5M2 => "e5m2",
78            FP8Kind::UE8M0 => "e8m0",
79        };
80        f.write_str(name)
81    }
82}
83
84impl<D: Dialect> Display for Elem<D> {
85    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86        D::compile_elem(f, self, false)
87    }
88}
89
90impl<D: Dialect> Elem<D> {
91    pub const fn size(&self) -> usize {
92        match self {
93            Elem::FP4(_) => core::mem::size_of::<e2m1>(),
94            Elem::FP4x2(_) => core::mem::size_of::<e2m1x2>(),
95            Elem::FP6(_) => core::mem::size_of::<e3m2>(),
96            Elem::FP6x2(_) => 2 * core::mem::size_of::<e3m2>(),
97            Elem::FP8(_) => core::mem::size_of::<e5m2>(),
98            Elem::FP8x2(_) => 2 * core::mem::size_of::<e5m2>(),
99            Elem::F16 => core::mem::size_of::<f16>(),
100            Elem::F16x2 => 2 * core::mem::size_of::<f16>(),
101            Elem::BF16x2 => 2 * core::mem::size_of::<bf16>(),
102            Elem::BF16 => core::mem::size_of::<bf16>(),
103            Elem::TF32 => core::mem::size_of::<tf32>(),
104            Elem::F32 => core::mem::size_of::<f32>(),
105            Elem::F64 => core::mem::size_of::<f64>(),
106            Elem::I8 => core::mem::size_of::<i8>(),
107            Elem::I16 => core::mem::size_of::<i16>(),
108            Elem::I32 => core::mem::size_of::<i32>(),
109            Elem::I64 => core::mem::size_of::<i64>(),
110            Elem::U8 => core::mem::size_of::<u8>(),
111            Elem::U16 => core::mem::size_of::<u16>(),
112            Elem::U32 => core::mem::size_of::<u32>(),
113            Elem::U64 => core::mem::size_of::<u64>(),
114            Elem::Bool => core::mem::size_of::<bool>(),
115            Elem::None => panic!("Can't get size of `None` element"),
116            Elem::_Dialect(_) => 0,
117        }
118    }
119
120    pub const fn size_bits(&self) -> usize {
121        match self {
122            Elem::FP4(_) => 4,
123            other => other.size() * 8,
124        }
125    }
126
127    pub const fn unpacked(&self) -> Self {
128        match self {
129            Elem::FP4x2(ty) => Elem::FP4(*ty),
130            Elem::FP6x2(ty) => Elem::FP6(*ty),
131            Elem::FP8x2(ty) => Elem::FP8(*ty),
132            Elem::F16x2 => Elem::F16,
133            Elem::BF16x2 => Elem::BF16,
134            elem => *elem,
135        }
136    }
137
138    /// Get the number of values packed into a single storage element. (i.e. `f16x2 -> 2`)
139    pub const fn packing_factor(&self) -> usize {
140        match self {
141            Elem::FP4x2(_) | Elem::FP6x2(_) | Elem::FP8x2(_) | Elem::F16x2 | Elem::BF16x2 => 2,
142            _ => 1,
143        }
144    }
145
146    pub const fn ident(&self) -> &str {
147        match self {
148            Elem::FP4(_) => "fp4",
149            Elem::FP4x2(_) => "fp4x2",
150            Elem::FP6(_) => "fp6",
151            Elem::FP6x2(_) => "fp6x2",
152            Elem::FP8(_) => "fp8",
153            Elem::FP8x2(_) => "fp8x2",
154            Elem::F16 => "f16",
155            Elem::F16x2 => "f16x2",
156            Elem::BF16x2 => "bf16x2",
157            Elem::BF16 => "bf16",
158            Elem::TF32 => "tf32",
159            Elem::F32 => "f32",
160            Elem::F64 => "f64",
161            Elem::I8 => "i8",
162            Elem::I16 => "i16",
163            Elem::I32 => "i32",
164            Elem::I64 => "i64",
165            Elem::U8 => "u8",
166            Elem::U16 => "u16",
167            Elem::U32 => "u32",
168            Elem::U64 => "u64",
169            Elem::Bool => "bool",
170            Elem::None => "<none>",
171            Elem::_Dialect(_) => "",
172        }
173    }
174}