Skip to main content

lance_index/vector/flat/
index.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Flat Vector Index.
5//!
6
7use lance_core::utils::row_addr_remap::RowAddrRemap;
8use std::collections::BinaryHeap;
9use std::sync::Arc;
10
11use arrow::array::AsArray;
12use arrow_array::{Array, ArrayRef, Float32Array, RecordBatch, UInt64Array};
13use arrow_schema::{DataType, Field, Schema, SchemaRef};
14use lance_core::deepsize::DeepSizeOf;
15use lance_core::{Error, ROW_ID_FIELD, Result};
16use lance_file::previous::reader::FileReader as PreviousFileReader;
17use lance_linalg::distance::DistanceType;
18use serde::{Deserialize, Serialize};
19
20use crate::{
21    metrics::MetricsCollector,
22    prefilter::PreFilter,
23    vector::{
24        ApproxMode, DIST_COL, Query,
25        graph::{OrderedFloat, OrderedNode},
26        quantizer::{Quantization, QuantizationType, Quantizer, QuantizerMetadata},
27        storage::{
28            DistCalculator, DistanceCalculatorOptions, QueryResidual, QueryScratch, VectorStore,
29        },
30        v3::subindex::IvfSubIndex,
31    },
32};
33
34use super::storage::{FLAT_COLUMN, FlatBinStorage, FlatFloatStorage};
35
36#[inline(always)]
37fn push_candidate_local(
38    res: &mut BinaryHeap<OrderedNode<u64>>,
39    k: usize,
40    row_id: u64,
41    dist: OrderedFloat,
42) {
43    if k == 0 {
44        return;
45    }
46    if res.len() < k {
47        res.push(OrderedNode::new(row_id, dist));
48    } else if res.peek().is_some_and(|node| node.dist > dist) {
49        res.pop();
50        res.push(OrderedNode::new(row_id, dist));
51    }
52}
53
54/// A Flat index is any index that stores no metadata, and
55/// during query, it simply scans over the storage and returns the top k results
56#[derive(Debug, Clone, Default, DeepSizeOf)]
57pub struct FlatIndex {}
58
59use std::sync::LazyLock;
60
61static ANN_SEARCH_SCHEMA: LazyLock<SchemaRef> = LazyLock::new(|| {
62    Schema::new(vec![
63        Field::new(DIST_COL, DataType::Float32, true),
64        ROW_ID_FIELD.clone(),
65    ])
66    .into()
67});
68
69#[derive(Default)]
70pub struct FlatQueryParams {
71    lower_bound: Option<f32>,
72    upper_bound: Option<f32>,
73    dist_q_c: f32,
74    approx_mode: ApproxMode,
75}
76
77impl From<&Query> for FlatQueryParams {
78    fn from(q: &Query) -> Self {
79        Self {
80            lower_bound: q.lower_bound,
81            upper_bound: q.upper_bound,
82            dist_q_c: q.dist_q_c,
83            approx_mode: q.approx_mode,
84        }
85    }
86}
87
88impl IvfSubIndex for FlatIndex {
89    type QueryParams = FlatQueryParams;
90    type BuildParams = ();
91
92    fn name() -> &'static str {
93        "FLAT"
94    }
95
96    fn metadata_key() -> &'static str {
97        "lance:flat"
98    }
99
100    fn schema() -> arrow_schema::SchemaRef {
101        Schema::new(vec![Field::new("__flat_marker", DataType::UInt64, false)]).into()
102    }
103
104    fn search(
105        &self,
106        query: ArrayRef,
107        k: usize,
108        params: Self::QueryParams,
109        storage: &impl VectorStore,
110        prefilter: Arc<dyn PreFilter>,
111        metrics: &dyn MetricsCollector,
112    ) -> Result<RecordBatch> {
113        let mut scratch = QueryScratch::new();
114        self.search_with_scratch(
115            query,
116            k,
117            params,
118            storage,
119            prefilter,
120            metrics,
121            None,
122            &mut scratch,
123        )
124    }
125
126    fn search_with_scratch(
127        &self,
128        query: ArrayRef,
129        k: usize,
130        params: Self::QueryParams,
131        storage: &impl VectorStore,
132        prefilter: Arc<dyn PreFilter>,
133        metrics: &dyn MetricsCollector,
134        residual: Option<QueryResidual<'_>>,
135        scratch: &mut QueryScratch,
136    ) -> Result<RecordBatch> {
137        let is_range_query = params.lower_bound.is_some() || params.upper_bound.is_some();
138        let row_ids = storage.row_ids();
139        let dist_calc = storage.dist_calculator_with_scratch(
140            query,
141            params.dist_q_c,
142            residual,
143            &mut scratch.query_f32,
144            DistanceCalculatorOptions {
145                approx_mode: params.approx_mode,
146            },
147        );
148        let mut res = BinaryHeap::with_capacity(k);
149        metrics.record_comparisons(storage.len());
150
151        match prefilter.is_empty() {
152            true => {
153                dist_calc.distance_all_with_scratch(
154                    k,
155                    &mut scratch.distances,
156                    &mut scratch.u16,
157                    &mut scratch.u8,
158                    &mut scratch.u32,
159                );
160                let dists = scratch.distances.iter().copied();
161
162                if is_range_query {
163                    let lower_bound = params.lower_bound.unwrap_or(f32::MIN).into();
164                    let upper_bound = params.upper_bound.unwrap_or(f32::MAX).into();
165
166                    for (&row_id, dist) in row_ids.zip(dists) {
167                        let dist = dist.into();
168                        if dist < lower_bound || dist >= upper_bound {
169                            continue;
170                        }
171                        push_candidate_local(&mut res, k, row_id, dist);
172                    }
173                } else {
174                    for (&row_id, dist) in row_ids.zip(dists) {
175                        let dist = dist.into();
176                        push_candidate_local(&mut res, k, row_id, dist);
177                    }
178                }
179            }
180            false => {
181                let row_addr_mask = prefilter.mask();
182                if is_range_query {
183                    let lower_bound = params.lower_bound.unwrap_or(f32::MIN).into();
184                    let upper_bound = params.upper_bound.unwrap_or(f32::MAX).into();
185                    for (id, &row_addr) in row_ids.enumerate() {
186                        if !row_addr_mask.selected(row_addr) {
187                            continue;
188                        }
189                        let dist = dist_calc.distance(id as u32).into();
190                        if dist < lower_bound || dist >= upper_bound {
191                            continue;
192                        }
193
194                        push_candidate_local(&mut res, k, row_addr, dist);
195                    }
196                } else {
197                    for (id, &row_addr) in row_ids.enumerate() {
198                        if !row_addr_mask.selected(row_addr) {
199                            continue;
200                        }
201
202                        let dist = dist_calc.distance(id as u32).into();
203                        push_candidate_local(&mut res, k, row_addr, dist);
204                    }
205                }
206            }
207        };
208
209        // we don't need to sort the results by distances here
210        // because there's a SortExec node in the query plan which sorts the results from all partitions
211        let (row_ids, dists): (Vec<_>, Vec<_>) = res.into_iter().map(|r| (r.id, r.dist.0)).unzip();
212        let (row_ids, dists) = (UInt64Array::from(row_ids), Float32Array::from(dists));
213
214        Ok(RecordBatch::try_new(
215            ANN_SEARCH_SCHEMA.clone(),
216            vec![Arc::new(dists), Arc::new(row_ids)],
217        )?)
218    }
219
220    fn supports_global_topk_heap() -> bool {
221        true
222    }
223
224    fn accumulate_topk(
225        &self,
226        query: ArrayRef,
227        k: usize,
228        params: Self::QueryParams,
229        storage: &impl VectorStore,
230        prefilter: Arc<dyn PreFilter>,
231        res: &mut BinaryHeap<OrderedNode<u64>>,
232        metrics: &dyn MetricsCollector,
233    ) -> Result<()> {
234        let mut scratch = QueryScratch::new();
235        self.accumulate_topk_with_scratch(
236            query,
237            k,
238            params,
239            storage,
240            prefilter,
241            res,
242            None,
243            &mut scratch,
244            metrics,
245        )
246    }
247
248    fn accumulate_topk_with_scratch(
249        &self,
250        query: ArrayRef,
251        k: usize,
252        params: Self::QueryParams,
253        storage: &impl VectorStore,
254        prefilter: Arc<dyn PreFilter>,
255        res: &mut BinaryHeap<OrderedNode<u64>>,
256        residual: Option<QueryResidual<'_>>,
257        scratch: &mut QueryScratch,
258        metrics: &dyn MetricsCollector,
259    ) -> Result<()> {
260        let row_ids = storage.row_ids();
261        let dist_calc = storage.dist_calculator_with_scratch(
262            query,
263            params.dist_q_c,
264            residual,
265            &mut scratch.query_f32,
266            DistanceCalculatorOptions {
267                approx_mode: params.approx_mode,
268            },
269        );
270        metrics.record_comparisons(storage.len());
271
272        match prefilter.is_empty() {
273            true => {
274                dist_calc.accumulate_topk_with_scratch(
275                    k,
276                    params.lower_bound,
277                    params.upper_bound,
278                    |id| storage.row_id(id),
279                    res,
280                    &mut scratch.distances,
281                    &mut scratch.u16,
282                    &mut scratch.u8,
283                    &mut scratch.u32,
284                );
285            }
286            false => {
287                let row_addr_mask = prefilter.mask();
288                dist_calc.accumulate_filtered_topk_with_scratch(
289                    k,
290                    params.lower_bound,
291                    params.upper_bound,
292                    row_ids.enumerate().map(|(id, &row_id)| (id as u32, row_id)),
293                    |row_id| row_addr_mask.selected(row_id),
294                    res,
295                    &mut scratch.distances,
296                    &mut scratch.u16,
297                    &mut scratch.u8,
298                    &mut scratch.u32,
299                );
300            }
301        };
302        Ok(())
303    }
304
305    fn load(_: RecordBatch) -> Result<Self> {
306        Ok(Self {})
307    }
308
309    fn index_vectors(_: &impl VectorStore, _: Self::BuildParams) -> Result<Self>
310    where
311        Self: Sized,
312    {
313        Ok(Self {})
314    }
315
316    fn remap(&self, _: &RowAddrRemap, _: &impl VectorStore) -> Result<Self> {
317        Ok(self.clone())
318    }
319
320    fn to_batch(&self) -> Result<RecordBatch> {
321        Ok(RecordBatch::new_empty(Schema::empty().into()))
322    }
323}
324
325#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
326pub struct FlatMetadata {
327    pub dim: usize,
328}
329
330#[async_trait::async_trait]
331impl QuantizerMetadata for FlatMetadata {
332    async fn load(_: &PreviousFileReader) -> Result<Self> {
333        unimplemented!("Flat will be used in new index builder which doesn't require this")
334    }
335}
336
337#[derive(Debug, Clone, DeepSizeOf)]
338pub struct FlatQuantizer {
339    dim: usize,
340    distance_type: DistanceType,
341}
342
343impl FlatQuantizer {
344    pub fn new(dim: usize, distance_type: DistanceType) -> Self {
345        Self { dim, distance_type }
346    }
347}
348
349impl Quantization for FlatQuantizer {
350    type BuildParams = ();
351    type Metadata = FlatMetadata;
352    type Storage = FlatFloatStorage;
353
354    fn build(data: &dyn Array, distance_type: DistanceType, _: &Self::BuildParams) -> Result<Self> {
355        let dim = data.as_fixed_size_list().value_length();
356        Ok(Self::new(dim as usize, distance_type))
357    }
358
359    fn retrain(&mut self, _: &dyn Array) -> Result<()> {
360        Ok(())
361    }
362
363    fn code_dim(&self) -> usize {
364        self.dim
365    }
366
367    fn column(&self) -> &'static str {
368        FLAT_COLUMN
369    }
370
371    fn from_metadata(metadata: &Self::Metadata, distance_type: DistanceType) -> Result<Quantizer> {
372        Ok(Quantizer::Flat(Self {
373            dim: metadata.dim,
374            distance_type,
375        }))
376    }
377
378    fn metadata(&self, _: Option<crate::vector::quantizer::QuantizationMetadata>) -> FlatMetadata {
379        FlatMetadata { dim: self.dim }
380    }
381
382    fn metadata_key() -> &'static str {
383        "flat"
384    }
385
386    fn quantization_type() -> QuantizationType {
387        QuantizationType::Flat
388    }
389
390    fn quantize(&self, vectors: &dyn Array) -> Result<ArrayRef> {
391        Ok(vectors.slice(0, vectors.len()))
392    }
393
394    fn field(&self) -> Field {
395        Field::new(
396            FLAT_COLUMN,
397            DataType::FixedSizeList(
398                Arc::new(Field::new("item", DataType::Float32, true)),
399                self.dim as i32,
400            ),
401            true,
402        )
403    }
404}
405
406impl From<FlatQuantizer> for Quantizer {
407    fn from(value: FlatQuantizer) -> Self {
408        Self::Flat(value)
409    }
410}
411
412impl TryFrom<Quantizer> for FlatQuantizer {
413    type Error = Error;
414
415    fn try_from(value: Quantizer) -> Result<Self> {
416        match value {
417            Quantizer::Flat(quantizer) => Ok(quantizer),
418            _ => Err(Error::invalid_input("quantizer is not FlatQuantizer")),
419        }
420    }
421}
422
423#[derive(Debug, Clone, DeepSizeOf)]
424pub struct FlatBinQuantizer {
425    dim: usize,
426    distance_type: DistanceType,
427}
428
429impl FlatBinQuantizer {
430    pub fn new(dim: usize, distance_type: DistanceType) -> Self {
431        Self { dim, distance_type }
432    }
433}
434
435impl Quantization for FlatBinQuantizer {
436    type BuildParams = ();
437    type Metadata = FlatMetadata;
438    type Storage = FlatBinStorage;
439
440    fn build(data: &dyn Array, distance_type: DistanceType, _: &Self::BuildParams) -> Result<Self> {
441        let dim = data.as_fixed_size_list().value_length();
442        Ok(Self::new(dim as usize, distance_type))
443    }
444
445    fn retrain(&mut self, _: &dyn Array) -> Result<()> {
446        Ok(())
447    }
448
449    fn code_dim(&self) -> usize {
450        self.dim
451    }
452
453    fn column(&self) -> &'static str {
454        FLAT_COLUMN
455    }
456
457    fn from_metadata(metadata: &Self::Metadata, distance_type: DistanceType) -> Result<Quantizer> {
458        Ok(Quantizer::FlatBin(Self {
459            dim: metadata.dim,
460            distance_type,
461        }))
462    }
463
464    fn metadata(&self, _: Option<crate::vector::quantizer::QuantizationMetadata>) -> FlatMetadata {
465        FlatMetadata { dim: self.dim }
466    }
467
468    fn metadata_key() -> &'static str {
469        "flat"
470    }
471
472    fn quantization_type() -> QuantizationType {
473        QuantizationType::FlatBin
474    }
475
476    fn quantize(&self, vectors: &dyn Array) -> Result<ArrayRef> {
477        Ok(vectors.slice(0, vectors.len()))
478    }
479
480    fn field(&self) -> Field {
481        Field::new(
482            FLAT_COLUMN,
483            DataType::FixedSizeList(
484                Arc::new(Field::new("item", DataType::UInt8, true)),
485                self.dim as i32,
486            ),
487            true,
488        )
489    }
490}
491
492impl From<FlatBinQuantizer> for Quantizer {
493    fn from(value: FlatBinQuantizer) -> Self {
494        Self::FlatBin(value)
495    }
496}
497
498impl TryFrom<Quantizer> for FlatBinQuantizer {
499    type Error = Error;
500
501    fn try_from(value: Quantizer) -> Result<Self> {
502        match value {
503            Quantizer::FlatBin(quantizer) => Ok(quantizer),
504            _ => Err(Error::invalid_input("quantizer is not FlatBinQuantizer")),
505        }
506    }
507}
508
509#[cfg(test)]
510mod tests {
511    use super::*;
512
513    use arrow_array::FixedSizeListArray;
514    use async_trait::async_trait;
515    use lance_arrow::FixedSizeListArrayExt;
516    use lance_select::{RowAddrMask, RowAddrTreeMap};
517
518    use crate::metrics::NoOpMetricsCollector;
519    use crate::prefilter::NoFilter;
520
521    struct MaskPreFilter {
522        mask: Arc<RowAddrMask>,
523    }
524
525    #[async_trait]
526    impl PreFilter for MaskPreFilter {
527        async fn wait_for_ready(&self) -> Result<()> {
528            Ok(())
529        }
530
531        fn is_empty(&self) -> bool {
532            false
533        }
534
535        fn mask(&self) -> Arc<RowAddrMask> {
536            self.mask.clone()
537        }
538
539        fn filter_row_ids<'a>(&self, row_ids: Box<dyn Iterator<Item = &'a u64> + 'a>) -> Vec<u64> {
540            self.mask.selected_indices(row_ids)
541        }
542    }
543
544    fn test_storage() -> FlatFloatStorage {
545        let values = Float32Array::from(vec![
546            0.0, 0.0, // row 0
547            1.0, 0.0, // row 1
548            1.0, 1.0, // row 2
549            3.0, 3.0, // row 3
550            4.0, 4.0, // row 4
551        ]);
552        let vectors = FixedSizeListArray::try_new_from_values(values, 2).unwrap();
553        FlatFloatStorage::new(vectors, DistanceType::L2)
554    }
555
556    fn query() -> ArrayRef {
557        Arc::new(Float32Array::from(vec![1.0, 1.0]))
558    }
559
560    fn batch_results(batch: RecordBatch) -> Vec<(u64, f32)> {
561        let dists = batch
562            .column(0)
563            .as_primitive::<arrow_array::types::Float32Type>();
564        let row_ids = batch
565            .column(1)
566            .as_primitive::<arrow_array::types::UInt64Type>();
567        let mut results = row_ids
568            .values()
569            .iter()
570            .zip(dists.values().iter())
571            .map(|(row_id, dist)| (*row_id, *dist))
572            .collect::<Vec<_>>();
573        results.sort_by_key(|left| left.0);
574        results
575    }
576
577    fn heap_results(heap: BinaryHeap<OrderedNode<u64>>) -> Vec<(u64, f32)> {
578        let mut results = heap
579            .into_iter()
580            .map(|node| (node.id, node.dist.0))
581            .collect::<Vec<_>>();
582        results.sort_by_key(|left| left.0);
583        results
584    }
585
586    #[test]
587    fn test_flat_search_matches_accumulate_topk_without_prefilter() {
588        let index = FlatIndex::default();
589        let storage = test_storage();
590        let k = 3;
591        let search_results = batch_results(
592            index
593                .search(
594                    query(),
595                    k,
596                    FlatQueryParams::default(),
597                    &storage,
598                    Arc::new(NoFilter),
599                    &NoOpMetricsCollector,
600                )
601                .unwrap(),
602        );
603
604        let mut heap = BinaryHeap::with_capacity(k);
605        index
606            .accumulate_topk(
607                query(),
608                k,
609                FlatQueryParams::default(),
610                &storage,
611                Arc::new(NoFilter),
612                &mut heap,
613                &NoOpMetricsCollector,
614            )
615            .unwrap();
616
617        assert_eq!(search_results, heap_results(heap));
618    }
619
620    #[test]
621    fn test_flat_search_matches_accumulate_topk_with_prefilter() {
622        let index = FlatIndex::default();
623        let storage = test_storage();
624        let k = 2;
625        let filter = Arc::new(MaskPreFilter {
626            mask: Arc::new(RowAddrMask::from_allowed(RowAddrTreeMap::from_iter([
627                0_u64, 3, 4,
628            ]))),
629        });
630        let search_results = batch_results(
631            index
632                .search(
633                    query(),
634                    k,
635                    FlatQueryParams::default(),
636                    &storage,
637                    filter.clone(),
638                    &NoOpMetricsCollector,
639                )
640                .unwrap(),
641        );
642
643        let mut heap = BinaryHeap::with_capacity(k);
644        index
645            .accumulate_topk(
646                query(),
647                k,
648                FlatQueryParams::default(),
649                &storage,
650                filter,
651                &mut heap,
652                &NoOpMetricsCollector,
653            )
654            .unwrap();
655
656        assert_eq!(search_results, heap_results(heap));
657        assert_eq!(
658            search_results.iter().map(|(id, _)| *id).collect::<Vec<_>>(),
659            vec![0, 3]
660        );
661    }
662}