1use std::fmt::Formatter;
7use std::slice;
8
9use arrow_array::{Array, FixedSizeBinaryArray, builder::BooleanBufferBuilder};
10use arrow_buffer::{Buffer, MutableBuffer};
11use arrow_data::ArrayData;
12use arrow_schema::{ArrowError, DataType, Field as ArrowField};
13use half::bf16;
14
15use crate::{ARROW_EXT_NAME_KEY, FloatArray};
16
17pub const BFLOAT16_EXT_NAME: &str = "lance.bfloat16";
19
20pub fn is_bfloat16_field(field: &ArrowField) -> bool {
25 field.data_type() == &DataType::FixedSizeBinary(2)
26 && field
27 .metadata()
28 .get(ARROW_EXT_NAME_KEY)
29 .map(|name| name == BFLOAT16_EXT_NAME)
30 .unwrap_or_default()
31}
32
33#[derive(Debug)]
37pub struct BFloat16Type {}
38
39#[derive(Clone)]
43pub struct BFloat16Array {
44 inner: FixedSizeBinaryArray,
45}
46
47impl std::fmt::Debug for BFloat16Array {
48 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
49 write!(f, "BFloat16Array\n[\n")?;
50 from_arrow::print_long_array(&self.inner, f, |array, i, f| {
51 if array.is_null(i) {
52 write!(f, "null")
53 } else {
54 let binary_values = array.value(i);
55 let value =
56 bf16::from_bits(u16::from_le_bytes([binary_values[0], binary_values[1]]));
57 write!(f, "{:?}", value)
58 }
59 })?;
60 write!(f, "]")
61 }
62}
63
64impl BFloat16Array {
65 pub fn from_iter_values(iter: impl IntoIterator<Item = bf16>) -> Self {
66 let values: Vec<bf16> = iter.into_iter().collect();
67 values.into()
68 }
69
70 pub fn len(&self) -> usize {
71 self.inner.len()
72 }
73
74 pub fn is_empty(&self) -> bool {
75 self.inner.is_empty()
76 }
77
78 pub fn is_null(&self, i: usize) -> bool {
79 self.inner.is_null(i)
80 }
81
82 pub fn null_count(&self) -> usize {
83 self.inner.null_count()
84 }
85
86 pub fn iter(&self) -> BFloat16Iter<'_> {
87 BFloat16Iter {
88 array: self,
89 index: 0,
90 }
91 }
92
93 pub fn value(&self, i: usize) -> bf16 {
94 assert!(
95 i < self.len(),
96 "Trying to access an element at index {} from a BFloat16Array of length {}",
97 i,
98 self.len()
99 );
100 unsafe { self.value_unchecked(i) }
103 }
104
105 pub unsafe fn value_unchecked(&self, i: usize) -> bf16 {
108 let binary_value = self.inner.value_unchecked(i);
109 bf16::from_bits(u16::from_le_bytes([binary_value[0], binary_value[1]]))
110 }
111
112 pub fn into_inner(self) -> FixedSizeBinaryArray {
113 self.inner
114 }
115}
116
117impl FromIterator<Option<bf16>> for BFloat16Array {
118 fn from_iter<I: IntoIterator<Item = Option<bf16>>>(iter: I) -> Self {
119 let mut buffer = MutableBuffer::new(10);
120 let mut nulls = BooleanBufferBuilder::new(10);
122 let mut len = 0;
123
124 for maybe_value in iter {
125 if let Some(value) = maybe_value {
126 let bytes = value.to_le_bytes();
127 buffer.extend(bytes);
128 } else {
129 buffer.extend([0u8, 0u8]);
130 }
131 nulls.append(maybe_value.is_some());
132 len += 1;
133 }
134
135 let null_buffer = nulls.finish();
136 let num_valid = null_buffer.count_set_bits();
137 let null_buffer = if num_valid == len {
138 None
139 } else {
140 Some(null_buffer.into_inner())
141 };
142
143 let array_data = ArrayData::builder(DataType::FixedSizeBinary(2))
144 .len(len)
145 .add_buffer(buffer.into())
146 .null_bit_buffer(null_buffer);
147 let array_data = unsafe { array_data.build_unchecked() };
153 Self {
154 inner: FixedSizeBinaryArray::from(array_data),
155 }
156 }
157}
158
159impl FromIterator<bf16> for BFloat16Array {
160 fn from_iter<I: IntoIterator<Item = bf16>>(iter: I) -> Self {
161 Self::from_iter_values(iter)
162 }
163}
164
165impl From<Vec<bf16>> for BFloat16Array {
166 fn from(data: Vec<bf16>) -> Self {
167 let len = data.len();
168 let raw: Vec<u16> = bytemuck::cast_vec(data);
174 let array_data = ArrayData::builder(DataType::FixedSizeBinary(2))
175 .len(len)
176 .add_buffer(Buffer::from_vec(raw));
177 let array_data = unsafe { array_data.build_unchecked() };
182 Self {
183 inner: FixedSizeBinaryArray::from(array_data),
184 }
185 }
186}
187
188impl TryFrom<FixedSizeBinaryArray> for BFloat16Array {
189 type Error = ArrowError;
190
191 fn try_from(value: FixedSizeBinaryArray) -> Result<Self, Self::Error> {
192 if value.value_length() == 2 {
193 Ok(Self { inner: value })
194 } else {
195 Err(ArrowError::InvalidArgumentError(
196 "FixedSizeBinaryArray must have a value length of 2".to_string(),
197 ))
198 }
199 }
200}
201
202impl PartialEq<Self> for BFloat16Array {
203 fn eq(&self, other: &Self) -> bool {
204 self.inner.eq(&other.inner)
205 }
206}
207
208pub struct BFloat16Iter<'a> {
209 array: &'a BFloat16Array,
210 index: usize,
211}
212
213impl<'a> Iterator for BFloat16Iter<'a> {
214 type Item = Option<bf16>;
215
216 fn next(&mut self) -> Option<Self::Item> {
217 if self.index >= self.array.len() {
218 return None;
219 }
220 let i = self.index;
221 self.index += 1;
222 if self.array.is_null(i) {
223 Some(None)
224 } else {
225 Some(Some(self.array.value(i)))
226 }
227 }
228}
229
230mod from_arrow {
232 use arrow_array::Array;
233
234 pub(super) fn print_long_array<A, F>(
236 array: &A,
237 f: &mut std::fmt::Formatter,
238 print_item: F,
239 ) -> std::fmt::Result
240 where
241 A: Array,
242 F: Fn(&A, usize, &mut std::fmt::Formatter) -> std::fmt::Result,
243 {
244 let head = std::cmp::min(10, array.len());
245
246 for i in 0..head {
247 if array.is_null(i) {
248 writeln!(f, " null,")?;
249 } else {
250 write!(f, " ")?;
251 print_item(array, i, f)?;
252 writeln!(f, ",")?;
253 }
254 }
255 if array.len() > 10 {
256 if array.len() > 20 {
257 writeln!(f, " ...{} elements...,", array.len() - 20)?;
258 }
259
260 let tail = std::cmp::max(head, array.len() - 10);
261
262 for i in tail..array.len() {
263 if array.is_null(i) {
264 writeln!(f, " null,")?;
265 } else {
266 write!(f, " ")?;
267 print_item(array, i, f)?;
268 writeln!(f, ",")?;
269 }
270 }
271 }
272 Ok(())
273 }
274}
275
276impl FloatArray<BFloat16Type> for FixedSizeBinaryArray {
277 type FloatType = BFloat16Type;
278
279 fn as_slice(&self) -> &[bf16] {
302 assert_eq!(
303 self.value_length(),
304 2,
305 "BFloat16 arrays must use FixedSizeBinary(2) storage"
306 );
307 debug_assert_eq!(
308 (self.value_data().as_ptr() as usize) % std::mem::align_of::<bf16>(),
309 0,
310 "BFloat16 value buffer must be at least 2-byte aligned"
311 );
312 unsafe {
338 slice::from_raw_parts(
339 self.value_data().as_ptr() as *const bf16,
340 self.value_data().len() / 2,
341 )
342 }
343 }
344
345 fn from_values(values: Vec<bf16>) -> Self {
346 BFloat16Array::from(values).into_inner()
347 }
348}
349
350#[cfg(test)]
351mod tests {
352 use super::*;
353
354 #[test]
355 fn test_basics() {
356 let values: Vec<f32> = vec![1.0, 2.0, 3.0];
357 let values: Vec<bf16> = values.iter().map(|v| bf16::from_f32(*v)).collect();
358
359 let array = BFloat16Array::from_iter_values(values.clone());
360 let array2 = BFloat16Array::from(values.clone());
361 assert_eq!(array, array2);
362 assert_eq!(array.len(), 3);
363
364 let inner = array2.clone().into_inner();
369 let raw_bytes: Vec<u8> = (0..inner.len())
370 .flat_map(|i| inner.value(i).to_vec())
371 .collect();
372 assert_eq!(raw_bytes, vec![0x80, 0x3F, 0x00, 0x40, 0x40, 0x40]);
373
374 let expected_fmt = "BFloat16Array\n[\n 1.0,\n 2.0,\n 3.0,\n]";
375 assert_eq!(expected_fmt, format!("{:?}", array));
376
377 for (expected, value) in values.iter().zip(array.iter()) {
378 assert_eq!(Some(*expected), value);
379 }
380
381 for (expected, value) in values.as_slice().iter().zip(array2.iter()) {
382 assert_eq!(Some(*expected), value);
383 }
384
385 let arrow_array = array.into_inner();
386 assert_eq!(arrow_array.as_slice(), values.as_slice());
387 }
388
389 #[test]
390 fn test_nulls() {
391 let values: Vec<Option<bf16>> =
392 vec![Some(bf16::from_f32(1.0)), None, Some(bf16::from_f32(3.0))];
393 let array = BFloat16Array::from_iter(values.clone());
394 assert_eq!(array.len(), 3);
395 assert_eq!(array.null_count(), 1);
396
397 let expected_fmt = "BFloat16Array\n[\n 1.0,\n null,\n 3.0,\n]";
398 assert_eq!(expected_fmt, format!("{:?}", array));
399
400 for (expected, value) in values.iter().zip(array.iter()) {
401 assert_eq!(*expected, value);
402 }
403 }
404}