1#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
8#[repr(i32)]
9pub enum DType {
10 #[default]
11 Unspecified = 0,
12 Bool = 1,
13 Uint8 = 2,
14 Uint16 = 3,
15 Uint32 = 4,
16 Uint64 = 5,
17 Int8 = 6,
18 Int16 = 7,
19 Int32 = 8,
20 Int64 = 9,
21 Float16 = 10,
22 Float32 = 11,
23 Float64 = 12,
24}
25
26impl TryFrom<i32> for DType {
27 type Error = &'static str;
28
29 fn try_from(value: i32) -> Result<Self, Self::Error> {
30 match value {
31 0 => Ok(Self::Unspecified),
32 1 => Ok(Self::Bool),
33 2 => Ok(Self::Uint8),
34 3 => Ok(Self::Uint16),
35 4 => Ok(Self::Uint32),
36 5 => Ok(Self::Uint64),
37 6 => Ok(Self::Int8),
38 7 => Ok(Self::Int16),
39 8 => Ok(Self::Int32),
40 9 => Ok(Self::Int64),
41 10 => Ok(Self::Float16),
42 11 => Ok(Self::Float32),
43 12 => Ok(Self::Float64),
44 _ => Err("invalid dtype"),
45 }
46 }
47}
48
49impl From<DType> for i32 {
50 fn from(value: DType) -> Self {
51 value as i32
52 }
53}
54
55impl DType {
56 pub const ALL: [DType; 13] = [
58 DType::Unspecified,
59 DType::Bool,
60 DType::Uint8,
61 DType::Uint16,
62 DType::Uint32,
63 DType::Uint64,
64 DType::Int8,
65 DType::Int16,
66 DType::Int32,
67 DType::Int64,
68 DType::Float16,
69 DType::Float32,
70 DType::Float64,
71 ];
72
73 pub const fn name(self) -> &'static str {
75 match self {
76 DType::Unspecified => "unspecified",
77 DType::Bool => "bool",
78 DType::Uint8 => "uint8",
79 DType::Int32 => "int32",
80 DType::Int64 => "int64",
81 DType::Float16 => "float16",
82 DType::Float32 => "float32",
83 DType::Float64 => "float64",
84 DType::Int8 => "int8",
85 DType::Int16 => "int16",
86 DType::Uint16 => "uint16",
87 DType::Uint32 => "uint32",
88 DType::Uint64 => "uint64",
89 }
90 }
91
92 pub fn from_name(name: &str) -> Option<Self> {
95 match name {
96 "bool" => Some(DType::Bool),
97 "uint8" => Some(DType::Uint8),
98 "int32" => Some(DType::Int32),
99 "int64" => Some(DType::Int64),
100 "float16" => Some(DType::Float16),
101 "float32" => Some(DType::Float32),
102 "float64" => Some(DType::Float64),
103 "int8" => Some(DType::Int8),
104 "int16" => Some(DType::Int16),
105 "uint16" => Some(DType::Uint16),
106 "uint32" => Some(DType::Uint32),
107 "uint64" => Some(DType::Uint64),
108 _ => None,
109 }
110 }
111
112 #[must_use]
115 pub const fn is_integer(self) -> bool {
116 matches!(
117 self,
118 DType::Uint8
119 | DType::Int8
120 | DType::Int16
121 | DType::Uint16
122 | DType::Int32
123 | DType::Uint32
124 | DType::Int64
125 | DType::Uint64
126 )
127 }
128
129 #[must_use]
131 pub const fn is_float(self) -> bool {
132 matches!(self, DType::Float16 | DType::Float32 | DType::Float64)
133 }
134}
135
136impl std::fmt::Display for DType {
137 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
138 f.write_str(self.name())
139 }
140}
141
142pub const fn dtype_size(dtype: DType) -> usize {
144 match dtype {
145 DType::Unspecified => 0,
146 DType::Bool | DType::Uint8 | DType::Int8 => 1,
147 DType::Float16 | DType::Int16 | DType::Uint16 => 2,
148 DType::Int32 | DType::Uint32 | DType::Float32 => 4,
149 DType::Int64 | DType::Uint64 | DType::Float64 => 8,
150 }
151}
152
153#[cfg(test)]
154mod tests {
155 use super::*;
156
157 #[test]
158 fn test_dtype_i32_roundtrip() {
159 for dtype in DType::ALL {
160 let raw = i32::from(dtype);
161 assert_eq!(DType::try_from(raw), Ok(dtype));
162 }
163 }
164
165 #[test]
166 fn test_dtype_rejects_unknown_values() {
167 assert!(DType::try_from(-1).is_err());
168 for unused in [13, 20, 30, 99] {
169 assert!(
170 DType::try_from(unused).is_err(),
171 "{unused} should be invalid"
172 );
173 }
174 }
175
176 #[test]
177 fn test_dtype_name_roundtrip() {
178 for dtype in DType::ALL {
179 if dtype == DType::Unspecified {
180 continue;
181 }
182 assert_eq!(DType::from_name(dtype.name()), Some(dtype));
183 assert_eq!(dtype.to_string(), dtype.name());
184 }
185 assert_eq!(DType::Unspecified.name(), "unspecified");
186 assert_eq!(DType::from_name("unspecified"), None);
187 assert_eq!(DType::from_name("complex64"), None);
188 assert_eq!(DType::from_name("Float32"), None);
189 }
190
191 #[test]
192 fn test_dtype_size_table() {
193 let expected = [
194 (DType::Unspecified, 0),
195 (DType::Bool, 1),
196 (DType::Uint8, 1),
197 (DType::Int8, 1),
198 (DType::Float16, 2),
199 (DType::Int16, 2),
200 (DType::Uint16, 2),
201 (DType::Int32, 4),
202 (DType::Uint32, 4),
203 (DType::Float32, 4),
204 (DType::Int64, 8),
205 (DType::Uint64, 8),
206 (DType::Float64, 8),
207 ];
208 for (dtype, size) in expected {
209 assert_eq!(dtype_size(dtype), size, "size mismatch for {dtype:?}");
210 }
211 }
212}