Skip to main content

lance_index/vector/hnsw/
builder.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Builder of Hnsw Graph.
5
6use 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/// Parameters of building HNSW index
48#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
49pub struct HnswBuildParams {
50    /// max level ofm
51    pub max_level: u16,
52
53    /// number of connections to establish while inserting new element
54    pub m: usize,
55
56    /// size of the dynamic list for the candidates
57    pub ef_construction: usize,
58
59    /// number of vectors ahead to prefetch while building the graph
60    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    /// The maximum level of the graph.
76    /// The default value is `8`.
77    pub fn max_level(mut self, max_level: u16) -> Self {
78        self.max_level = max_level;
79        self
80    }
81
82    /// The number of connections to establish while inserting new element
83    /// The default value is `30`.
84    pub fn num_edges(mut self, m: usize) -> Self {
85        self.m = m;
86        self
87    }
88
89    /// Number of candidates to be considered when searching for the nearest neighbors
90    /// during the construction of the graph.
91    ///
92    /// The default value is `100`.
93    pub fn ef_construction(mut self, ef_construction: usize) -> Self {
94        self.ef_construction = ef_construction;
95        self
96    }
97
98    /// Build the HNSW index from the given data.
99    ///
100    /// # Parameters
101    /// - `data`: A FixedSizeList to build the HNSW.
102    /// - `distance_type`: The distance type to use.
103    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/// Build a HNSW graph.
113///
114/// Currently, the HNSW graph is fully built in memory.
115///
116/// During the build, the graph is built layer by layer.
117///
118/// Each node in the graph has a global ID which is the index on the base layer.
119#[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        // Skipping the visited_generator_queue
138    }
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    /// Returns the metadata of this [`HNSW`].
334    pub fn metadata(&self) -> HnswMetadata {
335        // calculate the offsets of each level,
336        // start from 0
337        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        // Skipping the visited_generator_queue
374    }
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    /// Create a new [`HNSWBuilder`] with prepared params and in memory vector storage.
408    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    /// New node's level
453    ///
454    /// See paper `Algorithm 1`
455    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    /// Insert one node.
464    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        //
479        // Search for entry point in paper.
480        // ```
481        //   for l_c in (L..l+1) {
482        //     W = Search-Layer(q, ep, ef=1, l_c)
483        //    ep = Select-Neighbors(W, 1)
484        //  }
485        // ```
486        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(&current_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
577// View of a level in HNSW graph.
578// This is used to iterate over neighbors in a specific level.
579pub(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    /// Return the schema of the sub index
770    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, &params)
824        } else {
825            self.search_basic(query, k, &params, prefilter_bitset, storage)?
826        };
827        // if the queue is full, we just don't push it back, so ignore the error here
828        let _ = self.inner.visited_generator_queue.push(prefilter_generator);
829
830        // need to unique by row ids in case of searching multivector
831        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    /// Given a vector storage, containing all the data for the IVF partition, build the sub index.
843    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>>, // we don't need the mapping here because we rebuild the graph from remapped storage
878        store: &impl VectorStore,
879    ) -> Result<Self> {
880        // We can't simply remap the row ids in the graph because the vectors are changed,
881        // so the graph needs to be rebuilt.
882        Self::index_vectors(store, self.inner.params.clone())
883    }
884
885    /// Encode the sub index into a record batch
886    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, &params, None, store.as_ref())
1019            .unwrap();
1020        let loaded_results = loaded_hnsw
1021            .search_basic(query, k, &params, None, store.as_ref())
1022            .unwrap();
1023        assert_eq!(builder_results, loaded_results);
1024    }
1025}