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 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#[derive(Debug, Clone, Serialize, Deserialize, DeepSizeOf)]
48pub struct HnswBuildParams {
49 pub max_level: u16,
51
52 pub m: usize,
54
55 pub ef_construction: usize,
57
58 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 pub fn max_level(mut self, max_level: u16) -> Self {
77 self.max_level = max_level;
78 self
79 }
80
81 pub fn num_edges(mut self, m: usize) -> Self {
84 self.m = m;
85 self
86 }
87
88 pub fn ef_construction(mut self, ef_construction: usize) -> Self {
93 self.ef_construction = ef_construction;
94 self
95 }
96
97 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#[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 pub fn metadata(&self) -> HnswMetadata {
303 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 }
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 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 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 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 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(¤t_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 }
523
524 *neighbors_ranked = select_neighbors_heuristic(storage, &level_neighbors, m_max);
525 builder_node.update_from_ranked_neighbors(level);
526 }
527}
528
529pub(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 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, ¶ms)
736 } else {
737 self.search_basic(query, k, ¶ms, prefilter_bitset, storage)?
738 };
739 let _ = self.inner.visited_generator_queue.push(prefilter_generator);
741
742 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 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>>, store: &impl VectorStore,
794 ) -> Result<Self> {
795 Self::index_vectors(store, self.inner.params.clone())
798 }
799
800 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, ¶ms, None, store.as_ref())
935 .unwrap();
936 let loaded_results = loaded_hnsw
937 .search_basic(query, k, ¶ms, None, store.as_ref())
938 .unwrap();
939 assert_eq!(builder_results, loaded_results);
940 }
941}