Skip to main content

hurray_core/
lib.rs

1//! # hurray-core
2//!
3//! Core types for the hurray tensor interchange format.
4//!
5//! This crate provides the format types, tensor descriptor, buffer handle, and
6//! quantization descriptors. It has no I/O and no async dependencies — it is the
7//! foundation for all other hurray crates.
8//!
9//! ## Feature flags
10//!
11//! | Feature | Effect |
12//! |---------|--------|
13//! | `serde` | Derives `serde::Serialize` / `serde::Deserialize` for all public types |
14
15pub mod buffer;
16pub mod composite;
17pub mod descriptor;
18pub mod element_type;
19pub mod error;
20pub mod layout;
21pub mod quantization;
22pub mod shape;
23
24pub use buffer::{
25    validate_colocation, BufferHandle, DeviceTag, MemoryClass, PrivateMemoryClass, PrivateTag,
26    SyncMode, MIN_BUFFER_ALIGNMENT, PAGE_ALIGNMENT,
27};
28pub use composite::{CompositeTensor, CompositeValidator};
29pub use descriptor::{
30    CompositeMemberDescriptor, DescriptorFlags, ExtensionTypeDescriptor, MemberRole,
31    ShardDescriptor, Statistics, StatisticsMask, TensorDescriptor, DESCRIPTOR_VERSION_MAJOR,
32    DESCRIPTOR_VERSION_MINOR, MAGIC,
33};
34pub use element_type::ElementType;
35pub use error::{Error, Result};
36pub use layout::{
37    byte_address_from_element_offset, BlockPagedLayout, BlockTableIndexType, CombineOp,
38    CompositeLayout, CompositionRule, CsfLayout, ElementAddress, KvRole, LayoutDescriptor,
39};
40pub use quantization::{
41    validate_axis, validate_buffer_placement, Mxfp, Nf4, PerBlockAffine, PerChannelAffine,
42    PerTensorAffine, QuantizationDescriptor, QuantizationSchemeTag, MXFP_CANONICAL_BLOCK_SIZE,
43    MXFP_MAX_BLOCK_SIZE, MXFP_MIN_BLOCK_SIZE, NF4_LUT, NF4_MIN_BLOCK_SIZE,
44    PER_BLOCK_AFFINE_MIN_BLOCK_SIZE,
45};
46pub use shape::{Shape, DYNAMIC, MAX_RANK};
47
48/// Returns the minimum buffer size in bytes required to store `element_count`
49/// contiguous elements of the given type.
50///
51/// This implements the buffer-size formulas defined in
52/// `docs/spec/element-types.md § Buffer Size Calculation`:
53///
54/// - Whole-byte types (`bit_width ≥ 8`): `element_count × (bit_width / 8)`
55/// - 6-bit types (`float6_e2m3`, `float6_e3m2`): `⌈N / 4⌉ × 3`
56/// - Other sub-byte types (`bit_width ∈ {1, 2, 4}`): `⌈N × bit_width / 8⌉`
57///
58/// Uses 128-bit intermediate arithmetic to avoid overflow for large `element_count`.
59///
60/// # Examples
61///
62/// ```
63/// use hurray_core::{ElementType, buffer_size_bytes};
64///
65/// // float32: 60 elements × 4 bytes = 240 bytes
66/// assert_eq!(buffer_size_bytes(ElementType::Float32, 60), 240);
67///
68/// // int4: 7 elements → ⌈7/2⌉ = 4 bytes
69/// assert_eq!(buffer_size_bytes(ElementType::Int4, 7), 4);
70///
71/// // float6_e2m3: 5 elements → ⌈5/4⌉ × 3 = 6 bytes
72/// assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 5), 6);
73///
74/// // bool: 9 elements → ⌈9/8⌉ = 2 bytes
75/// assert_eq!(buffer_size_bytes(ElementType::Bool, 9), 2);
76///
77/// // Zero elements always yields zero bytes.
78/// assert_eq!(buffer_size_bytes(ElementType::Float32, 0), 0);
79/// ```
80pub fn buffer_size_bytes(ty: ElementType, element_count: u64) -> u64 {
81    let bits = ty.bit_width() as u64;
82    // Extension types declare bit_width == 0 (sentinel); caller must use
83    // ExtensionTypeDescriptor for buffer sizing — return 0 here.
84    if bits == 0 {
85        return 0;
86    }
87    if bits >= 8 {
88        // Whole-byte type: exact multiplication, no rounding needed.
89        element_count * (bits / 8)
90    } else if bits == 6 {
91        // 6-bit packing: 4 elements per 3 bytes (ceil(N/4)*3).
92        element_count.div_ceil(4) * 3
93    } else {
94        // Sub-byte power-of-two packing (bits ∈ {1, 2, 4}): ceil(N*bits/8).
95        (element_count as u128 * bits as u128).div_ceil(8) as u64
96    }
97}
98
99#[cfg(test)]
100mod tests {
101    use super::*;
102
103    // ── Whole-byte types ─────────────────────────────────────────────────────
104
105    /// Spec § element-types Buffer Size Calculation:
106    /// whole-byte types: element_count × (bit_width / 8).
107    #[test]
108    fn buffer_size_float32_60_elements() {
109        // 60 × 4 bytes = 240
110        assert_eq!(buffer_size_bytes(ElementType::Float32, 60), 240);
111    }
112
113    #[test]
114    fn buffer_size_int64_1_element() {
115        // 1 × 8 bytes = 8
116        assert_eq!(buffer_size_bytes(ElementType::Int64, 1), 8);
117    }
118
119    #[test]
120    fn buffer_size_float128_2_elements() {
121        // 2 × 16 bytes = 32
122        assert_eq!(buffer_size_bytes(ElementType::Float128, 2), 32);
123    }
124
125    #[test]
126    fn buffer_size_uint8_5_elements() {
127        // 5 × 1 byte = 5
128        assert_eq!(buffer_size_bytes(ElementType::Uint8, 5), 5);
129    }
130
131    #[test]
132    fn buffer_size_float16_4_elements() {
133        // 4 × 2 bytes = 8
134        assert_eq!(buffer_size_bytes(ElementType::Float16, 4), 8);
135    }
136
137    #[test]
138    fn buffer_size_complex64_3_elements() {
139        // Complex64 = 64 bits = 8 bytes; 3 × 8 = 24
140        assert_eq!(buffer_size_bytes(ElementType::Complex64, 3), 24);
141    }
142
143    #[test]
144    fn buffer_size_complex128_1_element() {
145        // Complex128 = 128 bits = 16 bytes; 1 × 16 = 16
146        assert_eq!(buffer_size_bytes(ElementType::Complex128, 1), 16);
147    }
148
149    // ── Sub-byte power-of-two types (bits ∈ {1, 2, 4}) ──────────────────────
150
151    /// Spec § element-types Buffer Size Calculation:
152    /// sub-byte power-of-two: ⌈N × bit_width / 8⌉.
153
154    // Bool (1 bit)
155    #[test]
156    fn buffer_size_bool_9_elements() {
157        // ⌈9 × 1 / 8⌉ = ⌈9/8⌉ = 2
158        assert_eq!(buffer_size_bytes(ElementType::Bool, 9), 2);
159    }
160
161    #[test]
162    fn buffer_size_bool_8_elements_exact_byte() {
163        // ⌈8 × 1 / 8⌉ = 1
164        assert_eq!(buffer_size_bytes(ElementType::Bool, 8), 1);
165    }
166
167    #[test]
168    fn buffer_size_bool_1_element() {
169        // ⌈1/8⌉ = 1
170        assert_eq!(buffer_size_bytes(ElementType::Bool, 1), 1);
171    }
172
173    #[test]
174    fn buffer_size_bool_16_elements() {
175        // ⌈16/8⌉ = 2
176        assert_eq!(buffer_size_bytes(ElementType::Bool, 16), 2);
177    }
178
179    // Int4 / Uint4 (4 bits)
180    #[test]
181    fn buffer_size_int4_7_elements() {
182        // ⌈7 × 4 / 8⌉ = ⌈28/8⌉ = 4 (rounds up from 3.5)
183        assert_eq!(buffer_size_bytes(ElementType::Int4, 7), 4);
184    }
185
186    #[test]
187    fn buffer_size_int4_8_elements_exact_byte() {
188        // ⌈8 × 4 / 8⌉ = 4 (exactly 4 bytes)
189        assert_eq!(buffer_size_bytes(ElementType::Int4, 8), 4);
190    }
191
192    #[test]
193    fn buffer_size_uint4_1_element() {
194        // ⌈1 × 4 / 8⌉ = 1
195        assert_eq!(buffer_size_bytes(ElementType::Uint4, 1), 1);
196    }
197
198    #[test]
199    fn buffer_size_uint4_2_elements() {
200        // ⌈2 × 4 / 8⌉ = 1 (exactly 1 byte)
201        assert_eq!(buffer_size_bytes(ElementType::Uint4, 2), 1);
202    }
203
204    // Int2 / Uint2 (2 bits)
205    #[test]
206    fn buffer_size_int2_5_elements() {
207        // ⌈5 × 2 / 8⌉ = ⌈10/8⌉ = 2
208        assert_eq!(buffer_size_bytes(ElementType::Int2, 5), 2);
209    }
210
211    #[test]
212    fn buffer_size_uint2_4_elements_exact_byte() {
213        // ⌈4 × 2 / 8⌉ = 1 (exactly 1 byte)
214        assert_eq!(buffer_size_bytes(ElementType::Uint2, 4), 1);
215    }
216
217    // Float4E2M1 (4 bits)
218    #[test]
219    fn buffer_size_float4e2m1_3_elements() {
220        // ⌈3 × 4 / 8⌉ = ⌈12/8⌉ = 2
221        assert_eq!(buffer_size_bytes(ElementType::Float4E2M1, 3), 2);
222    }
223
224    // ── 6-bit packing (Float6E2M3, Float6E3M2) ───────────────────────────────
225
226    /// Spec § element-types Buffer Size Calculation:
227    /// 6-bit types: ⌈N / 4⌉ × 3  (4 elements pack into 3 bytes).
228
229    #[test]
230    fn buffer_size_float6e2m3_4_elements_exact_group() {
231        // ⌈4/4⌉ × 3 = 1 × 3 = 3
232        assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 4), 3);
233    }
234
235    #[test]
236    fn buffer_size_float6e2m3_5_elements_partial_group() {
237        // ⌈5/4⌉ × 3 = 2 × 3 = 6
238        assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 5), 6);
239    }
240
241    #[test]
242    fn buffer_size_float6e2m3_8_elements_two_full_groups() {
243        // ⌈8/4⌉ × 3 = 2 × 3 = 6
244        assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 8), 6);
245    }
246
247    #[test]
248    fn buffer_size_float6e2m3_9_elements_partial_third_group() {
249        // ⌈9/4⌉ × 3 = 3 × 3 = 9
250        assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 9), 9);
251    }
252
253    #[test]
254    fn buffer_size_float6e3m2_4_elements() {
255        // ⌈4/4⌉ × 3 = 3
256        assert_eq!(buffer_size_bytes(ElementType::Float6E3M2, 4), 3);
257    }
258
259    #[test]
260    fn buffer_size_float6e3m2_1_element() {
261        // ⌈1/4⌉ × 3 = 1 × 3 = 3
262        assert_eq!(buffer_size_bytes(ElementType::Float6E3M2, 1), 3);
263    }
264
265    // ── Zero elements ────────────────────────────────────────────────────────
266
267    /// Zero elements must always produce 0 bytes, regardless of element type.
268    #[test]
269    fn buffer_size_zero_elements_whole_byte_type() {
270        assert_eq!(buffer_size_bytes(ElementType::Float32, 0), 0);
271    }
272
273    #[test]
274    fn buffer_size_zero_elements_bool() {
275        assert_eq!(buffer_size_bytes(ElementType::Bool, 0), 0);
276    }
277
278    #[test]
279    fn buffer_size_zero_elements_int4() {
280        assert_eq!(buffer_size_bytes(ElementType::Int4, 0), 0);
281    }
282
283    #[test]
284    fn buffer_size_zero_elements_int2() {
285        assert_eq!(buffer_size_bytes(ElementType::Int2, 0), 0);
286    }
287
288    #[test]
289    fn buffer_size_zero_elements_float6e2m3() {
290        assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 0), 0);
291    }
292
293    // ── Large element counts (no overflow) ───────────────────────────────────
294
295    /// A large element count for a byte-wide type must not overflow u64.
296    #[test]
297    fn buffer_size_large_count_uint8_no_overflow() {
298        // u64::MAX elements of uint8 (1 byte each) = u64::MAX bytes
299        assert_eq!(buffer_size_bytes(ElementType::Uint8, u64::MAX), u64::MAX);
300    }
301
302    /// A large element count for a 4-bit type uses 128-bit intermediate
303    /// arithmetic so must not overflow.
304    #[test]
305    fn buffer_size_large_count_int4_no_overflow() {
306        // 1 << 62 elements × 4 bits = 1 << 64 bits / 8 = 1 << 61 bytes — fits in u64
307        let n = 1u64 << 62;
308        let expected = n / 2; // ⌈n*4/8⌉ = n/2 (n is even)
309        assert_eq!(buffer_size_bytes(ElementType::Int4, n), expected);
310    }
311
312    /// A large element count for bool uses 128-bit intermediates.
313    #[test]
314    fn buffer_size_large_count_bool_no_overflow() {
315        // 1 << 63 bools = (1 << 63) / 8 bytes = 1 << 60 bytes — fits in u64
316        let n = 1u64 << 63;
317        let expected = n / 8;
318        assert_eq!(buffer_size_bytes(ElementType::Bool, n), expected);
319    }
320}