1pub 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
48pub fn buffer_size_bytes(ty: ElementType, element_count: u64) -> u64 {
81 let bits = ty.bit_width() as u64;
82 if bits == 0 {
85 return 0;
86 }
87 if bits >= 8 {
88 element_count * (bits / 8)
90 } else if bits == 6 {
91 element_count.div_ceil(4) * 3
93 } else {
94 (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 #[test]
108 fn buffer_size_float32_60_elements() {
109 assert_eq!(buffer_size_bytes(ElementType::Float32, 60), 240);
111 }
112
113 #[test]
114 fn buffer_size_int64_1_element() {
115 assert_eq!(buffer_size_bytes(ElementType::Int64, 1), 8);
117 }
118
119 #[test]
120 fn buffer_size_float128_2_elements() {
121 assert_eq!(buffer_size_bytes(ElementType::Float128, 2), 32);
123 }
124
125 #[test]
126 fn buffer_size_uint8_5_elements() {
127 assert_eq!(buffer_size_bytes(ElementType::Uint8, 5), 5);
129 }
130
131 #[test]
132 fn buffer_size_float16_4_elements() {
133 assert_eq!(buffer_size_bytes(ElementType::Float16, 4), 8);
135 }
136
137 #[test]
138 fn buffer_size_complex64_3_elements() {
139 assert_eq!(buffer_size_bytes(ElementType::Complex64, 3), 24);
141 }
142
143 #[test]
144 fn buffer_size_complex128_1_element() {
145 assert_eq!(buffer_size_bytes(ElementType::Complex128, 1), 16);
147 }
148
149 #[test]
156 fn buffer_size_bool_9_elements() {
157 assert_eq!(buffer_size_bytes(ElementType::Bool, 9), 2);
159 }
160
161 #[test]
162 fn buffer_size_bool_8_elements_exact_byte() {
163 assert_eq!(buffer_size_bytes(ElementType::Bool, 8), 1);
165 }
166
167 #[test]
168 fn buffer_size_bool_1_element() {
169 assert_eq!(buffer_size_bytes(ElementType::Bool, 1), 1);
171 }
172
173 #[test]
174 fn buffer_size_bool_16_elements() {
175 assert_eq!(buffer_size_bytes(ElementType::Bool, 16), 2);
177 }
178
179 #[test]
181 fn buffer_size_int4_7_elements() {
182 assert_eq!(buffer_size_bytes(ElementType::Int4, 7), 4);
184 }
185
186 #[test]
187 fn buffer_size_int4_8_elements_exact_byte() {
188 assert_eq!(buffer_size_bytes(ElementType::Int4, 8), 4);
190 }
191
192 #[test]
193 fn buffer_size_uint4_1_element() {
194 assert_eq!(buffer_size_bytes(ElementType::Uint4, 1), 1);
196 }
197
198 #[test]
199 fn buffer_size_uint4_2_elements() {
200 assert_eq!(buffer_size_bytes(ElementType::Uint4, 2), 1);
202 }
203
204 #[test]
206 fn buffer_size_int2_5_elements() {
207 assert_eq!(buffer_size_bytes(ElementType::Int2, 5), 2);
209 }
210
211 #[test]
212 fn buffer_size_uint2_4_elements_exact_byte() {
213 assert_eq!(buffer_size_bytes(ElementType::Uint2, 4), 1);
215 }
216
217 #[test]
219 fn buffer_size_float4e2m1_3_elements() {
220 assert_eq!(buffer_size_bytes(ElementType::Float4E2M1, 3), 2);
222 }
223
224 #[test]
230 fn buffer_size_float6e2m3_4_elements_exact_group() {
231 assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 4), 3);
233 }
234
235 #[test]
236 fn buffer_size_float6e2m3_5_elements_partial_group() {
237 assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 5), 6);
239 }
240
241 #[test]
242 fn buffer_size_float6e2m3_8_elements_two_full_groups() {
243 assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 8), 6);
245 }
246
247 #[test]
248 fn buffer_size_float6e2m3_9_elements_partial_third_group() {
249 assert_eq!(buffer_size_bytes(ElementType::Float6E2M3, 9), 9);
251 }
252
253 #[test]
254 fn buffer_size_float6e3m2_4_elements() {
255 assert_eq!(buffer_size_bytes(ElementType::Float6E3M2, 4), 3);
257 }
258
259 #[test]
260 fn buffer_size_float6e3m2_1_element() {
261 assert_eq!(buffer_size_bytes(ElementType::Float6E3M2, 1), 3);
263 }
264
265 #[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 #[test]
297 fn buffer_size_large_count_uint8_no_overflow() {
298 assert_eq!(buffer_size_bytes(ElementType::Uint8, u64::MAX), u64::MAX);
300 }
301
302 #[test]
305 fn buffer_size_large_count_int4_no_overflow() {
306 let n = 1u64 << 62;
308 let expected = n / 2; assert_eq!(buffer_size_bytes(ElementType::Int4, n), expected);
310 }
311
312 #[test]
314 fn buffer_size_large_count_bool_no_overflow() {
315 let n = 1u64 << 63;
317 let expected = n / 8;
318 assert_eq!(buffer_size_bytes(ElementType::Bool, n), expected);
319 }
320}