1use 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#[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 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, 1.0, 0.0, 1.0, 1.0, 3.0, 3.0, 4.0, 4.0, ]);
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}