1use arrow::array::{AsArray, ListBuilder, UInt32Builder};
7use arrow::compute::concat_batches;
8use arrow::datatypes::{Float32Type, UInt32Type};
9use arrow_array::{ArrayRef, Float32Array, RecordBatch, UInt64Array};
10use crossbeam_queue::ArrayQueue;
11use deepsize::DeepSizeOf;
12use itertools::Itertools;
13
14use lance_core::utils::tokio::get_num_compute_intensive_cpus;
15use lance_linalg::distance::DistanceType;
16use rayon::prelude::*;
17use std::cmp::min;
18use std::collections::{BinaryHeap, HashMap, VecDeque};
19use std::fmt::Debug;
20use std::iter;
21use std::sync::Arc;
22use std::sync::RwLock;
23use std::sync::atomic::{AtomicUsize, Ordering};
24use tracing::instrument;
25
26use lance_core::{Error, Result};
27use rand::{Rng, rng};
28use serde::{Deserialize, Serialize};
29
30use super::super::graph::beam_search;
31use super::{HNSW_TYPE, HnswMetadata, VECTOR_ID_COL, VECTOR_ID_FIELD, select_neighbors_heuristic};
32use crate::metrics::MetricsCollector;
33use crate::prefilter::PreFilter;
34use crate::vector::flat::storage::FlatFloatStorage;
35use crate::vector::graph::builder::GraphBuilderNode;
36use crate::vector::graph::{
37 BorrowingGraph, DISTS_FIELD, Graph, NEIGHBORS_COL, NEIGHBORS_FIELD, OrderedFloat, OrderedNode,
38 VisitedGenerator,
39};
40use crate::vector::graph::{Visited, beam_search_borrowed, greedy_search, greedy_search_borrowed};
41use crate::vector::storage::{DistCalculator, VectorStore};
42use crate::vector::v3::subindex::IvfSubIndex;
43use crate::vector::{DIST_COL, Query, VECTOR_RESULT_SCHEMA};
44
45pub const HNSW_METADATA_KEY: &str = "lance:hnsw";
46
47#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
49pub struct HnswBuildParams {
50 pub max_level: u16,
52
53 pub m: usize,
55
56 pub ef_construction: usize,
58
59 pub prefetch_distance: Option<usize>,
61}
62
63impl Default for HnswBuildParams {
64 fn default() -> Self {
65 Self {
66 max_level: 7,
67 m: 20,
68 ef_construction: 150,
69 prefetch_distance: Some(2),
70 }
71 }
72}
73
74impl HnswBuildParams {
75 pub fn max_level(mut self, max_level: u16) -> Self {
78 self.max_level = max_level;
79 self
80 }
81
82 pub fn num_edges(mut self, m: usize) -> Self {
85 self.m = m;
86 self
87 }
88
89 pub fn ef_construction(mut self, ef_construction: usize) -> Self {
94 self.ef_construction = ef_construction;
95 self
96 }
97
98 pub async fn build(self, data: ArrayRef, distance_type: DistanceType) -> Result<HNSW> {
104 let vec_store = Arc::new(FlatFloatStorage::new(
105 data.as_fixed_size_list().clone(),
106 distance_type,
107 ));
108 HNSW::index_vectors(vec_store.as_ref(), self)
109 }
110}
111
112#[derive(Clone, DeepSizeOf)]
120pub struct HNSW {
121 inner: Arc<HnswCore>,
122}
123
124struct HnswCore {
125 params: HnswBuildParams,
126 nodes: Arc<Vec<GraphBuilderNode>>,
127 level_count: Vec<usize>,
128 entry_point: u32,
129 visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
130}
131
132impl DeepSizeOf for HnswCore {
133 fn deep_size_of_children(&self, context: &mut deepsize::Context) -> usize {
134 self.params.deep_size_of_children(context)
135 + self.nodes.deep_size_of_children(context)
136 + self.level_count.deep_size_of_children(context)
137 }
139}
140
141impl HnswCore {
142 fn max_level(&self) -> u16 {
143 self.params.max_level
144 }
145
146 fn num_nodes(&self, level: usize) -> usize {
147 self.level_count[level]
148 }
149
150 fn nodes(&self) -> Arc<Vec<GraphBuilderNode>> {
151 self.nodes.clone()
152 }
153}
154
155impl Debug for HNSW {
156 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
157 write!(f, "HNSW(max_layers: {})", self.inner.max_level() as usize,)
158 }
159}
160
161impl HNSW {
162 pub fn empty() -> Self {
163 Self {
164 inner: Arc::new(HnswCore {
165 params: HnswBuildParams::default(),
166 nodes: Arc::new(Vec::new()),
167 level_count: Vec::new(),
168 entry_point: 0,
169 visited_generator_queue: Arc::new(ArrayQueue::new(1)),
170 }),
171 }
172 }
173
174 pub fn len(&self) -> usize {
175 self.inner.nodes.len()
176 }
177
178 pub fn is_empty(&self) -> bool {
179 self.len() == 0
180 }
181
182 pub fn max_level(&self) -> u16 {
183 self.inner.max_level()
184 }
185
186 pub fn num_nodes(&self, level: usize) -> usize {
187 self.inner.num_nodes(level)
188 }
189
190 pub fn nodes(&self) -> Arc<Vec<GraphBuilderNode>> {
191 self.inner.nodes()
192 }
193
194 #[allow(clippy::too_many_arguments)]
195 pub fn search_inner(
196 &self,
197 query: ArrayRef,
198 k: usize,
199 params: &HnswQueryParams,
200 bitset: Option<Visited>,
201 visited_generator: &mut VisitedGenerator,
202 storage: &impl VectorStore,
203 prefetch_distance: Option<usize>,
204 ) -> Result<Vec<OrderedNode>> {
205 let dist_calc = storage.dist_calculator(query, params.dist_q_c);
206 let mut ep = OrderedNode::new(0, dist_calc.distance(0).into());
207 let nodes = self.inner.nodes.as_ref();
208 for level in (0..self.max_level()).rev() {
209 let cur_level = ImmutableHnswLevelView::new(level, nodes);
210 ep = greedy_search_borrowed(
211 &cur_level,
212 ep,
213 &dist_calc,
214 self.inner.params.prefetch_distance,
215 );
216 }
217
218 let bottom_level = ImmutableHnswBottomView::new(nodes);
219 let mut visited = visited_generator.generate(storage.len());
220 Ok(beam_search_borrowed(
221 &bottom_level,
222 &ep,
223 params,
224 &dist_calc,
225 bitset.as_ref(),
226 prefetch_distance,
227 &mut visited,
228 )
229 .into_iter()
230 .take(k)
231 .collect())
232 }
233
234 #[instrument(level = "debug", skip(self, query, bitset, storage))]
235 pub fn search_basic(
236 &self,
237 query: ArrayRef,
238 k: usize,
239 params: &HnswQueryParams,
240 bitset: Option<Visited>,
241 storage: &impl VectorStore,
242 ) -> Result<Vec<OrderedNode>> {
243 let mut visited_generator = self
244 .inner
245 .visited_generator_queue
246 .pop()
247 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
248 let result = self.search_inner(
249 query,
250 k,
251 params,
252 bitset,
253 &mut visited_generator,
254 storage,
255 Some(2),
256 );
257
258 match self.inner.visited_generator_queue.push(visited_generator) {
259 Ok(_) => {}
260 Err(_) => {
261 log::warn!("visited_generator_queue is full");
262 }
263 }
264
265 result
266 }
267
268 #[instrument(level = "debug", skip(self, storage, query, prefilter_bitset))]
269 fn flat_search(
270 &self,
271 storage: &impl VectorStore,
272 query: ArrayRef,
273 k: usize,
274 prefilter_bitset: Visited,
275 params: &HnswQueryParams,
276 ) -> Vec<OrderedNode> {
277 let lower_bound: OrderedFloat = params.lower_bound.unwrap_or(f32::MIN).into();
278 let upper_bound: OrderedFloat = params.upper_bound.unwrap_or(f32::MAX).into();
279
280 let dist_calc = storage.dist_calculator(query, params.dist_q_c);
281 let mut heap = BinaryHeap::<OrderedNode>::with_capacity(k);
282
283 match self.inner.params.prefetch_distance {
284 Some(ahead) if ahead > 0 => {
285 let mut ids_iter = prefilter_bitset.iter_ones().map(|i| i as u32);
286 let mut buffer = VecDeque::with_capacity(ahead + 1);
287 for _ in 0..=ahead {
288 if let Some(id) = ids_iter.next() {
289 buffer.push_back(id);
290 } else {
291 break;
292 }
293 }
294
295 while let Some(node_id) = buffer.pop_front() {
296 if let Some(&prefetch_id) = buffer.get(ahead - 1) {
297 dist_calc.prefetch(prefetch_id);
298 }
299 if let Some(next) = ids_iter.next() {
300 buffer.push_back(next);
301 }
302
303 let dist: OrderedFloat = dist_calc.distance(node_id).into();
304 if dist <= lower_bound || dist > upper_bound {
305 continue;
306 }
307 if heap.len() < k {
308 heap.push((dist, node_id).into());
309 } else if dist < heap.peek().unwrap().dist {
310 heap.pop();
311 heap.push((dist, node_id).into());
312 }
313 }
314 }
315 _ => {
316 for node_id in prefilter_bitset.iter_ones().map(|i| i as u32) {
317 let dist: OrderedFloat = dist_calc.distance(node_id).into();
318 if dist <= lower_bound || dist > upper_bound {
319 continue;
320 }
321 if heap.len() < k {
322 heap.push((dist, node_id).into());
323 } else if dist < heap.peek().unwrap().dist {
324 heap.pop();
325 heap.push((dist, node_id).into());
326 }
327 }
328 }
329 };
330 heap.into_sorted_vec()
331 }
332
333 pub fn metadata(&self) -> HnswMetadata {
335 let level_offsets = self
338 .inner
339 .level_count
340 .iter()
341 .chain(iter::once(&0))
342 .scan(0, |state, x| {
343 let start = *state;
344 *state += *x;
345 Some(start)
346 })
347 .collect();
348
349 HnswMetadata {
350 entry_point: self.inner.entry_point,
351 params: self.inner.params.clone(),
352 level_offsets,
353 }
354 }
355}
356
357struct HnswBuilder {
358 params: HnswBuildParams,
359
360 nodes: Arc<Vec<RwLock<GraphBuilderNode>>>,
361 level_count: Vec<AtomicUsize>,
362
363 entry_point: u32,
364
365 visited_generator_queue: Arc<ArrayQueue<VisitedGenerator>>,
366}
367
368impl DeepSizeOf for HnswBuilder {
369 fn deep_size_of_children(&self, context: &mut deepsize::Context) -> usize {
370 self.params.deep_size_of_children(context)
371 + self.nodes.deep_size_of_children(context)
372 + self.level_count.deep_size_of_children(context)
373 }
375}
376
377impl HnswBuilder {
378 fn finish(self) -> HNSW {
379 let nodes = match Arc::try_unwrap(self.nodes) {
380 Ok(nodes) => nodes
381 .into_iter()
382 .map(|node| node.into_inner().expect("builder lock poisoned"))
383 .collect(),
384 Err(nodes) => nodes
385 .iter()
386 .map(|node| node.read().expect("builder lock poisoned").clone())
387 .collect(),
388 };
389
390 let level_count = self
391 .level_count
392 .into_iter()
393 .map(|count| count.load(Ordering::Relaxed))
394 .collect();
395
396 HNSW {
397 inner: Arc::new(HnswCore {
398 params: self.params,
399 nodes: Arc::new(nodes),
400 level_count,
401 entry_point: self.entry_point,
402 visited_generator_queue: self.visited_generator_queue,
403 }),
404 }
405 }
406
407 pub fn with_params(params: HnswBuildParams, storage: &impl VectorStore) -> Self {
409 let len = storage.len();
410 let max_level = params.max_level;
411
412 let level_count = (0..max_level)
413 .map(|_| AtomicUsize::new(0))
414 .collect::<Vec<_>>();
415
416 let visited_generator_queue = Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus()));
417 for _ in 0..get_num_compute_intensive_cpus() {
418 visited_generator_queue
419 .push(VisitedGenerator::new(0))
420 .unwrap();
421 }
422 let mut builder = Self {
423 params,
424 nodes: Arc::new(Vec::new()),
425 level_count,
426 entry_point: 0,
427 visited_generator_queue,
428 };
429
430 if storage.is_empty() {
431 return builder;
432 }
433
434 let mut nodes = Vec::with_capacity(len);
435 {
436 if len > 0 {
437 nodes.push(RwLock::new(GraphBuilderNode::new(0, max_level as usize)));
438 }
439 let mut level_rng = rng();
440 for i in 1..len {
441 nodes.push(RwLock::new(GraphBuilderNode::new(
442 i as u32,
443 builder.random_level(&mut level_rng) as usize + 1,
444 )));
445 }
446 }
447 builder.nodes = Arc::new(nodes);
448
449 builder
450 }
451
452 fn random_level<R: Rng + ?Sized>(&self, rng: &mut R) -> u16 {
456 let ml = 1.0 / (self.params.m as f32).ln();
457 min(
458 (-rng.random::<f32>().ln() * ml) as u16,
459 self.params.max_level - 1,
460 )
461 }
462
463 fn insert(
465 &self,
466 node: u32,
467 visited_generator: &mut VisitedGenerator,
468 storage: &impl VectorStore,
469 ) {
470 let nodes = &self.nodes;
471 let target_level = nodes[node as usize].read().unwrap().level_neighbors.len() as u16 - 1;
472 let dist_calc = storage.dist_calculator_from_id(node);
473 let mut ep = OrderedNode::new(
474 self.entry_point,
475 dist_calc.distance(self.entry_point).into(),
476 );
477
478 for level in (target_level + 1..self.params.max_level).rev() {
487 let cur_level = HnswLevelView::new(level, nodes);
488 ep = greedy_search(&cur_level, ep, &dist_calc, self.params.prefetch_distance);
489 }
490
491 let mut pruned_neighbors_per_level: Vec<Vec<_>> =
492 vec![Vec::new(); (target_level + 1) as usize];
493 {
494 let mut current_node = nodes[node as usize].write().unwrap();
495 for level in (0..=target_level).rev() {
496 self.level_count[level as usize].fetch_add(1, Ordering::Relaxed);
497
498 let neighbors = self.search_level(&ep, level, &dist_calc, nodes, visited_generator);
499 for neighbor in &neighbors {
500 current_node.add_neighbor(neighbor.id, neighbor.dist, level);
501 }
502 self.prune(storage, &mut current_node, level);
503 pruned_neighbors_per_level[level as usize]
504 .clone_from(¤t_node.level_neighbors_ranked[level as usize]);
505
506 ep = neighbors[0].clone();
507 }
508 }
509 for (level, pruned_neighbors) in pruned_neighbors_per_level.iter().enumerate() {
510 let _: Vec<_> = pruned_neighbors
511 .iter()
512 .map(|unpruned_edge| {
513 let level = level as u16;
514 let m_max = match level {
515 0 => self.params.m * 2,
516 _ => self.params.m,
517 };
518 if unpruned_edge.dist
519 < nodes[unpruned_edge.id as usize]
520 .read()
521 .unwrap()
522 .cutoff(level, m_max)
523 {
524 let mut chosen_node = nodes[unpruned_edge.id as usize].write().unwrap();
525 chosen_node.add_neighbor(node, unpruned_edge.dist, level);
526 self.prune(storage, &mut chosen_node, level);
527 }
528 })
529 .collect();
530 }
531 }
532
533 fn search_level(
534 &self,
535 ep: &OrderedNode,
536 level: u16,
537 dist_calc: &impl DistCalculator,
538 nodes: &[RwLock<GraphBuilderNode>],
539 visited_generator: &mut VisitedGenerator,
540 ) -> Vec<OrderedNode> {
541 let cur_level = HnswLevelView::new(level, nodes);
542 let mut visited = visited_generator.generate(nodes.len());
543 beam_search(
544 &cur_level,
545 ep,
546 &HnswQueryParams {
547 ef: self.params.ef_construction,
548 lower_bound: None,
549 upper_bound: None,
550 dist_q_c: 0.0,
551 },
552 dist_calc,
553 None,
554 self.params.prefetch_distance,
555 &mut visited,
556 )
557 }
558
559 fn prune(&self, storage: &impl VectorStore, builder_node: &mut GraphBuilderNode, level: u16) {
560 let m_max = match level {
561 0 => self.params.m * 2,
562 _ => self.params.m,
563 };
564
565 let neighbors_ranked = &mut builder_node.level_neighbors_ranked[level as usize];
566 let level_neighbors = neighbors_ranked.clone();
567 if level_neighbors.len() <= m_max {
568 builder_node.update_from_ranked_neighbors(level);
569 return;
570 }
571
572 *neighbors_ranked = select_neighbors_heuristic(storage, &level_neighbors, m_max);
573 builder_node.update_from_ranked_neighbors(level);
574 }
575}
576
577pub(crate) struct HnswLevelView<'a> {
580 level: u16,
581 nodes: &'a [RwLock<GraphBuilderNode>],
582}
583
584impl<'a> HnswLevelView<'a> {
585 pub fn new(level: u16, nodes: &'a [RwLock<GraphBuilderNode>]) -> Self {
586 Self { level, nodes }
587 }
588}
589
590impl Graph for HnswLevelView<'_> {
591 fn len(&self) -> usize {
592 self.nodes.len()
593 }
594
595 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
596 let node = &self.nodes[key as usize];
597 node.read().unwrap().level_neighbors[self.level as usize].clone()
598 }
599}
600
601pub(crate) struct ImmutableHnswLevelView<'a> {
602 level: u16,
603 nodes: &'a [GraphBuilderNode],
604}
605
606impl<'a> ImmutableHnswLevelView<'a> {
607 pub fn new(level: u16, nodes: &'a [GraphBuilderNode]) -> Self {
608 Self { level, nodes }
609 }
610}
611
612impl Graph for ImmutableHnswLevelView<'_> {
613 fn len(&self) -> usize {
614 self.nodes.len()
615 }
616
617 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
618 self.nodes[key as usize].level_neighbors[self.level as usize].clone()
619 }
620}
621
622impl BorrowingGraph for ImmutableHnswLevelView<'_> {
623 fn len(&self) -> usize {
624 self.nodes.len()
625 }
626
627 fn neighbors(&self, key: u32) -> &[u32] {
628 self.nodes[key as usize].level_neighbors[self.level as usize].as_slice()
629 }
630}
631
632pub(crate) struct ImmutableHnswBottomView<'a> {
633 nodes: &'a [GraphBuilderNode],
634}
635
636impl<'a> ImmutableHnswBottomView<'a> {
637 pub fn new(nodes: &'a [GraphBuilderNode]) -> Self {
638 Self { nodes }
639 }
640}
641
642impl Graph for ImmutableHnswBottomView<'_> {
643 fn len(&self) -> usize {
644 self.nodes.len()
645 }
646
647 fn neighbors(&self, key: u32) -> Arc<Vec<u32>> {
648 self.nodes[key as usize].bottom_neighbors.clone()
649 }
650}
651
652impl BorrowingGraph for ImmutableHnswBottomView<'_> {
653 fn len(&self) -> usize {
654 self.nodes.len()
655 }
656
657 fn neighbors(&self, key: u32) -> &[u32] {
658 self.nodes[key as usize].bottom_neighbors.as_slice()
659 }
660}
661
662#[derive(Debug, Clone, Copy)]
663pub struct HnswQueryParams {
664 pub ef: usize,
665 pub lower_bound: Option<f32>,
666 pub upper_bound: Option<f32>,
667 pub dist_q_c: f32,
668}
669
670impl From<&Query> for HnswQueryParams {
671 fn from(query: &Query) -> Self {
672 let k = query.k * query.refine_factor.unwrap_or(1) as usize;
673 Self {
674 ef: query.ef.unwrap_or(k + k / 2),
675 lower_bound: query.lower_bound,
676 upper_bound: query.upper_bound,
677 dist_q_c: query.dist_q_c,
678 }
679 }
680}
681
682impl IvfSubIndex for HNSW {
683 type BuildParams = HnswBuildParams;
684 type QueryParams = HnswQueryParams;
685
686 fn load(data: RecordBatch) -> Result<Self>
687 where
688 Self: Sized,
689 {
690 if data.num_rows() == 0 {
691 return Ok(Self::empty());
692 }
693
694 let hnsw_metadata = data
695 .schema_ref()
696 .metadata()
697 .get(HNSW_METADATA_KEY)
698 .ok_or(Error::index(format!("{} not found", HNSW_METADATA_KEY)))?;
699 let hnsw_metadata: HnswMetadata = serde_json::from_str(hnsw_metadata).map_err(|e| {
700 Error::index(format!(
701 "Failed to decode HNSW metadata: {}, json: {}",
702 e, hnsw_metadata
703 ))
704 })?;
705
706 let levels: Vec<_> = hnsw_metadata
707 .level_offsets
708 .iter()
709 .tuple_windows()
710 .map(|(start, end)| data.slice(*start, end - start))
711 .collect();
712
713 let level_count = levels.iter().map(|b| b.num_rows()).collect::<Vec<_>>();
714
715 let bottom_level_len = levels[0].num_rows();
716 let mut nodes = Vec::with_capacity(bottom_level_len);
717 for i in 0..bottom_level_len {
718 nodes.push(GraphBuilderNode::new(i as u32, levels.len()));
719 }
720 for (level, batch) in levels.into_iter().enumerate() {
721 let ids = batch[VECTOR_ID_COL].as_primitive::<UInt32Type>();
722 let neighbors = batch[NEIGHBORS_COL].as_list::<i32>();
723 let distances = batch[DIST_COL].as_list::<i32>();
724
725 for ((node, neighbors), distances) in
726 ids.iter().zip(neighbors.iter()).zip(distances.iter())
727 {
728 let node = node.unwrap();
729 let neighbors = neighbors.as_ref().unwrap().as_primitive::<UInt32Type>();
730 let distances = distances.as_ref().unwrap().as_primitive::<Float32Type>();
731
732 nodes[node as usize].level_neighbors_ranked[level] = neighbors
733 .iter()
734 .zip(distances.iter())
735 .map(|(n, dist)| OrderedNode::new(n.unwrap(), OrderedFloat(dist.unwrap())))
736 .collect();
737 nodes[node as usize].update_from_ranked_neighbors(level as u16);
738 }
739 }
740
741 let visited_generator_queue =
742 Arc::new(ArrayQueue::new(get_num_compute_intensive_cpus() * 2));
743 for _ in 0..get_num_compute_intensive_cpus() * 2 {
744 visited_generator_queue
745 .push(VisitedGenerator::new(0))
746 .unwrap();
747 }
748 let inner = HnswCore {
749 params: hnsw_metadata.params,
750 nodes: Arc::new(nodes),
751 level_count,
752 entry_point: hnsw_metadata.entry_point,
753 visited_generator_queue,
754 };
755
756 Ok(Self {
757 inner: Arc::new(inner),
758 })
759 }
760
761 fn name() -> &'static str {
762 HNSW_TYPE
763 }
764
765 fn metadata_key() -> &'static str {
766 "lance:hnsw"
767 }
768
769 fn schema() -> arrow_schema::SchemaRef {
771 arrow_schema::Schema::new(vec![
772 VECTOR_ID_FIELD.clone(),
773 NEIGHBORS_FIELD.clone(),
774 DISTS_FIELD.clone(),
775 ])
776 .into()
777 }
778
779 #[instrument(level = "debug", skip(self, query, storage, prefilter, _metrics))]
780 fn search(
781 &self,
782 query: ArrayRef,
783 k: usize,
784 params: Self::QueryParams,
785 storage: &impl VectorStore,
786 prefilter: Arc<dyn PreFilter>,
787 _metrics: &dyn MetricsCollector,
788 ) -> Result<RecordBatch> {
789 if params.ef < k {
790 return Err(Error::index(
791 "ef must be greater than or equal to k".to_string(),
792 ));
793 }
794
795 let schema = VECTOR_RESULT_SCHEMA.clone();
796 if self.is_empty() {
797 return Ok(RecordBatch::new_empty(schema));
798 }
799
800 let mut prefilter_generator = self
801 .inner
802 .visited_generator_queue
803 .pop()
804 .unwrap_or_else(|| VisitedGenerator::new(storage.len()));
805 let prefilter_bitset = if prefilter.is_empty() {
806 None
807 } else {
808 let indices = prefilter.filter_row_ids(Box::new(storage.row_ids()));
809 let mut bitset = prefilter_generator.generate(storage.len());
810 for indices in indices {
811 bitset.insert(indices as u32);
812 }
813 Some(bitset)
814 };
815
816 let remained = prefilter_bitset
817 .as_ref()
818 .map(|b| b.count_ones())
819 .unwrap_or(storage.len());
820 let results = if remained < self.len() * 10 / 100 {
821 let prefilter_bitset =
822 prefilter_bitset.expect("the prefilter bitset must be set for flat search");
823 self.flat_search(storage, query, k, prefilter_bitset, ¶ms)
824 } else {
825 self.search_basic(query, k, ¶ms, prefilter_bitset, storage)?
826 };
827 let _ = self.inner.visited_generator_queue.push(prefilter_generator);
829
830 let (row_ids, dists): (Vec<_>, Vec<_>) = results
832 .into_iter()
833 .map(|r| (storage.row_id(r.id), r.dist.0))
834 .unique_by(|r| r.0)
835 .unzip();
836 let row_ids = Arc::new(UInt64Array::from(row_ids));
837 let distances = Arc::new(Float32Array::from(dists));
838
839 Ok(RecordBatch::try_new(schema, vec![distances, row_ids])?)
840 }
841
842 fn index_vectors(storage: &impl VectorStore, params: Self::BuildParams) -> Result<Self>
844 where
845 Self: Sized,
846 {
847 let builder = HnswBuilder::with_params(params, storage);
848
849 log::debug!(
850 "Building HNSW graph: num={}, max_levels={}, m={}, ef_construction={}, distance_type:{}",
851 storage.len(),
852 builder.params.max_level,
853 builder.params.m,
854 builder.params.ef_construction,
855 storage.distance_type(),
856 );
857
858 if storage.is_empty() {
859 return Ok(builder.finish());
860 }
861
862 let len = storage.len();
863 builder.level_count[0].fetch_add(1, Ordering::Relaxed);
864 (1..len).into_par_iter().for_each_init(
865 || VisitedGenerator::new(len),
866 |visited_generator, node| {
867 builder.insert(node as u32, visited_generator, storage);
868 },
869 );
870
871 assert_eq!(builder.level_count[0].load(Ordering::Relaxed), len);
872 Ok(builder.finish())
873 }
874
875 fn remap(
876 &self,
877 _mapping: &HashMap<u64, Option<u64>>, store: &impl VectorStore,
879 ) -> Result<Self> {
880 Self::index_vectors(store, self.inner.params.clone())
883 }
884
885 fn to_batch(&self) -> Result<RecordBatch> {
887 let mut vector_id_builder = UInt32Builder::with_capacity(self.len());
888 let mut neighbors_builder = ListBuilder::with_capacity(UInt32Builder::new(), self.len());
889 let mut distances_builder =
890 ListBuilder::with_capacity(arrow_array::builder::Float32Builder::new(), self.len());
891 let mut batches = Vec::with_capacity(self.max_level() as usize);
892 for level in 0..self.max_level() {
893 let level = level as usize;
894 for (id, node) in self.inner.nodes.iter().enumerate() {
895 if level >= node.level_neighbors.len() {
896 continue;
897 }
898 let neighbors = node.level_neighbors[level].iter().map(|n| Some(*n));
899 let distances = node.level_neighbors_ranked[level]
900 .iter()
901 .map(|n| Some(n.dist.0));
902 vector_id_builder.append_value(id as u32);
903 neighbors_builder.append_value(neighbors);
904 distances_builder.append_value(distances);
905 }
906
907 let batch = RecordBatch::try_new(
908 Self::schema(),
909 vec![
910 Arc::new(vector_id_builder.finish()),
911 Arc::new(neighbors_builder.finish()),
912 Arc::new(distances_builder.finish()),
913 ],
914 )?;
915 batches.push(batch);
916 }
917
918 let metadata = self.metadata();
919 let metadata = serde_json::to_string(&metadata)?;
920 let schema = Self::schema()
921 .as_ref()
922 .clone()
923 .with_metadata(HashMap::from_iter(vec![(
924 HNSW_METADATA_KEY.to_string(),
925 metadata,
926 )]));
927 let batch = concat_batches(&Self::schema(), batches.iter())?;
928 let batch = batch.with_schema(Arc::new(schema))?;
929 Ok(batch)
930 }
931}
932
933#[cfg(test)]
934mod tests {
935 use std::sync::Arc;
936
937 use arrow_array::FixedSizeListArray;
938 use arrow_schema::Schema;
939 use lance_arrow::FixedSizeListArrayExt;
940 use lance_file::previous::{
941 reader::FileReader as PreviousFileReader,
942 writer::{
943 FileWriter as PreviousFileWriter, FileWriterOptions as PreviousFileWriterOptions,
944 },
945 };
946 use lance_io::object_store::ObjectStore;
947 use lance_linalg::distance::DistanceType;
948 use lance_table::format::SelfDescribingFileReader;
949 use lance_table::io::manifest::ManifestDescribing;
950 use lance_testing::datagen::generate_random_array;
951 use object_store::path::Path;
952
953 use crate::scalar::IndexWriter;
954 use crate::vector::v3::subindex::IvfSubIndex;
955 use crate::vector::{
956 flat::storage::FlatFloatStorage,
957 graph::{DISTS_FIELD, NEIGHBORS_FIELD},
958 hnsw::{
959 HNSW, VECTOR_ID_FIELD,
960 builder::{HnswBuildParams, HnswQueryParams},
961 },
962 };
963
964 #[tokio::test]
965 async fn test_builder_write_load() {
966 const DIM: usize = 32;
967 const TOTAL: usize = 2048;
968 const NUM_EDGES: usize = 20;
969 let data = generate_random_array(TOTAL * DIM);
970 let fsl = FixedSizeListArray::try_new_from_values(data, DIM as i32).unwrap();
971 let store = Arc::new(FlatFloatStorage::new(fsl.clone(), DistanceType::L2));
972 let builder = HNSW::index_vectors(
973 store.as_ref(),
974 HnswBuildParams::default()
975 .num_edges(NUM_EDGES)
976 .ef_construction(50),
977 )
978 .unwrap();
979
980 let object_store = ObjectStore::memory();
981 let path = Path::from("test_builder_write_load");
982 let writer = object_store.create(&path).await.unwrap();
983 let schema = Schema::new(vec![
984 VECTOR_ID_FIELD.clone(),
985 NEIGHBORS_FIELD.clone(),
986 DISTS_FIELD.clone(),
987 ]);
988 let schema = lance_core::datatypes::Schema::try_from(&schema).unwrap();
989 let mut writer = PreviousFileWriter::<ManifestDescribing>::with_object_writer(
990 writer,
991 schema,
992 &PreviousFileWriterOptions::default(),
993 )
994 .unwrap();
995 let batch = builder.to_batch().unwrap();
996 let metadata = batch.schema_ref().metadata().clone();
997 writer.write_record_batch(batch).await.unwrap();
998 writer.finish_with_metadata(&metadata).await.unwrap();
999
1000 let reader = PreviousFileReader::try_new_self_described(&object_store, &path, None)
1001 .await
1002 .unwrap();
1003 let batch = reader
1004 .read_range(0..reader.len(), reader.schema())
1005 .await
1006 .unwrap();
1007 let loaded_hnsw = HNSW::load(batch).unwrap();
1008
1009 let query = fsl.value(0);
1010 let k = 10;
1011 let params = HnswQueryParams {
1012 ef: 50,
1013 lower_bound: None,
1014 upper_bound: None,
1015 dist_q_c: 0.0,
1016 };
1017 let builder_results = builder
1018 .search_basic(query.clone(), k, ¶ms, None, store.as_ref())
1019 .unwrap();
1020 let loaded_results = loaded_hnsw
1021 .search_basic(query, k, ¶ms, None, store.as_ref())
1022 .unwrap();
1023 assert_eq!(builder_results, loaded_results);
1024 }
1025}