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 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}