Skip to main content

lance_index/vector/
utils.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use arrow::{
5    array::{ArrayData, make_array},
6    buffer::Buffer,
7    compute::cast,
8};
9use arrow_array::types::{Float16Type, Float32Type, Float64Type};
10use arrow_array::{Array, ArrayRef, BooleanArray, FixedSizeListArray, cast::AsArray};
11use arrow_schema::{DataType, Field};
12use lance_arrow::{BufferExt, DataTypeExt, FixedSizeListArrayExt};
13use lance_core::{Error, Result};
14use lance_linalg::distance::DistanceType;
15use prost::bytes;
16use std::sync::LazyLock;
17use std::{ops::Range, sync::Arc};
18
19use super::pb;
20use crate::pb::Tensor;
21use crate::vector::flat::storage::FlatBinStorage;
22use crate::vector::flat::storage::FlatFloatStorage;
23use crate::vector::hnsw::HNSW;
24use crate::vector::hnsw::builder::{HnswBuildParams, HnswQueryParams};
25use crate::vector::v3::subindex::IvfSubIndex;
26
27enum SimpleIndexStatus {
28    Auto,
29    Enabled,
30    Disabled,
31}
32
33static USE_HNSW_SPEEDUP_INDEXING: LazyLock<SimpleIndexStatus> = LazyLock::new(|| {
34    if let Ok(v) = std::env::var("LANCE_USE_HNSW_SPEEDUP_INDEXING") {
35        if v == "enabled" {
36            SimpleIndexStatus::Enabled
37        } else if v == "disabled" {
38            SimpleIndexStatus::Disabled
39        } else {
40            SimpleIndexStatus::Auto
41        }
42    } else {
43        SimpleIndexStatus::Auto
44    }
45});
46
47#[derive(Debug)]
48pub struct SimpleIndex {
49    store: SimpleStore,
50    index: HNSW,
51}
52
53#[derive(Debug)]
54enum SimpleStore {
55    Float(FlatFloatStorage),
56    Binary(FlatBinStorage),
57}
58
59impl SimpleIndex {
60    fn try_new(store: SimpleStore) -> Result<Self> {
61        let hnsw = match &store {
62            SimpleStore::Float(store) => HNSW::index_vectors(
63                store,
64                HnswBuildParams::default().ef_construction(15).num_edges(12),
65            )?,
66            SimpleStore::Binary(store) => HNSW::index_vectors(
67                store,
68                HnswBuildParams::default().ef_construction(15).num_edges(12),
69            )?,
70        };
71        Ok(Self { store, index: hnsw })
72    }
73
74    // train HNSW over the centroids to speed up finding the nearest clusters,
75    // only train if all conditions are met:
76    //  - the centroids are float16/float32 or uint8 with hamming distance
77    //  - `num_centroids * dimension >= 1_000_000`
78    //      we benchmarked that it's 2x faster in the case of 1024 centroids and 1024 dimensions,
79    //      so set the threshold to 1_000_000.
80    pub fn may_train_index(
81        centroids: ArrayRef,
82        dimension: usize,
83        distance_type: DistanceType,
84    ) -> Result<Option<Self>> {
85        match *USE_HNSW_SPEEDUP_INDEXING {
86            SimpleIndexStatus::Auto => {
87                if centroids.len() < 1_000_000 {
88                    return Ok(None);
89                }
90            }
91            SimpleIndexStatus::Disabled => return Ok(None),
92            _ => {}
93        }
94
95        let store = match (centroids.data_type(), distance_type) {
96            (DataType::Float16 | DataType::Float32 | DataType::Float64, _) => {
97                let fsl = FixedSizeListArray::try_new_from_values(centroids, dimension as i32)?;
98                SimpleStore::Float(FlatFloatStorage::new(fsl, distance_type))
99            }
100            (DataType::UInt8, DistanceType::Hamming) => {
101                let fsl = FixedSizeListArray::try_new_from_values(centroids, dimension as i32)?;
102                SimpleStore::Binary(FlatBinStorage::new(fsl, distance_type))
103            }
104            _ => return Ok(None),
105        };
106        Self::try_new(store).map(Some)
107    }
108
109    pub(crate) fn search(&self, query: ArrayRef) -> Result<(u32, f32)> {
110        let params = HnswQueryParams {
111            ef: 15,
112            lower_bound: None,
113            upper_bound: None,
114            dist_q_c: 0.0,
115            use_acorn: false,
116        };
117        let res = match &self.store {
118            SimpleStore::Float(store) => self.index.search_basic(query, 1, &params, None, store)?,
119            SimpleStore::Binary(store) => {
120                let query = if query.data_type() == &DataType::UInt8 {
121                    query
122                } else {
123                    cast(&query, &DataType::UInt8).map_err(|e| Error::index(e.to_string()))?
124                };
125                self.index.search_basic(query, 1, &params, None, store)?
126            }
127        };
128        Ok((res[0].id, res[0].dist.0))
129    }
130}
131
132#[inline]
133pub(crate) fn do_prefetch<T>(ptrs: Range<*const T>) {
134    // TODO use rust intrinsics instead of x86 intrinsics
135    // TODO finish this
136    unsafe {
137        let (ptr, end_ptr) = (ptrs.start as *const i8, ptrs.end as *const i8);
138        let mut current_ptr = ptr;
139        while current_ptr < end_ptr {
140            const CACHE_LINE_SIZE: usize = 64;
141            #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
142            {
143                use core::arch::x86_64::{_MM_HINT_T0, _mm_prefetch};
144                _mm_prefetch(current_ptr, _MM_HINT_T0);
145            }
146            current_ptr = current_ptr.add(CACHE_LINE_SIZE);
147        }
148    }
149}
150
151impl From<pb::tensor::DataType> for DataType {
152    fn from(dt: pb::tensor::DataType) -> Self {
153        match dt {
154            pb::tensor::DataType::Uint8 => Self::UInt8,
155            pb::tensor::DataType::Uint16 => Self::UInt16,
156            pb::tensor::DataType::Uint32 => Self::UInt32,
157            pb::tensor::DataType::Uint64 => Self::UInt64,
158            pb::tensor::DataType::Float16 => Self::Float16,
159            pb::tensor::DataType::Float32 => Self::Float32,
160            pb::tensor::DataType::Float64 => Self::Float64,
161            pb::tensor::DataType::Bfloat16 => unimplemented!(),
162        }
163    }
164}
165
166impl TryFrom<&DataType> for pb::tensor::DataType {
167    type Error = Error;
168
169    fn try_from(dt: &DataType) -> Result<Self> {
170        match dt {
171            DataType::UInt8 => Ok(Self::Uint8),
172            DataType::UInt16 => Ok(Self::Uint16),
173            DataType::UInt32 => Ok(Self::Uint32),
174            DataType::UInt64 => Ok(Self::Uint64),
175            DataType::Float16 => Ok(Self::Float16),
176            DataType::Float32 => Ok(Self::Float32),
177            DataType::Float64 => Ok(Self::Float64),
178            _ => Err(Error::index(format!(
179                "pb tensor type not supported: {:?}",
180                dt
181            ))),
182        }
183    }
184}
185
186impl TryFrom<DataType> for pb::tensor::DataType {
187    type Error = Error;
188
189    fn try_from(dt: DataType) -> Result<Self> {
190        (&dt).try_into()
191    }
192}
193
194impl TryFrom<&FixedSizeListArray> for pb::Tensor {
195    type Error = Error;
196
197    fn try_from(array: &FixedSizeListArray) -> Result<Self> {
198        let mut tensor = Self::default();
199        tensor.data_type = pb::tensor::DataType::try_from(array.value_type())? as i32;
200        tensor.shape = vec![Array::len(array) as u32, array.value_length() as u32];
201        let flat_array = array.values();
202        tensor.data = flat_array.into_data().buffers()[0].to_vec();
203        Ok(tensor)
204    }
205}
206
207impl TryFrom<&pb::Tensor> for FixedSizeListArray {
208    type Error = Error;
209
210    fn try_from(tensor: &Tensor) -> Result<Self> {
211        if tensor.shape.len() != 2 {
212            return Err(Error::index(format!(
213                "only accept 2-D tensor shape, got: {:?}",
214                tensor.shape
215            )));
216        }
217        let dim = tensor.shape[1] as usize;
218        let num_rows = tensor.shape[0] as usize;
219        let num_values = dim.checked_mul(num_rows).ok_or_else(|| {
220            Error::index(format!(
221                "Tensor shape {:?} exceeds the supported size",
222                tensor.shape
223            ))
224        })?;
225        let data_type = DataType::from(pb::tensor::DataType::try_from(tensor.data_type).unwrap());
226        let expected_data_len =
227            num_values
228                .checked_mul(data_type.byte_width())
229                .ok_or_else(|| {
230                    Error::index(format!(
231                        "Tensor shape {:?} exceeds the supported byte length",
232                        tensor.shape
233                    ))
234                })?;
235        if tensor.data.len() != expected_data_len {
236            return Err(Error::index(format!(
237                "Tensor shape {:?} with data type {data_type} requires {expected_data_len} bytes, got {}",
238                tensor.shape,
239                tensor.data.len()
240            )));
241        }
242
243        let buffer = Buffer::from_bytes_bytes(
244            bytes::Bytes::from(tensor.data.clone()),
245            data_type.byte_width() as u64,
246        );
247        let data = ArrayData::builder(data_type)
248            .len(num_values)
249            .null_count(0)
250            .add_buffer(buffer)
251            .build()?;
252        let flat_array = make_array(data);
253        let field = Field::new("item", flat_array.data_type().clone(), true);
254        Ok(Self::try_new(
255            Arc::new(field),
256            dim as i32,
257            flat_array,
258            None,
259        )?)
260    }
261}
262
263/// Check if all vectors in the FixedSizeListArray are finite
264/// null values are considered as not finite
265/// returns a BooleanArray
266/// with the same length as the FixedSizeListArray
267/// with true for finite values and false for non-finite values
268pub fn is_finite(fsl: &FixedSizeListArray) -> BooleanArray {
269    let is_finite = fsl
270        .iter()
271        .map(|v| match v {
272            Some(v) => match v.data_type() {
273                DataType::Float16 => {
274                    let v = v.as_primitive::<Float16Type>();
275                    Array::null_count(v) == 0 && v.values().iter().all(|v| v.is_finite())
276                }
277                DataType::Float32 => {
278                    let v = v.as_primitive::<Float32Type>();
279                    Array::null_count(v) == 0 && v.values().iter().all(|v| v.is_finite())
280                }
281                DataType::Float64 => {
282                    let v = v.as_primitive::<Float64Type>();
283                    Array::null_count(v) == 0 && v.values().iter().all(|v| v.is_finite())
284                }
285                _ => Array::null_count(&v) == 0,
286            },
287            None => false,
288        })
289        .collect::<Vec<_>>();
290    BooleanArray::from(is_finite)
291}
292
293#[cfg(test)]
294mod tests {
295    use super::*;
296
297    use arrow_array::{Float16Array, Float32Array, Float64Array, UInt8Array};
298    use half::f16;
299    use lance_arrow::FixedSizeListArrayExt;
300    use num_traits::identities::Zero;
301    use rayon::ThreadPoolBuilder;
302
303    use arrow::compute::cast;
304    use rstest::rstest;
305
306    fn build_index(centroids: ArrayRef, dim: usize) -> SimpleIndex {
307        let f32_centroids = cast(&centroids, &DataType::Float32).unwrap();
308        let fsl = FixedSizeListArray::try_new_from_values(f32_centroids, dim as i32).unwrap();
309        let store = SimpleStore::Float(FlatFloatStorage::new(fsl, DistanceType::L2));
310        SimpleIndex::try_new(store).unwrap()
311    }
312
313    fn build_binary_index(centroids: ArrayRef, dim: usize) -> SimpleIndex {
314        let u8_centroids = if centroids.data_type() == &DataType::UInt8 {
315            centroids
316        } else {
317            cast(&centroids, &DataType::UInt8).unwrap()
318        };
319        let fsl = FixedSizeListArray::try_new_from_values(u8_centroids, dim as i32).unwrap();
320        let store = SimpleStore::Binary(FlatBinStorage::new(fsl, DistanceType::Hamming));
321        SimpleIndex::try_new(store).unwrap()
322    }
323
324    #[rstest]
325    #[case::f16(Arc::new(Float16Array::from(
326        (0..100).flat_map(|i| std::iter::repeat_n(f16::from_f32(i as f32), 16)).collect::<Vec<_>>(),
327    )) as ArrayRef, 42.0f32)]
328    #[case::f32(Arc::new(Float32Array::from(
329        (0..100).flat_map(|i| std::iter::repeat_n(i as f32, 16)).collect::<Vec<_>>(),
330    )) as ArrayRef, 42.0f32)]
331    fn test_simple_index_nearest_centroid(#[case] centroids: ArrayRef, #[case] query_val: f32) {
332        let thread_pool = ThreadPoolBuilder::new().num_threads(1).build().unwrap();
333        let index = thread_pool.install(|| build_index(centroids, 16));
334        let query: ArrayRef = Arc::new(Float32Array::from(vec![query_val; 16]));
335        let (id, dist) = index.search(query).unwrap();
336        assert_eq!(id, 42);
337        assert_eq!(dist, 0.0);
338    }
339
340    #[test]
341    fn test_simple_index_nearest_centroid_binary() {
342        let centroids: ArrayRef = Arc::new(UInt8Array::from(
343            (0..100)
344                .flat_map(|i| std::iter::repeat_n(i as u8, 16))
345                .collect::<Vec<_>>(),
346        ));
347        let index = build_binary_index(centroids, 16);
348        let query: ArrayRef = Arc::new(UInt8Array::from(vec![42u8; 16]));
349        let (id, dist) = index.search(query).unwrap();
350        assert_eq!(id, 42);
351        assert_eq!(dist, 0.0);
352    }
353
354    #[test]
355    fn test_simple_index_rejects_f64() {
356        let centroids: ArrayRef = Arc::new(Float64Array::from(vec![0.0; 1600]));
357        let result = SimpleIndex::may_train_index(centroids, 16, DistanceType::L2).unwrap();
358        assert!(result.is_none());
359    }
360
361    #[test]
362    fn test_simple_index_rejects_uint8_non_hamming() {
363        let centroids: ArrayRef = Arc::new(UInt8Array::from(vec![0u8; 1600]));
364        let result = SimpleIndex::may_train_index(centroids, 16, DistanceType::L2).unwrap();
365        assert!(result.is_none());
366    }
367
368    #[test]
369    fn test_fsl_to_tensor() {
370        let fsl =
371            FixedSizeListArray::try_new_from_values(Float16Array::from(vec![f16::zero(); 20]), 5)
372                .unwrap();
373        let tensor = pb::Tensor::try_from(&fsl).unwrap();
374        assert_eq!(tensor.data_type, pb::tensor::DataType::Float16 as i32);
375        assert_eq!(tensor.shape, vec![4, 5]);
376        assert_eq!(tensor.data.len(), 20 * 2);
377        let decoded = FixedSizeListArray::try_from(&tensor).unwrap();
378        assert_eq!(decoded.values().to_data(), fsl.values().to_data());
379
380        let fsl =
381            FixedSizeListArray::try_new_from_values(Float32Array::from(vec![0.0; 20]), 5).unwrap();
382        let tensor = pb::Tensor::try_from(&fsl).unwrap();
383        assert_eq!(tensor.data_type, pb::tensor::DataType::Float32 as i32);
384        assert_eq!(tensor.shape, vec![4, 5]);
385        assert_eq!(tensor.data.len(), 20 * 4);
386        let decoded = FixedSizeListArray::try_from(&tensor).unwrap();
387        assert_eq!(decoded.values().to_data(), fsl.values().to_data());
388
389        let fsl =
390            FixedSizeListArray::try_new_from_values(Float64Array::from(vec![0.0; 20]), 5).unwrap();
391        let tensor = pb::Tensor::try_from(&fsl).unwrap();
392        assert_eq!(tensor.data_type, pb::tensor::DataType::Float64 as i32);
393        assert_eq!(tensor.shape, vec![4, 5]);
394        assert_eq!(tensor.data.len(), 20 * 8);
395        let decoded = FixedSizeListArray::try_from(&tensor).unwrap();
396        assert_eq!(decoded.values().to_data(), fsl.values().to_data());
397    }
398
399    #[rstest]
400    #[case::too_short(vec![0; 7])]
401    #[case::too_long(vec![0; 9])]
402    fn test_tensor_to_fsl_rejects_invalid_data_length(#[case] data: Vec<u8>) {
403        let tensor = pb::Tensor {
404            data_type: pb::tensor::DataType::Uint32 as i32,
405            shape: vec![1, 2],
406            data,
407        };
408
409        let error = FixedSizeListArray::try_from(&tensor).unwrap_err();
410        assert!(error.to_string().contains("requires 8 bytes"));
411    }
412}