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