1use super::DataType;
4
5#[allow(clippy::match_same_arms)]
6impl DataType {
7 #[must_use]
9 pub const fn min_bytes(&self) -> usize {
10 match self {
11 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => 2,
12 Self::Bool | Self::U32 | Self::I32 | Self::F32 | Self::Handle(_) => 4,
13 Self::I64 | Self::U64 | Self::Vec2U32 | Self::F64 => 8,
14 Self::Vec4U32 => 16,
15 Self::Vec { element, count } => element.min_bytes().saturating_mul(*count as usize),
16 Self::Bytes | Self::Array { .. } | Self::Tensor | Self::TensorShaped { .. } => 0,
17 Self::U8
21 | Self::I8
22 | Self::F8E4M3
23 | Self::F8E5M2
24 | Self::I4
25 | Self::FP4
26 | Self::NF4 => 1,
27 Self::SparseCsr { .. } | Self::SparseCoo { .. } | Self::SparseBsr { .. } => 0,
30 Self::DeviceMesh { .. } => 0,
31 Self::Quantized { storage, .. } => storage.min_bytes(),
32 Self::Opaque(_) => 0,
34 }
35 }
36
37 #[must_use]
42 pub const fn max_bytes(&self) -> Option<usize> {
43 match self {
44 Self::U8 | Self::I8 => Some(1),
45 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => Some(2),
46 Self::U32 | Self::I32 | Self::Bool => Some(4),
47 Self::I64 | Self::U64 | Self::Vec2U32 | Self::F64 => Some(8),
48 Self::Vec4U32 => Some(16),
49 Self::F32 => Some(4),
50 Self::Handle(_) => Some(4),
51 Self::Vec { element, count } => match element.max_bytes() {
52 Some(bytes) => bytes.checked_mul(*count as usize),
53 None => None,
54 },
55 Self::Bytes => Some(64 * 1024 * 1024),
56 Self::Array { .. } | Self::Tensor => Some(256 * 1024 * 1024),
57 Self::TensorShaped { .. } => None,
58 Self::F8E4M3 | Self::F8E5M2 => Some(1),
59 Self::I4 | Self::FP4 | Self::NF4 => Some(1),
60 Self::SparseCsr { .. } | Self::SparseCoo { .. } | Self::SparseBsr { .. } => None,
61 Self::DeviceMesh { .. } => Some(4),
62 Self::Quantized { storage, .. } => storage.max_bytes(),
63 Self::Opaque(_) => None,
65 }
66 }
67
68 #[must_use]
70 pub const fn element_size(&self) -> Option<usize> {
71 match self {
72 Self::Array { element_size } => Some(*element_size),
73 Self::Vec { element, .. }
74 | Self::TensorShaped { element, .. }
75 | Self::SparseCsr { element }
76 | Self::SparseCoo { element }
77 | Self::SparseBsr { element, .. } => element.size_bytes(),
78 Self::Quantized { storage, .. } => storage.size_bytes(),
79 Self::Opaque(_) => None,
80 _ => None,
81 }
82 }
83
84 #[must_use]
91 pub const fn size_bytes(&self) -> Option<usize> {
92 match self {
93 Self::U8 | Self::I8 => Some(1),
94 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => Some(2),
95 Self::Bool | Self::U32 | Self::I32 | Self::F32 => Some(4),
96 Self::I64 | Self::U64 | Self::Vec2U32 | Self::F64 => Some(8),
97 Self::Vec4U32 => Some(16),
98 Self::Handle(_) => Some(4),
99 Self::Bytes => Some(1),
100 Self::Array { element_size } => Some(*element_size),
101 Self::Vec { element, count } => match element.size_bytes() {
102 Some(bytes) => bytes.checked_mul(*count as usize),
103 None => None,
104 },
105 Self::Tensor | Self::TensorShaped { .. } => None,
106 Self::F8E4M3 | Self::F8E5M2 => Some(1),
107 Self::I4 | Self::FP4 | Self::NF4 => Some(1),
108 Self::SparseCsr { .. } | Self::SparseCoo { .. } | Self::SparseBsr { .. } => None,
109 Self::DeviceMesh { .. } => Some(4),
110 Self::Quantized { storage, .. } => storage.size_bytes(),
111 Self::Opaque(_) => None,
113 }
114 }
115
116 #[must_use]
130 pub const fn bit_width(&self) -> Option<usize> {
131 match self {
132 Self::I4 | Self::FP4 | Self::NF4 => Some(4),
133 Self::F8E4M3 | Self::F8E5M2 | Self::U8 | Self::I8 => Some(8),
134 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => Some(16),
135 Self::Bool | Self::U32 | Self::I32 | Self::F32 | Self::Handle(_) => Some(32),
136 Self::I64 | Self::U64 | Self::F64 | Self::Vec2U32 => Some(64),
137 Self::Vec4U32 => Some(128),
138 Self::DeviceMesh { .. } => Some(32),
139 Self::Quantized { storage, .. } => storage.bit_width(),
140 Self::Bytes => Some(8),
141 Self::Vec { element, count } => match element.bit_width() {
144 Some(bits) => bits.checked_mul(*count as usize),
145 None => None,
146 },
147 Self::Array { .. }
150 | Self::Tensor
151 | Self::TensorShaped { .. }
152 | Self::SparseCsr { .. }
153 | Self::SparseCoo { .. }
154 | Self::SparseBsr { .. }
155 | Self::Opaque(_) => None,
156 }
157 }
158
159 pub fn packed_size_bytes(&self, element_count: usize) -> Result<Option<usize>, String> {
172 if let Some(bits) = self.checked_bit_width_for_packed_size()? {
173 let total_bits = bits.checked_mul(element_count).ok_or_else(|| {
174 format!(
175 "Fix: packed byte sizing overflowed bits for {self} with {element_count} logical element(s)."
176 )
177 })?;
178 return total_bits
179 .checked_add(7)
180 .map(|rounded_bits| Some(rounded_bits / 8))
181 .ok_or_else(|| {
182 format!(
183 "Fix: packed byte sizing overflowed byte rounding for {self} with {element_count} logical element(s)."
184 )
185 });
186 }
187 if let Some(bytes) = self.checked_size_bytes_for_packed_size()? {
188 return bytes
189 .checked_mul(element_count)
190 .map(Some)
191 .ok_or_else(|| {
192 format!(
193 "Fix: packed byte sizing overflowed bytes for {self} with {element_count} logical element(s)."
194 )
195 });
196 }
197 Ok(None)
198 }
199
200 fn checked_bit_width_for_packed_size(&self) -> Result<Option<usize>, String> {
201 match self {
202 Self::I4 | Self::FP4 | Self::NF4 => Ok(Some(4)),
203 Self::F8E4M3 | Self::F8E5M2 | Self::U8 | Self::I8 => Ok(Some(8)),
204 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => Ok(Some(16)),
205 Self::Bool | Self::U32 | Self::I32 | Self::F32 | Self::Handle(_) => Ok(Some(32)),
206 Self::I64 | Self::U64 | Self::F64 | Self::Vec2U32 => Ok(Some(64)),
207 Self::Vec4U32 => Ok(Some(128)),
208 Self::DeviceMesh { .. } => Ok(Some(32)),
209 Self::Quantized { storage, .. } => storage.checked_bit_width_for_packed_size(),
210 Self::Bytes => Ok(Some(8)),
211 Self::Vec { element, count } => {
212 let Some(bits) = element.checked_bit_width_for_packed_size()? else {
213 return Ok(None);
214 };
215 bits.checked_mul(*count as usize).map(Some).ok_or_else(|| {
216 format!("Fix: packed byte sizing overflowed nested bit width for {self}.")
217 })
218 }
219 Self::Array { .. }
220 | Self::Tensor
221 | Self::TensorShaped { .. }
222 | Self::SparseCsr { .. }
223 | Self::SparseCoo { .. }
224 | Self::SparseBsr { .. }
225 | Self::Opaque(_) => Ok(None),
226 }
227 }
228
229 fn checked_size_bytes_for_packed_size(&self) -> Result<Option<usize>, String> {
230 match self {
231 Self::U8 | Self::I8 => Ok(Some(1)),
232 Self::U16 | Self::I16 | Self::F16 | Self::BF16 => Ok(Some(2)),
233 Self::Bool | Self::U32 | Self::I32 | Self::F32 => Ok(Some(4)),
234 Self::I64 | Self::U64 | Self::Vec2U32 | Self::F64 => Ok(Some(8)),
235 Self::Vec4U32 => Ok(Some(16)),
236 Self::Handle(_) => Ok(Some(4)),
237 Self::Bytes => Ok(Some(1)),
238 Self::Array { element_size } => Ok(Some(*element_size)),
239 Self::Vec { element, count } => {
240 let Some(bytes) = element.checked_size_bytes_for_packed_size()? else {
241 return Ok(None);
242 };
243 bytes.checked_mul(*count as usize).map(Some).ok_or_else(|| {
244 format!("Fix: packed byte sizing overflowed nested byte width for {self}.")
245 })
246 }
247 Self::Tensor | Self::TensorShaped { .. } => Ok(None),
248 Self::F8E4M3 | Self::F8E5M2 => Ok(Some(1)),
249 Self::I4 | Self::FP4 | Self::NF4 => Ok(Some(1)),
250 Self::SparseCsr { .. } | Self::SparseCoo { .. } | Self::SparseBsr { .. } => Ok(None),
251 Self::DeviceMesh { .. } => Ok(Some(4)),
252 Self::Quantized { storage, .. } => storage.checked_size_bytes_for_packed_size(),
253 Self::Opaque(_) => Ok(None),
254 }
255 }
256}