1use arrow::array::{AsArray, ListBuilder, UInt32Builder};
7use arrow::compute::concat_batches;
8use arrow::datatypes::{DataType, UInt32Type};
9use arrow_array::{ArrayRef, Float32Array, ListArray, RecordBatch, UInt64Array};
10use crossbeam_queue::ArrayQueue;
11use itertools::Itertools;
12use lance_core::deepsize::DeepSizeOf;
13use lance_core::utils::row_addr_remap::RowAddrRemap;
14
15use lance_core::utils::tokio::get_num_compute_intensive_cpus;
16use lance_linalg::distance::DistanceType;
17use rayon::prelude::*;
18use std::cmp::min;
19use std::collections::{BinaryHeap, HashMap, VecDeque};
20use std::fmt::Debug;
21use std::iter;
22use std::sync::Arc;
23use std::sync::RwLock;
24use std::sync::atomic::{AtomicUsize, Ordering};
25use tracing::instrument;
26
27use lance_core::{Error, Result};
28use rand::{Rng, SeedableRng, rngs::SmallRng};
29use serde::{Deserialize, Serialize};
30
31use super::super::graph::beam_search;
32use super::{
33 HNSW_TYPE, HnswMetadata, VECTOR_ID_COL, VECTOR_ID_FIELD, select_neighbors_heuristic_owned,
34};
35use crate::metrics::MetricsCollector;
36use crate::prefilter::PreFilter;
37use crate::vector::flat::storage::{FlatBinStorage, FlatFloatStorage};
38use crate::vector::graph::builder::GraphBuilderNode;
39use crate::vector::graph::{
40 BorrowingGraph, DISTS_FIELD, Graph, NEIGHBORS_COL, NEIGHBORS_FIELD, OrderedFloat, OrderedNode,
41 VisitedGenerator,
42};
43use crate::vector::graph::{
44 Visited, beam_search_acorn, beam_search_borrowed, greedy_search, greedy_search_borrowed,
45};
46use crate::vector::storage::{DistCalculator, VectorStore};
47use crate::vector::v3::subindex::IvfSubIndex;
48use crate::vector::{ApproxMode, Query, VECTOR_RESULT_SCHEMA};
49
50pub const HNSW_METADATA_KEY: &str = "lance:hnsw";
51
52pub(crate) const HNSW_LEVEL_RNG_SEED: u64 = 42;
60
61#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
63pub struct HnswBuildParams {
64 pub max_level: u16,
66
67 pub m: usize,
69
70 pub ef_construction: usize,
72
73 pub prefetch_distance: Option<usize>,
75}
76
77impl From<&HnswBuildParams> for crate::pb::HnswParameters {
78 fn from(params: &HnswBuildParams) -> Self {
79 Self {
80 max_connections: params.m as u32,
81 construction_ef: params.ef_construction as u32,
82 max_level: params.max_level as u32,
83 }
84 }
85}
86
87impl Default for HnswBuildParams {
88 fn default() -> Self {
89 Self {
90 max_level: 7,
91 m: 20,
92 ef_construction: 150,
93 prefetch_distance: Some(2),
94 }
95 }
96}
97
98impl HnswBuildParams {
99 pub fn max_level(mut self, max_level: u16) -> Self {
102 self.max_level = max_level;
103 self
104 }
105
106 pub fn num_edges(mut self, m: usize) -> Self {
109 self.m = m;
110 self
111 }
112
113 pub fn ef_construction(mut self, ef_construction: usize) -> Self {
118 self.ef_construction = ef_construction;
119 self
120 }
121
122 pub async fn build(self, data: ArrayRef, distance_type: DistanceType) -> Result<HNSW> {
128 let vectors = data.as_fixed_size_list().clone();
129 match (vectors.value_type(), distance_type) {
130 (DataType::UInt8, DistanceType::Hamming) => {
131 let vec_store = Arc::new(FlatBinStorage::new(vectors, distance_type));
132 HNSW::index_vectors(vec_store.as_ref(), self)
133 }
134 (DataType::UInt8, _) => Err(Error::invalid_input(format!(
135 "HNSW only supports hamming distance for UInt8 vectors, got {}",
136 distance_type
137 ))),
138 (_, DistanceType::Hamming) => Err(Error::invalid_input(format!(
139 "HNSW hamming distance only supports UInt8 vectors, got {}",
140 vectors.value_type()
141 ))),
142 _ => {
143 let vec_store = Arc::new(FlatFloatStorage::new(vectors, distance_type));
144 HNSW::index_vectors(vec_store.as_ref(), self)
145 }
146 }
147 }
148}
149
150#[derive(Clone, DeepSizeOf)]
158pub struct HNSW {
159 inner: Arc<HnswCore>,
160}
161
162struct HnswCore {
163 params: HnswBuildParams,
164 graph: HnswGraph,
165 level_count: Vec<usize>,
166 entry_point: u32,
167 visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
168}
169
170impl DeepSizeOf for HnswCore {
171 fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
172 self.params.deep_size_of_children(context)
173 + self.graph.deep_size_of_children(context)
174 + self.level_count.deep_size_of_children(context)
175 }
177}
178
179impl HnswCore {
180 fn max_level(&self) -> u16 {
181 self.params.max_level
182 }
183
184 fn num_nodes(&self, level: usize) -> usize {
185 self.level_count[level]
186 }
187}
188
189impl Debug for HNSW {
190 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
191 write!(f, "HNSW(max_layers: {})", self.inner.max_level() as usize,)
192 }
193}
194
195impl HNSW {
196 pub(crate) fn from_parts(
199 params: HnswBuildParams,
200 nodes: Vec<GraphBuilderNode>,
201 level_count: Vec<usize>,
202 entry_point: u32,
203 ) -> Self {
204 let queue_size = get_num_compute_intensive_cpus().max(1) * 2;
205 let visited_generator_queue = Arc::new(ArrayQueue::new(queue_size));
206 for _ in 0..queue_size {
207 let _ = visited_generator_queue.push(VisitedGenerator::new(0));
208 }
209 Self {
210 inner: Arc::new(HnswCore {
211 params,
212 graph: HnswGraph::Built(Arc::new(nodes)),
213 level_count,
214 entry_point,
215 visited_generator_queue,
216 }),
217 }
218 }
219
220 pub fn empty() -> Self {
221 Self {
222 inner: Arc::new(HnswCore {
223 params: HnswBuildParams::default(),
224 graph: HnswGraph::Built(Arc::new(Vec::new())),
225 level_count: Vec::new(),
226 entry_point: 0,
227 visited_generator_queue: Arc::new(ArrayQueue::new(1)),
228 }),
229 }
230 }
231
232 pub fn len(&self) -> usize {
233 match &self.inner.graph {
234 HnswGraph::Built(nodes) => nodes.len(),
235 HnswGraph::Loaded(graph) => graph.level_count[0],
237 }
238 }
239
240 pub fn is_empty(&self) -> bool {
241 self.len() == 0
242 }
243
244 pub fn max_level(&self) -> u16 {
245 self.inner.max_level()
246 }
247
248 pub fn num_nodes(&self, level: usize) -> usize {
249 self.inner.num_nodes(level)
250 }
251
252 pub fn nodes(&self) -> Option<Arc<Vec<GraphBuilderNode>>> {
257 match &self.inner.graph {
258 HnswGraph::Built(nodes) => Some(nodes.clone()),
259 HnswGraph::Loaded(_) => None,
260 }
261 }
262
263 #[allow(clippy::too_many_arguments)]
264 pub fn search_inner(
265 &self,
266 query: ArrayRef,
267 k: usize,
268 params: &HnswQueryParams,
269 bitset: Option<Visited>,
270 visited_generator: &mut VisitedGenerator,
271 storage: &impl VectorStore,
272 prefetch_distance: Option<usize>,
273 ) -> Result<Vec<OrderedNode>> {
274 let dist_calc = storage.dist_calculator(query, params.dist_q_c);
275 let entry = self.inner.entry_point;
276 let ep = OrderedNode::new(entry, dist_calc.distance(entry).into());
277
278 let result = match &self.inner.graph {
284 HnswGraph::Built(nodes) => {
285 let nodes = nodes.as_slice();
286 self.run_search(
287 ep,
288 k,
289 params,
290 bitset.as_ref(),
291 visited_generator,
292 storage.len(),
293 prefetch_distance,
294 &dist_calc,
295 |level| ImmutableHnswLevelView::new(level, nodes),
296 ImmutableHnswBottomView::new(nodes),
297 )
298 }
299 HnswGraph::Loaded(graph) => {
300 let graph = graph.as_ref();
301 self.run_search(
302 ep,
303 k,
304 params,
305 bitset.as_ref(),
306 visited_generator,
307 storage.len(),
308 prefetch_distance,
309 &dist_calc,
310 |level| LoadedHnswLevelView::new(level, graph),
311 LoadedHnswBottomView::new(graph),
312 )
313 }
314 };
315 Ok(result)
316 }
317
318 #[allow(clippy::too_many_arguments)]
326 fn run_search<L, B>(
327 &self,
328 ep: OrderedNode,
329 k: usize,
330 params: &HnswQueryParams,
331 bitset: Option<&Visited>,
332 visited_generator: &mut VisitedGenerator,
333 storage_len: usize,
334 prefetch_distance: Option<usize>,
335 dist_calc: &impl DistCalculator,
336 make_level: impl Fn(u16) -> L,
337 bottom: B,
338 ) -> Vec<OrderedNode>
339 where
340 L: BorrowingGraph,
341 B: BorrowingGraph,
342 {
343 let mut ep = ep;
344 for level in (0..self.max_level()).rev() {
345 let cur_level = make_level(level);
346 ep = greedy_search_borrowed(
347 &cur_level,
348 ep,
349 dist_calc,
350 self.inner.params.prefetch_distance,
351 );
352 }
353 let mut visited = visited_generator.generate(storage_len);
354 beam_search_borrowed(
355 &bottom,
356 &ep,
357 params,
358 dist_calc,
359 bitset,
360 prefetch_distance,
361 &mut visited,
362 )
363 .into_iter()
364 .take(k)
365 .collect::<Vec<OrderedNode>>()
366 }
367
368 #[instrument(level = "debug", skip(self, query, bitset, storage))]
369 pub fn search_basic(
370 &self,
371 query: ArrayRef,
372 k: usize,
373 params: &HnswQueryParams,
374 bitset: Option<Visited>,
375 storage: &impl VectorStore,
376 ) -> Result<Vec<OrderedNode>> {
377 let mut visited_generator = self
378 .inner
379 .visited_generator_queue
380 .pop()
381 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
382 let result = self.search_inner(
383 query,
384 k,
385 params,
386 bitset,
387 &mut visited_generator,
388 storage,
389 Some(2),
390 );
391
392 match self.inner.visited_generator_queue.push(visited_generator) {
393 Ok(_) => {}
394 Err(_) => {
395 log::warn!("visited_generator_queue is full");
396 }
397 }
398
399 result
400 }
401
402 pub fn search_acorn(
405 &self,
406 query: ArrayRef,
407 k: usize,
408 params: &HnswQueryParams,
409 bitset: &Visited,
410 storage: &impl VectorStore,
411 ) -> Result<Vec<OrderedNode>> {
412 let mut visited_generator = self
413 .inner
414 .visited_generator_queue
415 .pop()
416 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
417 let mut expanded_generator = self
418 .inner
419 .visited_generator_queue
420 .pop()
421 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
422
423 let result = self.search_acorn_inner(
424 query,
425 k,
426 params,
427 bitset,
428 &mut visited_generator,
429 &mut expanded_generator,
430 storage,
431 Some(2),
432 );
433
434 let _ = self.inner.visited_generator_queue.push(visited_generator);
436 let _ = self.inner.visited_generator_queue.push(expanded_generator);
437 result
438 }
439
440 #[allow(clippy::too_many_arguments)]
441 fn search_acorn_inner(
442 &self,
443 query: ArrayRef,
444 k: usize,
445 params: &HnswQueryParams,
446 bitset: &Visited,
447 visited_generator: &mut VisitedGenerator,
448 expanded_generator: &mut VisitedGenerator,
449 storage: &impl VectorStore,
450 prefetch_distance: Option<usize>,
451 ) -> Result<Vec<OrderedNode>> {
452 let dist_calc = storage.dist_calculator(query, params.dist_q_c);
453 let entry = self.inner.entry_point;
454 let ep = OrderedNode::new(entry, dist_calc.distance(entry).into());
455
456 let result = match &self.inner.graph {
457 HnswGraph::Built(nodes) => {
458 let nodes = nodes.as_slice();
459 self.run_search_acorn(
460 ep,
461 params,
462 bitset,
463 visited_generator,
464 expanded_generator,
465 storage.len(),
466 prefetch_distance,
467 &dist_calc,
468 |level| ImmutableHnswLevelView::new(level, nodes),
469 ImmutableHnswBottomView::new(nodes),
470 )
471 }
472 HnswGraph::Loaded(graph) => {
473 let graph = graph.as_ref();
474 self.run_search_acorn(
475 ep,
476 params,
477 bitset,
478 visited_generator,
479 expanded_generator,
480 storage.len(),
481 prefetch_distance,
482 &dist_calc,
483 |level| LoadedHnswLevelView::new(level, graph),
484 LoadedHnswBottomView::new(graph),
485 )
486 }
487 };
488 Ok(result.into_iter().take(k).collect())
489 }
490
491 #[allow(clippy::too_many_arguments)]
494 fn run_search_acorn<L, B>(
495 &self,
496 ep: OrderedNode,
497 params: &HnswQueryParams,
498 bitset: &Visited,
499 visited_generator: &mut VisitedGenerator,
500 expanded_generator: &mut VisitedGenerator,
501 storage_len: usize,
502 prefetch_distance: Option<usize>,
503 dist_calc: &impl DistCalculator,
504 make_level: impl Fn(u16) -> L,
505 bottom: B,
506 ) -> Vec<OrderedNode>
507 where
508 L: BorrowingGraph,
509 B: BorrowingGraph,
510 {
511 let mut ep = ep;
512 for level in (0..self.max_level()).rev() {
513 let cur_level = make_level(level);
514 ep = greedy_search_borrowed(
515 &cur_level,
516 ep,
517 dist_calc,
518 self.inner.params.prefetch_distance,
519 );
520 }
521 let mut visited = visited_generator.generate(storage_len);
522 let mut expanded = expanded_generator.generate(storage_len);
523 beam_search_acorn(
524 &bottom,
525 &ep,
526 params,
527 dist_calc,
528 bitset,
529 prefetch_distance,
530 &mut visited,
531 &mut expanded,
532 )
533 }
534
535 #[instrument(level = "debug", skip(self, storage, query, prefilter_bitset))]
536 fn flat_search(
537 &self,
538 storage: &impl VectorStore,
539 query: ArrayRef,
540 k: usize,
541 prefilter_bitset: Visited,
542 params: &HnswQueryParams,
543 ) -> Vec<OrderedNode> {
544 let lower_bound: OrderedFloat = params.lower_bound.unwrap_or(f32::MIN).into();
545 let upper_bound: OrderedFloat = params.upper_bound.unwrap_or(f32::MAX).into();
546
547 let dist_calc = storage.dist_calculator(query, params.dist_q_c);
548 let mut heap = BinaryHeap::<OrderedNode>::with_capacity(k);
549
550 match self.inner.params.prefetch_distance {
551 Some(ahead) if ahead > 0 => {
552 let mut ids_iter = prefilter_bitset.iter_ones().map(|i| i as u32);
553 let mut buffer = VecDeque::with_capacity(ahead + 1);
554 for _ in 0..=ahead {
555 if let Some(id) = ids_iter.next() {
556 buffer.push_back(id);
557 } else {
558 break;
559 }
560 }
561
562 while let Some(node_id) = buffer.pop_front() {
563 if let Some(&prefetch_id) = buffer.get(ahead - 1) {
564 dist_calc.prefetch(prefetch_id);
565 }
566 if let Some(next) = ids_iter.next() {
567 buffer.push_back(next);
568 }
569
570 let dist: OrderedFloat = dist_calc.distance(node_id).into();
571 if dist <= lower_bound || dist > upper_bound {
572 continue;
573 }
574 if heap.len() < k {
575 heap.push((dist, node_id).into());
576 } else if dist < heap.peek().unwrap().dist {
577 heap.pop();
578 heap.push((dist, node_id).into());
579 }
580 }
581 }
582 _ => {
583 for node_id in prefilter_bitset.iter_ones().map(|i| i as u32) {
584 let dist: OrderedFloat = dist_calc.distance(node_id).into();
585 if dist <= lower_bound || dist > upper_bound {
586 continue;
587 }
588 if heap.len() < k {
589 heap.push((dist, node_id).into());
590 } else if dist < heap.peek().unwrap().dist {
591 heap.pop();
592 heap.push((dist, node_id).into());
593 }
594 }
595 }
596 };
597 heap.into_sorted_vec()
598 }
599
600 pub fn metadata(&self) -> HnswMetadata {
602 let level_offsets = self
605 .inner
606 .level_count
607 .iter()
608 .chain(iter::once(&0))
609 .scan(0, |state, x| {
610 let start = *state;
611 *state += *x;
612 Some(start)
613 })
614 .collect();
615
616 HnswMetadata {
617 entry_point: self.inner.entry_point,
618 params: self.inner.params.clone(),
619 level_offsets,
620 }
621 }
622}
623
624struct HnswBuilder {
625 params: HnswBuildParams,
626
627 nodes: Arc<Vec<RwLock<GraphBuilderNode>>>,
628 level_count: Vec<AtomicUsize>,
629
630 entry_point: u32,
631
632 visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
633}
634
635impl DeepSizeOf for HnswBuilder {
636 fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
637 self.params.deep_size_of_children(context)
638 + self.nodes.deep_size_of_children(context)
639 + self.level_count.deep_size_of_children(context)
640 }
642}
643
644impl HnswBuilder {
645 fn finish(self) -> HNSW {
646 let nodes = match Arc::try_unwrap(self.nodes) {
647 Ok(nodes) => nodes
648 .into_iter()
649 .map(|node| node.into_inner().expect("builder lock poisoned"))
650 .collect(),
651 Err(nodes) => nodes
652 .iter()
653 .map(|node| node.read().expect("builder lock poisoned").clone())
654 .collect(),
655 };
656
657 let level_count = self
658 .level_count
659 .into_iter()
660 .map(|count| count.load(Ordering::Relaxed))
661 .collect();
662
663 HNSW {
664 inner: Arc::new(HnswCore {
665 params: self.params,
666 graph: HnswGraph::Built(Arc::new(nodes)),
667 level_count,
668 entry_point: self.entry_point,
669 visited_generator_queue: self.visited_generator_queue,
670 }),
671 }
672 }
673
674 pub fn with_params(params: HnswBuildParams, storage: &impl VectorStore) -> Self {
676 let len = storage.len();
677 let max_level = params.max_level;
678
679 let level_count = (0..max_level)
680 .map(|_| AtomicUsize::new(0))
681 .collect::<Vec<_>>();
682
683 let visited_generator_queue = Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus()));
684 for _ in 0..get_num_compute_intensive_cpus() {
685 visited_generator_queue
686 .push(VisitedGenerator::new(0))
687 .unwrap();
688 }
689 let mut builder = Self {
690 params,
691 nodes: Arc::new(Vec::new()),
692 level_count,
693 entry_point: 0,
694 visited_generator_queue,
695 };
696
697 if storage.is_empty() {
698 return builder;
699 }
700
701 let mut nodes = Vec::with_capacity(len);
702 {
703 if len > 0 {
704 nodes.push(RwLock::new(GraphBuilderNode::new(0, max_level as usize)));
705 }
706 let mut level_rng = SmallRng::seed_from_u64(HNSW_LEVEL_RNG_SEED);
707 for i in 1..len {
708 nodes.push(RwLock::new(GraphBuilderNode::new(
709 i as u32,
710 builder.random_level(&mut level_rng) as usize + 1,
711 )));
712 }
713 }
714 builder.nodes = Arc::new(nodes);
715
716 builder
717 }
718
719 fn random_level<R: Rng + ?Sized>(&self, rng: &mut R) -> u16 {
723 let ml = 1.0 / (self.params.m as f32).ln();
724 min(
725 (-rng.random::<f32>().ln() * ml) as u16,
726 self.params.max_level - 1,
727 )
728 }
729
730 fn insert(
732 &self,
733 node: u32,
734 visited_generator: &mut VisitedGenerator,
735 storage: &impl VectorStore,
736 ) {
737 let nodes = &self.nodes;
738 let target_level = nodes[node as usize].read().unwrap().level_neighbors.len() as u16 - 1;
739 let dist_calc = storage.dist_calculator_from_id(node);
740 let mut ep = OrderedNode::new(
741 self.entry_point,
742 dist_calc.distance(self.entry_point).into(),
743 );
744
745 for level in (target_level + 1..self.params.max_level).rev() {
754 let cur_level = HnswLevelView::new(level, nodes);
755 ep = greedy_search(&cur_level, ep, &dist_calc, self.params.prefetch_distance);
756 }
757
758 let mut pruned_neighbors_per_level: Vec<Vec<_>> =
759 vec![Vec::new(); (target_level + 1) as usize];
760 {
761 let mut current_node = nodes[node as usize].write().unwrap();
762 for level in (0..=target_level).rev() {
763 self.level_count[level as usize].fetch_add(1, Ordering::Relaxed);
764
765 let neighbors = self.search_level(&ep, level, &dist_calc, nodes, visited_generator);
766 for neighbor in &neighbors {
767 current_node.add_neighbor(neighbor.id, neighbor.dist, level);
768 }
769 self.prune(storage, &mut current_node, level);
770 pruned_neighbors_per_level[level as usize]
771 .clone_from(¤t_node.level_neighbors_ranked[level as usize]);
772
773 ep = neighbors[0].clone();
774 }
775 }
776 for (level, pruned_neighbors) in pruned_neighbors_per_level.iter().enumerate() {
777 for unpruned_edge in pruned_neighbors {
778 let level = level as u16;
779 let m_max = match level {
780 0 => self.params.m * 2,
781 _ => self.params.m,
782 };
783 if unpruned_edge.dist
784 < nodes[unpruned_edge.id as usize]
785 .read()
786 .unwrap()
787 .cutoff(level, m_max)
788 {
789 let mut chosen_node = nodes[unpruned_edge.id as usize].write().unwrap();
790 chosen_node.add_neighbor(node, unpruned_edge.dist, level);
791 self.prune(storage, &mut chosen_node, level);
792 }
793 }
794 }
795 }
796
797 fn search_level(
798 &self,
799 ep: &OrderedNode,
800 level: u16,
801 dist_calc: &impl DistCalculator,
802 nodes: &[RwLock<GraphBuilderNode>],
803 visited_generator: &mut VisitedGenerator,
804 ) -> Vec<OrderedNode> {
805 let cur_level = HnswLevelView::new(level, nodes);
806 let mut visited = visited_generator.generate(nodes.len());
807 beam_search(
808 &cur_level,
809 ep,
810 &HnswQueryParams {
811 ef: self.params.ef_construction,
812 lower_bound: None,
813 upper_bound: None,
814 dist_q_c: 0.0,
815 use_acorn: false,
816 },
817 dist_calc,
818 None,
819 self.params.prefetch_distance,
820 &mut visited,
821 )
822 }
823
824 fn prune(&self, storage: &impl VectorStore, builder_node: &mut GraphBuilderNode, level: u16) {
825 let m_max = match level {
826 0 => self.params.m * 2,
827 _ => self.params.m,
828 };
829
830 let neighbors_ranked = &mut builder_node.level_neighbors_ranked[level as usize];
831 if neighbors_ranked.len() <= m_max {
832 builder_node.update_from_ranked_neighbors(level);
833 return;
834 }
835
836 let level_neighbors = std::mem::take(neighbors_ranked);
837 *neighbors_ranked = select_neighbors_heuristic_owned(storage, level_neighbors, m_max);
838 builder_node.update_from_ranked_neighbors(level);
839 }
840}
841
842pub(crate) struct HnswLevelView<'a> {
845 level: u16,
846 nodes: &'a [RwLock<GraphBuilderNode>],
847}
848
849impl<'a> HnswLevelView<'a> {
850 pub fn new(level: u16, nodes: &'a [RwLock<GraphBuilderNode>]) -> Self {
851 Self { level, nodes }
852 }
853}
854
855impl Graph for HnswLevelView<'_> {
856 fn len(&self) -> usize {
857 self.nodes.len()
858 }
859
860 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
861 let node = &self.nodes[key as usize];
862 node.read().unwrap().level_neighbors[self.level as usize].clone()
863 }
864}
865
866pub(crate) struct ImmutableHnswLevelView<'a> {
867 level: u16,
868 nodes: &'a [GraphBuilderNode],
869}
870
871impl<'a> ImmutableHnswLevelView<'a> {
872 pub fn new(level: u16, nodes: &'a [GraphBuilderNode]) -> Self {
873 Self { level, nodes }
874 }
875}
876
877impl Graph for ImmutableHnswLevelView<'_> {
878 fn len(&self) -> usize {
879 self.nodes.len()
880 }
881
882 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
883 self.nodes[key as usize].level_neighbors[self.level as usize].clone()
884 }
885}
886
887impl BorrowingGraph for ImmutableHnswLevelView<'_> {
888 fn len(&self) -> usize {
889 self.nodes.len()
890 }
891
892 fn neighbors(&self, key: u32) -> &[u32] {
893 self.nodes[key as usize].level_neighbors[self.level as usize].as_slice()
894 }
895}
896
897pub(crate) struct ImmutableHnswBottomView<'a> {
898 nodes: &'a [GraphBuilderNode],
899}
900
901impl<'a> ImmutableHnswBottomView<'a> {
902 pub fn new(nodes: &'a [GraphBuilderNode]) -> Self {
903 Self { nodes }
904 }
905}
906
907impl Graph for ImmutableHnswBottomView<'_> {
908 fn len(&self) -> usize {
909 self.nodes.len()
910 }
911
912 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
913 self.nodes[key as usize].bottom_neighbors.clone()
914 }
915}
916
917impl BorrowingGraph for ImmutableHnswBottomView<'_> {
918 fn len(&self) -> usize {
919 self.nodes.len()
920 }
921
922 fn neighbors(&self, key: u32) -> &[u32] {
923 self.nodes[key as usize].bottom_neighbors.as_slice()
924 }
925}
926
927enum LevelLookup {
929 Dense,
933 Sparse(HashMap<u32, u32>),
944}
945
946struct LoadedHnswGraph {
958 batch: RecordBatch,
961 level_neighbors: Vec<ListArray>,
963 level_lookup: Vec<LevelLookup>,
965 level_count: Vec<usize>,
967}
968
969impl DeepSizeOf for LoadedHnswGraph {
970 fn deep_size_of_children(&self, _context: &mut lance_core::deepsize::Context) -> usize {
971 self.batch.get_array_memory_size()
977 }
978}
979
980impl LoadedHnswGraph {
981 #[inline]
984 fn neighbors_at(&self, level: usize, key: u32) -> &[u32] {
985 let row = match &self.level_lookup[level] {
986 LevelLookup::Dense => key as usize,
987 LevelLookup::Sparse(id_to_row) => match id_to_row.get(&key) {
988 Some(&row) => row as usize,
989 None => return &[],
995 },
996 };
997 let list = &self.level_neighbors[level];
998 let offsets = list.value_offsets();
999 let start = offsets[row] as usize;
1000 let end = offsets[row + 1] as usize;
1001 let values = list.values().as_primitive::<UInt32Type>();
1005 &values.values()[start..end]
1006 }
1007}
1008
1009pub(crate) struct LoadedHnswLevelView<'a> {
1011 level: usize,
1012 graph: &'a LoadedHnswGraph,
1013}
1014
1015impl<'a> LoadedHnswLevelView<'a> {
1016 fn new(level: u16, graph: &'a LoadedHnswGraph) -> Self {
1017 Self {
1018 level: level as usize,
1019 graph,
1020 }
1021 }
1022}
1023
1024impl Graph for LoadedHnswLevelView<'_> {
1025 fn len(&self) -> usize {
1026 self.graph.level_count[0]
1028 }
1029
1030 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
1031 Arc::new(self.graph.neighbors_at(self.level, key).to_vec())
1035 }
1036}
1037
1038impl BorrowingGraph for LoadedHnswLevelView<'_> {
1039 fn len(&self) -> usize {
1040 self.graph.level_count[0]
1041 }
1042
1043 fn neighbors(&self, key: u32) -> &[u32] {
1044 self.graph.neighbors_at(self.level, key)
1045 }
1046}
1047
1048pub(crate) struct LoadedHnswBottomView<'a> {
1050 graph: &'a LoadedHnswGraph,
1051}
1052
1053impl<'a> LoadedHnswBottomView<'a> {
1054 fn new(graph: &'a LoadedHnswGraph) -> Self {
1055 Self { graph }
1056 }
1057}
1058
1059impl Graph for LoadedHnswBottomView<'_> {
1060 fn len(&self) -> usize {
1061 self.graph.level_count[0]
1062 }
1063
1064 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
1065 Arc::new(self.graph.neighbors_at(0, key).to_vec())
1066 }
1067}
1068
1069impl BorrowingGraph for LoadedHnswBottomView<'_> {
1070 fn len(&self) -> usize {
1071 self.graph.level_count[0]
1072 }
1073
1074 fn neighbors(&self, key: u32) -> &[u32] {
1075 self.graph.neighbors_at(0, key)
1076 }
1077}
1078
1079enum HnswGraph {
1081 Built(Arc<Vec<GraphBuilderNode>>),
1085 Loaded(Arc<LoadedHnswGraph>),
1087}
1088
1089impl DeepSizeOf for HnswGraph {
1090 fn deep_size_of_children(&self, context: &mut lance_core::deepsize::Context) -> usize {
1091 match self {
1092 Self::Built(nodes) => nodes.deep_size_of_children(context),
1093 Self::Loaded(graph) => graph.deep_size_of_children(context),
1094 }
1095 }
1096}
1097
1098#[derive(Debug, Clone, Copy)]
1099pub struct HnswQueryParams {
1100 pub ef: usize,
1101 pub lower_bound: Option<f32>,
1102 pub upper_bound: Option<f32>,
1103 pub dist_q_c: f32,
1104 pub use_acorn: bool,
1105}
1106
1107impl From<&Query> for HnswQueryParams {
1108 fn from(query: &Query) -> Self {
1109 let k = query.k * query.refine_factor.unwrap_or(1) as usize;
1110 Self {
1111 ef: query.ef.unwrap_or(k + k / 2),
1112 lower_bound: query.lower_bound,
1113 upper_bound: query.upper_bound,
1114 dist_q_c: query.dist_q_c,
1115 use_acorn: query.approx_mode == ApproxMode::Fast,
1116 }
1117 }
1118}
1119
1120impl IvfSubIndex for HNSW {
1121 type BuildParams = HnswBuildParams;
1122 type QueryParams = HnswQueryParams;
1123
1124 fn load(data: RecordBatch) -> Result<Self>
1125 where
1126 Self: Sized,
1127 {
1128 if data.num_rows() == 0 {
1129 return Ok(Self::empty());
1130 }
1131
1132 let hnsw_metadata = data
1133 .schema_ref()
1134 .metadata()
1135 .get(HNSW_METADATA_KEY)
1136 .ok_or(Error::index(format!("{} not found", HNSW_METADATA_KEY)))?;
1137 let hnsw_metadata: HnswMetadata = serde_json::from_str(hnsw_metadata).map_err(|e| {
1138 Error::index(format!(
1139 "Failed to decode HNSW metadata: {}, json: {}",
1140 e, hnsw_metadata
1141 ))
1142 })?;
1143
1144 let level_batches: Vec<RecordBatch> = hnsw_metadata
1146 .level_offsets
1147 .iter()
1148 .tuple_windows()
1149 .map(|(start, end)| data.slice(*start, end - start))
1150 .collect();
1151
1152 let level_count = level_batches
1153 .iter()
1154 .map(|b| b.num_rows())
1155 .collect::<Vec<_>>();
1156
1157 let mut level_neighbors = Vec::with_capacity(level_batches.len());
1162 let mut level_lookup = Vec::with_capacity(level_batches.len());
1163 for (level, batch) in level_batches.iter().enumerate() {
1164 let neighbors = batch[NEIGHBORS_COL].as_list::<i32>().clone();
1167 let ids = batch[VECTOR_ID_COL].as_primitive::<UInt32Type>();
1168 if level == 0 {
1169 if let Some((row, id)) = ids
1178 .values()
1179 .iter()
1180 .enumerate()
1181 .find(|&(row, id)| *id != row as u32)
1182 {
1183 return Err(Error::index(format!(
1184 "HNSW level-0 __vector_id must equal the row index, but \
1185 row {row} has __vector_id {id}; the on-disk batch is \
1186 malformed or was written by an incompatible version"
1187 )));
1188 }
1189 level_lookup.push(LevelLookup::Dense);
1190 } else {
1191 let id_to_row: HashMap<u32, u32> = ids
1197 .values()
1198 .iter()
1199 .enumerate()
1200 .map(|(row, id)| (*id, row as u32))
1201 .collect();
1202 level_lookup.push(LevelLookup::Sparse(id_to_row));
1203 }
1204 level_neighbors.push(neighbors);
1205 }
1206
1207 let num_nodes = level_count[0];
1212 if hnsw_metadata.entry_point as usize >= num_nodes {
1213 return Err(Error::index(format!(
1214 "HNSW entry_point {} is out of range for a graph with {num_nodes} \
1215 nodes; the on-disk batch is malformed or was written by an \
1216 incompatible version",
1217 hnsw_metadata.entry_point
1218 )));
1219 }
1220
1221 let visited_generator_queue =
1222 Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus() * 2));
1223 for _ in 0..get_num_compute_intensive_cpus() * 2 {
1224 visited_generator_queue
1225 .push(VisitedGenerator::new(0))
1226 .unwrap();
1227 }
1228
1229 let graph = LoadedHnswGraph {
1230 batch: data,
1231 level_neighbors,
1232 level_lookup,
1233 level_count: level_count.clone(),
1234 };
1235 let inner = HnswCore {
1236 params: hnsw_metadata.params,
1237 graph: HnswGraph::Loaded(Arc::new(graph)),
1238 level_count,
1239 entry_point: hnsw_metadata.entry_point,
1240 visited_generator_queue,
1241 };
1242
1243 Ok(Self {
1244 inner: Arc::new(inner),
1245 })
1246 }
1247
1248 fn name() -> &'static str {
1249 HNSW_TYPE
1250 }
1251
1252 fn metadata_key() -> &'static str {
1253 "lance:hnsw"
1254 }
1255
1256 fn schema() -> arrow_schema::SchemaRef {
1258 arrow_schema::Schema::new(vec![
1259 VECTOR_ID_FIELD.clone(),
1260 NEIGHBORS_FIELD.clone(),
1261 DISTS_FIELD.clone(),
1262 ])
1263 .into()
1264 }
1265
1266 #[instrument(level = "debug", skip(self, query, storage, prefilter, _metrics))]
1267 fn search(
1268 &self,
1269 query: ArrayRef,
1270 k: usize,
1271 params: Self::QueryParams,
1272 storage: &impl VectorStore,
1273 prefilter: Arc<dyn PreFilter>,
1274 _metrics: &dyn MetricsCollector,
1275 ) -> Result<RecordBatch> {
1276 if params.ef < k {
1277 return Err(Error::index(
1278 "ef must be greater than or equal to k".to_string(),
1279 ));
1280 }
1281
1282 let schema = VECTOR_RESULT_SCHEMA.clone();
1283 if self.is_empty() {
1284 return Ok(RecordBatch::new_empty(schema));
1285 }
1286
1287 let mut prefilter_generator = self
1288 .inner
1289 .visited_generator_queue
1290 .pop()
1291 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
1292 let results = if prefilter.is_empty() {
1293 self.search_basic(query, k, ¶ms, None, storage)?
1294 } else {
1295 let indices = prefilter.filter_row_ids(Box::new(storage.row_ids()));
1298 let mut prefilter_bitset = prefilter_generator.generate(storage.len());
1299 for index in indices {
1300 prefilter_bitset.insert(index as u32);
1301 }
1302 let remained = prefilter_bitset.count_ones();
1303 if remained == storage.len() {
1304 drop(prefilter_bitset);
1306 self.search_basic(query, k, ¶ms, None, storage)?
1307 } else if remained < self.len() * 10 / 100 {
1308 self.flat_search(storage, query, k, prefilter_bitset, ¶ms)
1310 } else if params.use_acorn {
1311 let acorn_results =
1312 self.search_acorn(query.clone(), k, ¶ms, &prefilter_bitset, storage)?;
1313 let bounded = params.lower_bound.is_some() || params.upper_bound.is_some();
1317 if !bounded && acorn_results.len() < k.min(remained) {
1318 self.search_basic(query, k, ¶ms, Some(prefilter_bitset), storage)?
1319 } else {
1320 drop(prefilter_bitset);
1321 acorn_results
1322 }
1323 } else {
1324 self.search_basic(query, k, ¶ms, Some(prefilter_bitset), storage)?
1325 }
1326 };
1327 let _ = self.inner.visited_generator_queue.push(prefilter_generator);
1329
1330 let (row_ids, dists): (Vec<_>, Vec<_>) = results
1332 .into_iter()
1333 .map(|r| (storage.row_id(r.id), r.dist.0))
1334 .unique_by(|r| r.0)
1335 .unzip();
1336 let row_ids = Arc::new(UInt64Array::from(row_ids));
1337 let distances = Arc::new(Float32Array::from(dists));
1338
1339 Ok(RecordBatch::try_new(schema, vec![distances, row_ids])?)
1340 }
1341
1342 fn index_vectors(storage: &impl VectorStore, params: Self::BuildParams) -> Result<Self>
1344 where
1345 Self: Sized,
1346 {
1347 let builder = HnswBuilder::with_params(params, storage);
1348
1349 log::debug!(
1350 "Building HNSW graph: num={}, max_levels={}, m={}, ef_construction={}, distance_type:{}",
1351 storage.len(),
1352 builder.params.max_level,
1353 builder.params.m,
1354 builder.params.ef_construction,
1355 storage.distance_type(),
1356 );
1357
1358 if storage.is_empty() {
1359 return Ok(builder.finish());
1360 }
1361
1362 let len = storage.len();
1363 builder.level_count[0].fetch_add(1, Ordering::Relaxed);
1364 (1..len).into_par_iter().for_each_init(
1365 || VisitedGenerator::new(len),
1366 |visited_generator, node| {
1367 builder.insert(node as u32, visited_generator, storage);
1368 },
1369 );
1370
1371 assert_eq!(builder.level_count[0].load(Ordering::Relaxed), len);
1372 Ok(builder.finish())
1373 }
1374
1375 fn remap(
1376 &self,
1377 _mapping: &RowAddrRemap, store: &impl VectorStore,
1379 ) -> Result<Self> {
1380 Self::index_vectors(store, self.inner.params.clone())
1383 }
1384
1385 fn to_batch(&self) -> Result<RecordBatch> {
1387 let nodes = match &self.inner.graph {
1388 HnswGraph::Built(nodes) => nodes,
1389 HnswGraph::Loaded(graph) => {
1390 let metadata = serde_json::to_string(&self.metadata())?;
1403 let mut schema_metadata = graph.batch.schema_ref().metadata().clone();
1404 schema_metadata.insert(HNSW_METADATA_KEY.to_string(), metadata);
1405 let schema = graph
1406 .batch
1407 .schema()
1408 .as_ref()
1409 .clone()
1410 .with_metadata(schema_metadata);
1411 return Ok(graph.batch.clone().with_schema(Arc::new(schema))?);
1412 }
1413 };
1414
1415 let mut vector_id_builder = UInt32Builder::with_capacity(self.len());
1416 let mut neighbors_builder = ListBuilder::with_capacity(UInt32Builder::new(), self.len());
1417 let mut distances_builder =
1418 ListBuilder::with_capacity(arrow_array::builder::Float32Builder::new(), self.len());
1419 let mut batches = Vec::with_capacity(self.max_level() as usize);
1420 for level in 0..self.max_level() {
1421 let level = level as usize;
1422 for (id, node) in nodes.iter().enumerate() {
1423 if level >= node.level_neighbors.len() {
1424 continue;
1425 }
1426 let neighbors = node.level_neighbors[level].iter().map(|n| Some(*n));
1427 let distances = node.level_neighbors_ranked[level]
1428 .iter()
1429 .map(|n| Some(n.dist.0));
1430 vector_id_builder.append_value(id as u32);
1431 neighbors_builder.append_value(neighbors);
1432 distances_builder.append_value(distances);
1433 }
1434
1435 let batch = RecordBatch::try_new(
1436 Self::schema(),
1437 vec![
1438 Arc::new(vector_id_builder.finish()),
1439 Arc::new(neighbors_builder.finish()),
1440 Arc::new(distances_builder.finish()),
1441 ],
1442 )?;
1443 batches.push(batch);
1444 }
1445
1446 let metadata = self.metadata();
1447 let metadata = serde_json::to_string(&metadata)?;
1448 let schema = Self::schema()
1449 .as_ref()
1450 .clone()
1451 .with_metadata(HashMap::from_iter(vec![(
1452 HNSW_METADATA_KEY.to_string(),
1453 metadata,
1454 )]));
1455 let batch = concat_batches(&Self::schema(), batches.iter())?;
1456 let batch = batch.with_schema(Arc::new(schema))?;
1457 Ok(batch)
1458 }
1459}
1460
1461#[cfg(test)]
1462mod tests {
1463 use std::sync::Arc;
1464
1465 use arrow_array::{ArrayRef, FixedSizeListArray, RecordBatch, UInt8Array, UInt32Array};
1466 use arrow_schema::Schema;
1467 use lance_arrow::FixedSizeListArrayExt;
1468 use lance_core::deepsize::DeepSizeOf;
1469 use lance_file::versions::v1::{
1470 reader::FileReader as V1FileReader,
1471 writer::{FileWriter as V1FileWriter, FileWriterOptions as V1FileWriterOptions},
1472 };
1473 use lance_io::object_store::ObjectStore;
1474 use lance_linalg::distance::DistanceType;
1475 use lance_table::format::SelfDescribingFileReader;
1476 use lance_table::io::manifest::ManifestDescribing;
1477 use lance_testing::datagen::generate_random_array;
1478 use object_store::path::Path;
1479 use rstest::rstest;
1480
1481 use super::HnswGraph;
1482 use crate::vector::storage::{DistCalculator, VectorStore};
1483 use crate::vector::v3::subindex::IvfSubIndex;
1484 use crate::vector::{
1485 flat::storage::{FlatBinStorage, FlatFloatStorage},
1486 graph::{DISTS_FIELD, NEIGHBORS_FIELD, VisitedGenerator},
1487 hnsw::{
1488 HNSW, VECTOR_ID_FIELD,
1489 builder::{HnswBuildParams, HnswQueryParams},
1490 },
1491 };
1492
1493 #[tokio::test]
1494 async fn test_builder_write_load() {
1495 const DIM: usize = 32;
1496 const TOTAL: usize = 2048;
1497 const NUM_EDGES: usize = 20;
1498 let data = generate_random_array(TOTAL * DIM);
1499 let fsl = FixedSizeListArray::try_new_from_values(data, DIM as i32).unwrap();
1500 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1501 let builder = HNSW::index_vectors(
1502 store.as_ref(),
1503 HnswBuildParams::default()
1504 .num_edges(NUM_EDGES)
1505 .ef_construction(50),
1506 )
1507 .unwrap();
1508
1509 let object_store = ObjectStore::memory();
1510 let path = Path::from("test_builder_write_load");
1511 let writer = object_store.create(&path).await.unwrap();
1512 let schema = Schema::new(vec![
1513 VECTOR_ID_FIELD.clone(),
1514 NEIGHBORS_FIELD.clone(),
1515 DISTS_FIELD.clone(),
1516 ]);
1517 let schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
1518 let mut writer = V1FileWriter::<ManifestDescribing>::with_object_writer(
1519 writer,
1520 schema,
1521 &V1FileWriterOptions::default(),
1522 )
1523 .unwrap();
1524 let batch = builder.to_batch().unwrap();
1525 let metadata = batch.schema_ref().metadata().clone();
1526 writer.write(&[batch]).await.unwrap();
1527 writer.finish_with_metadata(&metadata).await.unwrap();
1528
1529 let reader = V1FileReader::try_new_self_described(&object_store, &path, None)
1530 .await
1531 .unwrap();
1532 let batch = reader
1533 .read_range(0..reader.len(), reader.schema())
1534 .await
1535 .unwrap();
1536 let loaded_hnsw = HNSW::load(batch).unwrap();
1537
1538 let query = fsl.value(0);
1539 let k = 10;
1540 let params = HnswQueryParams {
1541 ef: 50,
1542 lower_bound: None,
1543 upper_bound: None,
1544 dist_q_c: 0.0,
1545 use_acorn: false,
1546 };
1547 let builder_results = builder
1548 .search_basic(query.clone(), k, ¶ms, None, store.as_ref())
1549 .unwrap();
1550 let loaded_results = loaded_hnsw
1551 .search_basic(query, k, ¶ms, None, store.as_ref())
1552 .unwrap();
1553 assert_eq!(builder_results, loaded_results);
1554 }
1555
1556 #[tokio::test]
1557 async fn test_builder_write_load_binary_hamming() {
1558 const DIM: usize = 8;
1559 const TOTAL: usize = 256;
1560 const NUM_EDGES: usize = 20;
1561 let data = UInt8Array::from_iter_values((0..TOTAL * DIM).map(|v| (v % 16) as u8));
1562 let fsl = FixedSizeListArray::try_new_from_values(data, DIM as i32).unwrap();
1563 let store = Arc::new(FlatBinStorage::new(fsl.clone(), DistanceType::Hamming));
1564 let builder = HnswBuildParams::default()
1565 .num_edges(NUM_EDGES)
1566 .ef_construction(50)
1567 .build(Arc::new(fsl.clone()), DistanceType::Hamming)
1568 .await
1569 .unwrap();
1570
1571 let object_store = ObjectStore::memory();
1572 let path = Path::from("test_builder_write_load_binary_hamming");
1573 let writer = object_store.create(&path).await.unwrap();
1574 let schema = Schema::new(vec![
1575 VECTOR_ID_FIELD.clone(),
1576 NEIGHBORS_FIELD.clone(),
1577 DISTS_FIELD.clone(),
1578 ]);
1579 let schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
1580 let mut writer = V1FileWriter::<ManifestDescribing>::with_object_writer(
1581 writer,
1582 schema,
1583 &V1FileWriterOptions::default(),
1584 )
1585 .unwrap();
1586 let batch = builder.to_batch().unwrap();
1587 let metadata = batch.schema_ref().metadata().clone();
1588 writer.write(&[batch]).await.unwrap();
1589 writer.finish_with_metadata(&metadata).await.unwrap();
1590
1591 let reader = V1FileReader::try_new_self_described(&object_store, &path, None)
1592 .await
1593 .unwrap();
1594 let batch = reader
1595 .read_range(0..reader.len(), reader.schema())
1596 .await
1597 .unwrap();
1598 let loaded_hnsw = HNSW::load(batch).unwrap();
1599
1600 let query = fsl.value(0);
1601 let k = 10;
1602 let params = HnswQueryParams {
1603 ef: 50,
1604 lower_bound: None,
1605 upper_bound: None,
1606 dist_q_c: 0.0,
1607 use_acorn: false,
1608 };
1609 let builder_results = builder
1610 .search_basic(query.clone(), k, ¶ms, None, store.as_ref())
1611 .unwrap();
1612 let loaded_results = loaded_hnsw
1613 .search_basic(query, k, ¶ms, None, store.as_ref())
1614 .unwrap();
1615 assert_eq!(builder_results, loaded_results);
1616 }
1617
1618 fn brute_force_topk(store: &FlatFloatStorage, query: ArrayRef, k: usize) -> Vec<u32> {
1620 let dist_calc = store.dist_calculator(query, 0.0);
1621 let mut all: Vec<(f32, u32)> = (0..store.len() as u32)
1622 .map(|id| (dist_calc.distance(id), id))
1623 .collect();
1624 all.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
1625 all.into_iter().take(k).map(|(_, id)| id).collect()
1626 }
1627
1628 #[rstest]
1633 #[case::l2_single(DistanceType::L2, 1)]
1634 #[case::l2_pair(DistanceType::L2, 2)]
1635 #[case::l2_multi_level(DistanceType::L2, 2048)]
1636 #[case::dot_multi_level(DistanceType::Dot, 2048)]
1637 #[tokio::test]
1638 async fn test_loaded_search_parity_and_recall(
1639 #[case] distance_type: DistanceType,
1640 #[case] total: usize,
1641 ) {
1642 const DIM: usize = 32;
1643 let fsl =
1644 FixedSizeListArray::try_new_from_values(generate_random_array(total * DIM), DIM as i32)
1645 .unwrap();
1646 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), distance_type));
1647 let builder = HNSW::index_vectors(
1648 store.as_ref(),
1649 HnswBuildParams::default().num_edges(20).ef_construction(50),
1650 )
1651 .unwrap();
1652 assert!(!matches!(builder.inner.graph, HnswGraph::Loaded(_)));
1653
1654 let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
1655 assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
1656 assert_eq!(loaded.len(), total);
1657
1658 let k = total.min(10);
1659 let params = HnswQueryParams {
1660 ef: 50,
1661 lower_bound: None,
1662 upper_bound: None,
1663 dist_q_c: 0.0,
1664 use_acorn: false,
1665 };
1666 let query = fsl.value(0);
1667
1668 let builder_results = builder
1669 .search_basic(query.clone(), k, ¶ms, None, store.as_ref())
1670 .unwrap();
1671 let loaded_results = loaded
1672 .search_basic(query.clone(), k, ¶ms, None, store.as_ref())
1673 .unwrap();
1674 assert_eq!(builder_results, loaded_results);
1675
1676 let truth: std::collections::HashSet<u32> = brute_force_topk(store.as_ref(), query, k)
1678 .into_iter()
1679 .collect();
1680 let hits = loaded_results
1681 .iter()
1682 .filter(|n| truth.contains(&n.id))
1683 .count();
1684 let recall = hits as f32 / k as f32;
1685 assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1686 }
1687
1688 fn brute_force_topk_masked(
1690 store: &FlatFloatStorage,
1691 query: ArrayRef,
1692 k: usize,
1693 passes: impl Fn(u32) -> bool,
1694 ) -> Vec<u32> {
1695 let dist_calc = store.dist_calculator(query, 0.0);
1696 let mut matching: Vec<(f32, u32)> = (0..store.len() as u32)
1697 .filter(|id| passes(*id))
1698 .map(|id| (dist_calc.distance(id), id))
1699 .collect();
1700 matching.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
1701 matching.into_iter().take(k).map(|(_, id)| id).collect()
1702 }
1703
1704 #[tokio::test]
1707 async fn test_acorn_filtered_search() {
1708 const DIM: usize = 32;
1709 const TOTAL: usize = 2048;
1710 let fsl =
1711 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1712 .unwrap();
1713 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1714 let builder = HNSW::index_vectors(
1715 store.as_ref(),
1716 HnswBuildParams::default().num_edges(20).ef_construction(50),
1717 )
1718 .unwrap();
1719 let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
1720
1721 let mut mask_generator = VisitedGenerator::new(TOTAL);
1722 let k = 10;
1723 let params = HnswQueryParams {
1724 ef: 50,
1725 lower_bound: None,
1726 upper_bound: None,
1727 dist_q_c: 0.0,
1728 use_acorn: false,
1729 };
1730 let query = fsl.value(0);
1731 let truth: std::collections::HashSet<u32> =
1732 brute_force_topk_masked(store.as_ref(), query.clone(), k, |id| id % 2 == 0)
1733 .into_iter()
1734 .collect();
1735
1736 let mut all_results = vec![];
1737 for hnsw in [&builder, &loaded] {
1738 let mut bitset = mask_generator.generate(TOTAL);
1739 for id in (0..TOTAL as u32).step_by(2) {
1740 bitset.insert(id);
1741 }
1742 let results = hnsw
1743 .search_acorn(query.clone(), k, ¶ms, &bitset, store.as_ref())
1744 .unwrap();
1745 assert_eq!(results.len(), k);
1746 assert!(results.iter().all(|node| node.id % 2 == 0));
1747 assert!(results.windows(2).all(|w| w[0].dist <= w[1].dist));
1748 let hits = results.iter().filter(|n| truth.contains(&n.id)).count();
1749 let recall = hits as f32 / k as f32;
1750 assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1751 all_results.push(results);
1752 }
1753 assert_eq!(all_results[0], all_results[1]);
1754
1755 let default_ef_params = HnswQueryParams {
1757 ef: k + k / 2,
1758 ..params
1759 };
1760 for excluded_stride in [2, 400] {
1761 let passes = |id: u32| id % excluded_stride != 1;
1762 let mut bitset = mask_generator.generate(TOTAL);
1763 for id in (0..TOTAL as u32).filter(|id| passes(*id)) {
1764 bitset.insert(id);
1765 }
1766 let truth: std::collections::HashSet<u32> =
1767 brute_force_topk_masked(store.as_ref(), query.clone(), k, passes)
1768 .into_iter()
1769 .collect();
1770 let results = builder
1771 .search_acorn(
1772 query.clone(),
1773 k,
1774 &default_ef_params,
1775 &bitset,
1776 store.as_ref(),
1777 )
1778 .unwrap();
1779 assert_eq!(results.len(), k);
1780 assert!(results.iter().all(|node| passes(node.id)));
1781 let hits = results.iter().filter(|n| truth.contains(&n.id)).count();
1782 let recall = hits as f32 / k as f32;
1783 assert!(recall >= 0.5, "recall {recall} below 0.5 (k={k})");
1784 }
1785 }
1786
1787 #[tokio::test]
1790 async fn test_subindex_prefilter_dispatch() {
1791 use arrow_array::cast::AsArray;
1792 use async_trait::async_trait;
1793 use lance_core::Result;
1794 use lance_select::{RowAddrMask, RowAddrTreeMap};
1795
1796 use crate::metrics::NoOpMetricsCollector;
1797 use crate::prefilter::PreFilter;
1798
1799 struct MaskPreFilter {
1800 mask: Arc<RowAddrMask>,
1801 }
1802
1803 #[async_trait]
1804 impl PreFilter for MaskPreFilter {
1805 async fn wait_for_ready(&self) -> Result<()> {
1806 Ok(())
1807 }
1808 fn is_empty(&self) -> bool {
1809 false
1810 }
1811 fn mask(&self) -> Arc<RowAddrMask> {
1812 self.mask.clone()
1813 }
1814 fn filter_row_ids<'a>(
1815 &self,
1816 row_ids: Box<dyn Iterator<Item = &'a u64> + 'a>,
1817 ) -> Vec<u64> {
1818 self.mask.selected_indices(row_ids)
1819 }
1820 }
1821
1822 const DIM: usize = 32;
1823 const TOTAL: usize = 2048;
1824 let fsl =
1825 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1826 .unwrap();
1827 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1828 let hnsw = HNSW::index_vectors(
1829 store.as_ref(),
1830 HnswBuildParams::default().num_edges(20).ef_construction(50),
1831 )
1832 .unwrap();
1833
1834 let k = 10;
1835 let query_key = fsl.value(0);
1836
1837 let search_row_ids = |allowed: Vec<u64>, use_acorn: bool| {
1838 let params = HnswQueryParams {
1839 ef: 50,
1840 lower_bound: None,
1841 upper_bound: None,
1842 dist_q_c: 0.0,
1843 use_acorn,
1844 };
1845 let filter = Arc::new(MaskPreFilter {
1846 mask: Arc::new(RowAddrMask::from_allowed(RowAddrTreeMap::from_iter(
1847 allowed,
1848 ))),
1849 });
1850 let batch = hnsw
1851 .search(
1852 query_key.clone(),
1853 k,
1854 params,
1855 store.as_ref(),
1856 filter,
1857 &NoOpMetricsCollector,
1858 )
1859 .unwrap();
1860 batch[lance_core::ROW_ID]
1861 .as_primitive::<arrow_array::types::UInt64Type>()
1862 .values()
1863 .to_vec()
1864 };
1865
1866 let dense: Vec<u64> = (0..TOTAL as u64).step_by(2).collect();
1868 for use_acorn in [false, true] {
1869 let row_ids = search_row_ids(dense.clone(), use_acorn);
1870 assert_eq!(row_ids.len(), k);
1871 assert!(row_ids.iter().all(|id| id % 2 == 0));
1872 }
1873
1874 let all: Vec<u64> = (0..TOTAL as u64).collect();
1876 let unfiltered = hnsw
1877 .search_basic(
1878 query_key.clone(),
1879 k,
1880 &HnswQueryParams {
1881 ef: 50,
1882 lower_bound: None,
1883 upper_bound: None,
1884 dist_q_c: 0.0,
1885 use_acorn: false,
1886 },
1887 None,
1888 store.as_ref(),
1889 )
1890 .unwrap();
1891 let row_ids = search_row_ids(all, true);
1892 assert_eq!(
1893 row_ids,
1894 unfiltered.iter().map(|n| n.id as u64).collect::<Vec<_>>()
1895 );
1896
1897 let sparse: Vec<u64> = (0..TOTAL as u64).step_by(25).collect();
1899 let row_ids = search_row_ids(sparse.clone(), true);
1900 assert_eq!(row_ids.len(), k);
1901 let truth = brute_force_topk_masked(store.as_ref(), query_key.clone(), k, |id| {
1902 sparse.contains(&(id as u64))
1903 });
1904 let mut got: Vec<u32> = row_ids.iter().map(|id| *id as u32).collect();
1905 got.sort_unstable();
1906 let mut expected = truth;
1907 expected.sort_unstable();
1908 assert_eq!(got, expected);
1909 }
1910
1911 #[tokio::test]
1921 async fn test_loaded_level_offsets_misalignment_invariant() {
1922 use arrow::array::AsArray;
1923 use arrow::datatypes::UInt32Type;
1924
1925 const DIM: usize = 32;
1926 const TOTAL: usize = 2048;
1927 let fsl =
1928 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
1929 .unwrap();
1930 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
1931 let builder = HNSW::index_vectors(
1932 store.as_ref(),
1933 HnswBuildParams::default().num_edges(20).ef_construction(50),
1934 )
1935 .unwrap();
1936
1937 assert!(
1939 builder.max_level() >= 2,
1940 "expected a multi-level graph (got max_level {})",
1941 builder.max_level()
1942 );
1943
1944 let batch = builder.to_batch().unwrap();
1945 let md = builder.metadata();
1946 let total_counted = *md.level_offsets.last().unwrap();
1947
1948 assert!(
1952 batch.num_rows() > total_counted,
1953 "expected serialized rows ({}) to exceed sum(level_count) ({}) -- \
1954 entry point should be written at every level",
1955 batch.num_rows(),
1956 total_counted,
1957 );
1958
1959 let n = md.level_offsets[1];
1962 assert_eq!(n, TOTAL);
1963 let level0 = batch.slice(0, n);
1964 let ids = level0.column(0).as_primitive::<UInt32Type>();
1965 assert!(
1966 ids.values()
1967 .iter()
1968 .enumerate()
1969 .all(|(row, id)| *id == row as u32),
1970 "level-0 __vector_id must equal the row index",
1971 );
1972
1973 let loaded = HNSW::load(batch).unwrap();
1977 assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
1978 let params = HnswQueryParams {
1979 ef: 50,
1980 lower_bound: None,
1981 upper_bound: None,
1982 dist_q_c: 0.0,
1983 use_acorn: false,
1984 };
1985 let query = fsl.value(0);
1986 let builder_results = builder
1987 .search_basic(query.clone(), 10, ¶ms, None, store.as_ref())
1988 .unwrap();
1989 let loaded_results = loaded
1990 .search_basic(query, 10, ¶ms, None, store.as_ref())
1991 .unwrap();
1992 assert_eq!(builder_results, loaded_results);
1993 }
1994
1995 #[tokio::test]
2001 async fn test_load_rejects_misaligned_level0_id() {
2002 use arrow::array::AsArray;
2003 use arrow::datatypes::UInt32Type;
2004
2005 const DIM: usize = 16;
2006 const TOTAL: usize = 256;
2007 let fsl =
2008 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2009 .unwrap();
2010 let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2011 let builder = HNSW::index_vectors(
2012 store.as_ref(),
2013 HnswBuildParams::default().num_edges(20).ef_construction(50),
2014 )
2015 .unwrap();
2016
2017 let batch = builder.to_batch().unwrap();
2018 let mut ids = batch
2021 .column(0)
2022 .as_primitive::<UInt32Type>()
2023 .values()
2024 .to_vec();
2025 ids[0] = ids.len() as u32;
2026 let mut columns = batch.columns().to_vec();
2027 columns[0] = Arc::new(UInt32Array::from(ids));
2028 let corrupted = RecordBatch::try_new(batch.schema(), columns).unwrap();
2029
2030 assert!(
2031 HNSW::load(corrupted).is_err(),
2032 "load() must reject a misaligned level-0 __vector_id"
2033 );
2034 }
2035
2036 #[tokio::test]
2041 async fn test_load_rejects_out_of_range_entry_point() {
2042 use super::{HNSW_METADATA_KEY, HnswMetadata};
2043
2044 const DIM: usize = 16;
2045 const TOTAL: usize = 256;
2046 let fsl =
2047 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2048 .unwrap();
2049 let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2050 let builder = HNSW::index_vectors(
2051 store.as_ref(),
2052 HnswBuildParams::default().num_edges(20).ef_construction(50),
2053 )
2054 .unwrap();
2055
2056 let batch = builder.to_batch().unwrap();
2057 let mut metadata = batch.schema_ref().metadata().clone();
2058 let mut md: HnswMetadata =
2059 serde_json::from_str(metadata.get(HNSW_METADATA_KEY).unwrap()).unwrap();
2060 let n = md.level_offsets[1];
2062 md.entry_point = n as u32;
2063 metadata.insert(
2064 HNSW_METADATA_KEY.to_string(),
2065 serde_json::to_string(&md).unwrap(),
2066 );
2067 let schema = batch.schema().as_ref().clone().with_metadata(metadata);
2071 let corrupted = RecordBatch::try_new(Arc::new(schema), batch.columns().to_vec()).unwrap();
2072
2073 assert!(
2074 HNSW::load(corrupted).is_err(),
2075 "load() must reject an out-of-range entry_point"
2076 );
2077 }
2078
2079 #[tokio::test]
2081 async fn test_loaded_empty_index() {
2082 const DIM: usize = 16;
2083 let fsl =
2084 FixedSizeListArray::try_new_from_values(generate_random_array(0), DIM as i32).unwrap();
2085 let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2086 let builder = HNSW::index_vectors(store.as_ref(), HnswBuildParams::default()).unwrap();
2087 assert!(builder.is_empty());
2088
2089 let batch = builder.to_batch().unwrap();
2090 assert_eq!(batch.num_rows(), 0);
2091
2092 let loaded = HNSW::load(batch).unwrap();
2093 assert!(loaded.is_empty());
2094 assert_eq!(loaded.len(), 0);
2095 assert!(!matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2097 assert_eq!(loaded.to_batch().unwrap().num_rows(), 0);
2098 }
2099
2100 #[tokio::test]
2105 async fn test_to_batch_roundtrip_loaded() {
2106 const DIM: usize = 24;
2107 const TOTAL: usize = 1500;
2108 let fsl =
2109 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2110 .unwrap();
2111 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
2112 let builder = HNSW::index_vectors(
2113 store.as_ref(),
2114 HnswBuildParams::default().num_edges(16).ef_construction(50),
2115 )
2116 .unwrap();
2117
2118 let b1 = builder.to_batch().unwrap();
2119 let loaded = HNSW::load(b1.clone()).unwrap();
2120 assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2121 let b2 = loaded.to_batch().unwrap();
2122 assert_eq!(b1, b2);
2123
2124 let reloaded = HNSW::load(b2).unwrap();
2125 let params = HnswQueryParams {
2126 ef: 50,
2127 lower_bound: None,
2128 upper_bound: None,
2129 dist_q_c: 0.0,
2130 use_acorn: false,
2131 };
2132 let query = fsl.value(7);
2133 let a = builder
2134 .search_basic(query.clone(), 10, ¶ms, None, store.as_ref())
2135 .unwrap();
2136 let b = reloaded
2137 .search_basic(query, 10, ¶ms, None, store.as_ref())
2138 .unwrap();
2139 assert_eq!(a, b);
2140 }
2141
2142 #[tokio::test]
2150 async fn test_to_batch_loaded_preserves_extra_schema_metadata() {
2151 use super::HNSW_METADATA_KEY;
2152
2153 const DIM: usize = 24;
2154 const TOTAL: usize = 512;
2155 let fsl =
2156 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2157 .unwrap();
2158 let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2159 let builder = HNSW::index_vectors(
2160 store.as_ref(),
2161 HnswBuildParams::default().num_edges(16).ef_construction(50),
2162 )
2163 .unwrap();
2164
2165 let built_batch = builder.to_batch().unwrap();
2169 let mut metadata = built_batch.schema_ref().metadata().clone();
2170 metadata.insert(
2171 "lance:index_metadata".to_string(),
2172 "{\"distance_type\":\"l2\"}".to_string(),
2173 );
2174 metadata.insert("lance:ivf".to_string(), "42".to_string());
2175 let schema = built_batch
2176 .schema()
2177 .as_ref()
2178 .clone()
2179 .with_metadata(metadata);
2180 let batch_with_extra =
2181 RecordBatch::try_new(Arc::new(schema), built_batch.columns().to_vec()).unwrap();
2182
2183 let loaded = HNSW::load(batch_with_extra).unwrap();
2184 assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2185
2186 let out = loaded.to_batch().unwrap();
2189 let out_metadata = out.schema_ref().metadata();
2190 assert!(out_metadata.contains_key(HNSW_METADATA_KEY));
2191 assert_eq!(
2192 out_metadata.get("lance:index_metadata").map(String::as_str),
2193 Some("{\"distance_type\":\"l2\"}"),
2194 );
2195 assert_eq!(
2196 out_metadata.get("lance:ivf").map(String::as_str),
2197 Some("42")
2198 );
2199
2200 let reloaded = HNSW::load(out).unwrap();
2202 assert_eq!(reloaded.len(), loaded.len());
2203 }
2204
2205 #[tokio::test]
2209 async fn test_loaded_graph_is_arrow_backed() {
2210 const DIM: usize = 32;
2211 const TOTAL: usize = 2048;
2212 let fsl =
2213 FixedSizeListArray::try_new_from_values(generate_random_array(TOTAL * DIM), DIM as i32)
2214 .unwrap();
2215 let store = Arc::new(FlatFloatStorage::new(fsl, DistanceType::L2));
2216 let builder = HNSW::index_vectors(
2217 store.as_ref(),
2218 HnswBuildParams::default().num_edges(20).ef_construction(50),
2219 )
2220 .unwrap();
2221 assert!(!matches!(builder.inner.graph, HnswGraph::Loaded(_)));
2222
2223 let loaded = HNSW::load(builder.to_batch().unwrap()).unwrap();
2224 assert!(matches!(loaded.inner.graph, HnswGraph::Loaded(_)));
2225 assert!(
2226 loaded.deep_size_of() < builder.deep_size_of(),
2227 "loaded graph ({}) should be lighter than built ({})",
2228 loaded.deep_size_of(),
2229 builder.deep_size_of(),
2230 );
2231 }
2232}