Skip to main content

vyre_spec/data_type/
layout.rs

1//! Size and packing rules for frozen data-type contracts.
2
3use super::DataType;
4
5#[allow(clippy::match_same_arms)]
6impl DataType {
7    /// Minimum byte count to represent one value of this type.
8    #[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            // Quantized / compressed scalar families. F8/F4 = 1 byte rounded up;
18            // I4 / NF4 = 1 byte rounded up (two values share a byte in practice,
19            // but the conservative minimum is one byte per logical value).
20            Self::U8
21            | Self::I8
22            | Self::F8E4M3
23            | Self::F8E5M2
24            | Self::I4
25            | Self::FP4
26            | Self::NF4 => 1,
27            // Sparse layouts + device-mesh handles are unbounded at the
28            // spec level; runtime asks the extension for a concrete size.
29            Self::SparseCsr { .. } | Self::SparseCoo { .. } | Self::SparseBsr { .. } => 0,
30            Self::DeviceMesh { .. } => 0,
31            Self::Quantized { storage, .. } => storage.min_bytes(),
32            // Opaque: conservative sentinel. Real value via ExtensionDataType::min_bytes.
33            Self::Opaque(_) => 0,
34        }
35    }
36
37    /// Maximum byte count for one value of this type.
38    ///
39    /// Returns `None` for truly unbounded types; currently all variants
40    /// have a hard ceiling. Fixed-width types return `Some(min_bytes())`.
41    #[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            // Opaque: unbounded at the spec level. Real ceiling via ExtensionDataType::max_bytes.
64            Self::Opaque(_) => None,
65        }
66    }
67
68    /// Element size for array-typed outputs, or `None` for scalar types.
69    #[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    /// Fixed scalar element size in bytes, or `None` for variable-size types.
85    ///
86    /// Scalar types return their natural width (`U32` -> `Some(4)`, `Vec4U32` ->
87    /// `Some(16)`). `Bytes` returns `Some(1)` because each element is one byte.
88    /// `Array` returns `Some(element_size)`. `Tensor` returns `None` because it
89    /// has no fixed per-element size.
90    #[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            // Opaque: real size via ExtensionDataType::size_bytes (runtime).
112            Self::Opaque(_) => None,
113        }
114    }
115
116    /// True element bit width for fixed-width scalar types. Returns `None`
117    /// for variable / dynamically-shaped / extension-defined types.
118    ///
119    /// Sub-byte types (`I4`, `FP4`, `NF4`) report `4` here; `size_bytes`
120    /// over-rounds to `1` for safety. Callers that pack two `I4` per
121    /// byte (the standard layout for INT4 quantization) need
122    /// `bit_width()` to compute correct packed-buffer sizes:
123    ///
124    /// ```ignore
125    /// // Allocate enough bytes to hold `count` packed `I4` values.
126    /// let bits = count.checked_mul(DataType::I4.bit_width().unwrap_or(8)).unwrap();
127    /// let bytes = bits.div_ceil(8);
128    /// ```
129    #[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            // Vec packs `count` elements: total bits scale with the inner
142            // element's bit width.
143            Self::Vec { element, count } => match element.bit_width() {
144                Some(bits) => bits.checked_mul(*count as usize),
145                None => None,
146            },
147            // Variable / extension-defined / dynamically-shaped: no
148            // compile-time width.
149            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    /// Checked packed byte count for `element_count` logical values.
160    ///
161    /// This is the sizing helper CUDA and wire codecs should use for sub-byte
162    /// quantized storage. `I4`, `FP4`, and `NF4` pack two logical elements per
163    /// byte, so `packed_size_bytes(3)` returns `Some(2)` instead of the
164    /// conservative `size_bytes() * 3 == 3`. Variable-width types return
165    /// `Ok(None)`; arithmetic overflow returns an actionable error instead of
166    /// saturating.
167    ///
168    /// # Errors
169    ///
170    /// Returns an error when bit or byte arithmetic overflows host `usize`.
171    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}