1use 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 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, ¶ms, 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, ¶ms, 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 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
263pub 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(¢roids, &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(¢roids, &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}