1use std::cell::RefCell;
12use std::cmp::{Ordering, Reverse};
13use std::collections::BinaryHeap;
14use std::ops::Range;
15use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering as AtomicOrdering};
16use std::sync::{Mutex, MutexGuard, PoisonError};
17
18use crate::error::{Error, Result};
19use crate::metric::Metric;
20use crate::rng::SplitMix64;
21
22pub(crate) const MAX_LEVEL: usize = 16;
24
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
27pub struct HnswParams {
28 pub m: usize,
30 pub ef_construction: usize,
33 pub ef_search: usize,
36}
37
38impl Default for HnswParams {
39 fn default() -> Self {
40 Self {
41 m: 16,
42 ef_construction: 200,
43 ef_search: 64,
44 }
45 }
46}
47
48impl HnswParams {
49 pub(crate) fn validate(&self) -> Result<()> {
50 if !(2..=256).contains(&self.m) {
51 return Err(Error::InvalidArgument("m must be between 2 and 256".into()));
52 }
53 if self.ef_construction == 0 || self.ef_search == 0 {
54 return Err(Error::InvalidArgument(
55 "ef values must be greater than zero".into(),
56 ));
57 }
58 Ok(())
59 }
60
61 pub(crate) fn max_links(&self, layer: usize) -> usize {
62 if layer == 0 { self.m * 2 } else { self.m }
63 }
64}
65
66pub(crate) struct Vectors {
68 pub(crate) dim: usize,
69 pub(crate) data: Vec<f32>,
70}
71
72impl Vectors {
73 pub(crate) fn new(dim: usize) -> Self {
74 Self {
75 dim,
76 data: Vec::new(),
77 }
78 }
79
80 #[inline]
81 pub(crate) fn get(&self, node: u32) -> &[f32] {
82 let start = node as usize * self.dim;
83 &self.data[start..start + self.dim]
84 }
85
86 pub(crate) fn push(&mut self, vector: &[f32]) {
87 debug_assert_eq!(vector.len(), self.dim);
88 self.data.extend_from_slice(vector);
89 }
90}
91
92#[derive(Clone, Copy, Debug)]
93pub(crate) struct Candidate {
94 pub(crate) dist: f32,
95 pub(crate) id: u32,
96}
97
98impl PartialEq for Candidate {
99 fn eq(&self, other: &Self) -> bool {
100 self.cmp(other) == Ordering::Equal
101 }
102}
103
104impl Eq for Candidate {}
105
106impl PartialOrd for Candidate {
107 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
108 Some(self.cmp(other))
109 }
110}
111
112impl Ord for Candidate {
113 fn cmp(&self, other: &Self) -> Ordering {
114 self.dist
115 .total_cmp(&other.dist)
116 .then(self.id.cmp(&other.id))
117 }
118}
119
120#[derive(Clone, Copy, Debug, Default)]
122pub(crate) struct Counters {
123 pub(crate) visited: usize,
124 pub(crate) distance_computations: usize,
125}
126
127#[derive(Default)]
131pub(crate) struct Visited {
132 bits: Vec<u64>,
133 touched: Vec<u32>,
134 fresh: Vec<u32>,
136}
137
138impl Visited {
139 fn reset(&mut self, len: usize) {
140 for word in self.touched.drain(..) {
141 self.bits[word as usize] = 0;
142 }
143 let words = len.div_ceil(64);
144 if self.bits.len() < words {
145 self.bits.resize(words, 0);
146 }
147 }
148
149 #[inline]
151 fn insert(&mut self, node: u32) -> bool {
152 let (word, bit) = (node as usize / 64, 1u64 << (node % 64));
153 let bits = &mut self.bits[word];
154 if *bits & bit != 0 {
155 return false;
156 }
157 if *bits == 0 {
158 self.touched.push(word as u32);
159 }
160 *bits |= bit;
161 true
162 }
163}
164
165thread_local! {
166 static VISITED: RefCell<Visited> = RefCell::new(Visited::default());
167}
168
169pub(crate) trait ReadLinks {
171 fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R;
172
173 fn is_ready(&self, _node: u32) -> bool {
176 true
177 }
178}
179
180pub(crate) trait Links: ReadLinks {
182 fn set_neighbors(
185 &mut self,
186 node: u32,
187 layer: usize,
188 neighbors: Vec<u32>,
189 max: usize,
190 vectors: &Vectors,
191 metric: Metric,
192 );
193
194 fn add_link(
196 &mut self,
197 node: u32,
198 layer: usize,
199 new: u32,
200 max: usize,
201 vectors: &Vectors,
202 metric: Metric,
203 );
204}
205
206impl ReadLinks for [Vec<Vec<u32>>] {
207 #[inline]
208 fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
209 f(&self[node as usize][layer])
210 }
211}
212
213struct Exclusive<'a>(&'a mut [Vec<Vec<u32>>]);
215
216impl ReadLinks for Exclusive<'_> {
217 fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
218 self.0.with_neighbors(node, layer, f)
219 }
220}
221
222impl Links for Exclusive<'_> {
223 fn set_neighbors(
224 &mut self,
225 node: u32,
226 layer: usize,
227 neighbors: Vec<u32>,
228 _: usize,
229 _: &Vectors,
230 _: Metric,
231 ) {
232 self.0[node as usize][layer] = neighbors;
234 }
235
236 fn add_link(
237 &mut self,
238 node: u32,
239 layer: usize,
240 new: u32,
241 max: usize,
242 vectors: &Vectors,
243 metric: Metric,
244 ) {
245 push_and_prune(
246 &mut self.0[node as usize][layer],
247 node,
248 new,
249 max,
250 vectors,
251 metric,
252 );
253 }
254}
255
256#[derive(Clone, Copy)]
258struct Shared<'a> {
259 links: &'a [Mutex<Vec<Vec<u32>>>],
260 ready: &'a [AtomicBool],
261}
262
263impl Shared<'_> {
264 fn lock(&self, node: u32) -> MutexGuard<'_, Vec<Vec<u32>>> {
265 self.links[node as usize]
266 .lock()
267 .unwrap_or_else(PoisonError::into_inner)
268 }
269}
270
271impl ReadLinks for Shared<'_> {
272 fn with_neighbors<R>(&self, node: u32, layer: usize, f: impl FnOnce(&[u32]) -> R) -> R {
273 f(&self.lock(node)[layer])
274 }
275
276 fn is_ready(&self, node: u32) -> bool {
277 self.ready[node as usize].load(AtomicOrdering::Acquire)
278 }
279}
280
281impl Links for Shared<'_> {
282 fn set_neighbors(
283 &mut self,
284 node: u32,
285 layer: usize,
286 neighbors: Vec<u32>,
287 max: usize,
288 vectors: &Vectors,
289 metric: Metric,
290 ) {
291 let mut links = self.lock(node);
292 let list = &mut links[layer];
293 if list.is_empty() {
294 *list = neighbors;
295 return;
296 }
297 for neighbor in neighbors {
298 if !list.contains(&neighbor) {
299 list.push(neighbor);
300 }
301 }
302 prune(list, node, max, vectors, metric);
303 }
304
305 fn add_link(
306 &mut self,
307 node: u32,
308 layer: usize,
309 new: u32,
310 max: usize,
311 vectors: &Vectors,
312 metric: Metric,
313 ) {
314 push_and_prune(&mut self.lock(node)[layer], node, new, max, vectors, metric);
315 }
316}
317
318pub(crate) struct Graph {
319 pub(crate) params: HnswParams,
320 pub(crate) links: Vec<Vec<Vec<u32>>>,
322 pub(crate) entry: Option<u32>,
323 pub(crate) max_level: usize,
324 pub(crate) rng: SplitMix64,
325}
326
327impl Graph {
328 pub(crate) fn new(params: HnswParams, seed: u64) -> Self {
329 Self {
330 params,
331 links: Vec::new(),
332 entry: None,
333 max_level: 0,
334 rng: SplitMix64::new(seed),
335 }
336 }
337
338 pub(crate) fn len(&self) -> usize {
339 self.links.len()
340 }
341
342 fn random_level(&mut self) -> usize {
343 let ml = 1.0 / (self.params.m as f64).ln();
344 let level = (-self.rng.next_f64().ln() * ml).floor() as usize;
345 level.min(MAX_LEVEL)
346 }
347
348 pub(crate) fn insert(&mut self, id: u32, vectors: &Vectors, metric: Metric) {
351 debug_assert_eq!(id as usize, self.links.len());
352 self.insert_batch(id..id + 1, vectors, metric, 1);
353 }
354
355 pub(crate) fn insert_batch(
360 &mut self,
361 nodes: Range<u32>,
362 vectors: &Vectors,
363 metric: Metric,
364 threads: usize,
365 ) {
366 debug_assert_eq!(nodes.start as usize, self.links.len());
367 let levels: Vec<usize> = nodes.clone().map(|_| self.random_level()).collect();
368 for &level in &levels {
369 self.links.push(vec![Vec::new(); level + 1]);
370 }
371 let mut rest = nodes.clone();
372 if self.entry.is_none() {
373 let Some(first) = rest.next() else { return };
374 self.entry = Some(first);
375 self.max_level = levels[0];
376 }
377 let level_of = |node: u32| levels[(node - nodes.start) as usize];
378 let params = self.params;
379
380 if threads <= 1 || rest.len() < 1000 {
382 VISITED.with_borrow_mut(|visited| {
383 for node in rest {
384 let level = level_of(node);
385 let entry = self.entry.expect("entry is set above");
386 let mut links = Exclusive(&mut self.links);
387 link_node(
388 &mut links,
389 ¶ms,
390 node,
391 level,
392 entry,
393 self.max_level,
394 vectors,
395 metric,
396 visited,
397 );
398 if level > self.max_level {
399 self.max_level = level;
400 self.entry = Some(node);
401 }
402 }
403 });
404 return;
405 }
406
407 let locked: Vec<Mutex<Vec<Vec<u32>>>> = std::mem::take(&mut self.links)
408 .into_iter()
409 .map(Mutex::new)
410 .collect();
411 let top = Mutex::new((self.entry.expect("entry is set above"), self.max_level));
412 let next = AtomicUsize::new(rest.start as usize);
413 let ready: Vec<AtomicBool> = (0..locked.len())
414 .map(|node| AtomicBool::new(node < rest.start as usize))
415 .collect();
416 std::thread::scope(|scope| {
417 for _ in 0..threads {
418 scope.spawn(|| {
419 let mut links = Shared {
420 links: &locked,
421 ready: &ready,
422 };
423 let mut visited = Visited::default();
424 loop {
425 let node = next.fetch_add(1, AtomicOrdering::Relaxed);
426 if node >= rest.end as usize {
427 break;
428 }
429 let (node, level) = (node as u32, level_of(node as u32));
430 let mut top = top.lock().unwrap_or_else(PoisonError::into_inner);
431 let (entry, max_level) = *top;
432 if level > max_level {
433 link_node(
436 &mut links,
437 ¶ms,
438 node,
439 level,
440 entry,
441 max_level,
442 vectors,
443 metric,
444 &mut visited,
445 );
446 *top = (node, level);
447 } else {
448 drop(top);
449 link_node(
450 &mut links,
451 ¶ms,
452 node,
453 level,
454 entry,
455 max_level,
456 vectors,
457 metric,
458 &mut visited,
459 );
460 }
461 ready[node as usize].store(true, AtomicOrdering::Release);
462 }
463 });
464 }
465 });
466 self.links = locked
467 .into_iter()
468 .map(|m| m.into_inner().unwrap_or_else(PoisonError::into_inner))
469 .collect();
470 let (entry, max_level) = top.into_inner().unwrap_or_else(PoisonError::into_inner);
471 self.entry = Some(entry);
472 self.max_level = max_level;
473 self.repair_unreachable(vectors, metric);
474 }
475
476 pub(crate) fn repair_unreachable(&mut self, vectors: &Vectors, metric: Metric) -> usize {
480 let reachable = self.reachable();
481 let orphans: Vec<u32> = (0..self.links.len() as u32)
482 .filter(|&node| !reachable[node as usize])
483 .collect();
484 let Some(entry) = self.entry else { return 0 };
485 let max = self.params.max_links(0);
486 VISITED.with_borrow_mut(|visited| {
487 for &orphan in &orphans {
488 let query = vectors.get(orphan);
489 let found = {
490 let links: &[Vec<Vec<u32>>] = &self.links;
491 let mut counters = Counters::default();
492 let mut ep = Candidate {
493 dist: metric.distance(query, vectors.get(entry)),
494 id: entry,
495 };
496 for layer in (1..=self.max_level).rev() {
497 ep = greedy(links, query, ep, layer, vectors, metric, &mut counters);
498 }
499 let ef = self.params.ef_construction;
500 let accept = |n: u32| n != orphan;
501 search_layer(
502 links,
503 query,
504 &[ep],
505 ef,
506 0,
507 vectors,
508 metric,
509 &accept,
510 &mut counters,
511 visited,
512 )
513 };
514 let mut linked = false;
515 for candidate in found.iter().take(self.params.m) {
516 let list = &mut self.links[candidate.id as usize][0];
517 push_and_prune(list, candidate.id, orphan, max, vectors, metric);
518 linked |= list.contains(&orphan);
519 }
520 if !linked {
521 if let Some(nearest) = found.first() {
522 self.links[nearest.id as usize][0].push(orphan);
524 }
525 }
526 }
527 });
528 orphans.len()
529 }
530
531 #[allow(clippy::too_many_arguments)]
533 pub(crate) fn search(
534 &self,
535 query: &[f32],
536 k: usize,
537 ef: usize,
538 vectors: &Vectors,
539 metric: Metric,
540 accept: &dyn Fn(u32) -> bool,
541 counters: &mut Counters,
542 ) -> Vec<Candidate> {
543 let Some(entry) = self.entry else {
544 return Vec::new();
545 };
546 let links: &[Vec<Vec<u32>>] = &self.links;
547 counters.distance_computations += 1;
548 let mut ep = Candidate {
549 dist: metric.distance(query, vectors.get(entry)),
550 id: entry,
551 };
552 for layer in (1..=self.max_level).rev() {
553 ep = greedy(links, query, ep, layer, vectors, metric, counters);
554 }
555 let mut found = VISITED.with_borrow_mut(|visited| {
556 search_layer(
557 links,
558 query,
559 &[ep],
560 ef.max(k),
561 0,
562 vectors,
563 metric,
564 accept,
565 counters,
566 visited,
567 )
568 });
569 found.truncate(k);
570 found
571 }
572
573 pub(crate) fn nodes_per_layer(&self) -> Vec<usize> {
575 if self.links.is_empty() {
576 return Vec::new();
577 }
578 let mut counts = vec![0; self.max_level + 1];
579 for layers in &self.links {
580 for count in counts.iter_mut().take(layers.len()) {
581 *count += 1;
582 }
583 }
584 counts
585 }
586
587 pub(crate) fn reachable(&self) -> Vec<bool> {
589 let mut seen = vec![false; self.links.len()];
590 let Some(entry) = self.entry else { return seen };
591 let mut stack = vec![entry];
592 seen[entry as usize] = true;
593 while let Some(node) = stack.pop() {
594 for &neighbor in &self.links[node as usize][0] {
595 if !seen[neighbor as usize] {
596 seen[neighbor as usize] = true;
597 stack.push(neighbor);
598 }
599 }
600 }
601 seen
602 }
603}
604
605#[allow(clippy::too_many_arguments)]
609fn link_node<L: Links>(
610 links: &mut L,
611 params: &HnswParams,
612 node: u32,
613 level: usize,
614 entry: u32,
615 max_level: usize,
616 vectors: &Vectors,
617 metric: Metric,
618 visited: &mut Visited,
619) {
620 let query = vectors.get(node);
621 let mut counters = Counters::default();
622 let mut ep = Candidate {
623 dist: metric.distance(query, vectors.get(entry)),
624 id: entry,
625 };
626 for layer in (level + 1..=max_level).rev() {
627 ep = greedy(&*links, query, ep, layer, vectors, metric, &mut counters);
628 }
629
630 let mut entry_points = vec![ep];
631 for layer in (0..=level.min(max_level)).rev() {
632 let found = search_layer(
633 &*links,
634 query,
635 &entry_points,
636 params.ef_construction,
637 layer,
638 vectors,
639 metric,
640 &|n| n != node,
641 &mut counters,
642 visited,
643 );
644 let max = params.max_links(layer);
645 let neighbors = select_neighbors(&found, max, vectors, metric);
646 links.set_neighbors(node, layer, neighbors.clone(), max, vectors, metric);
647 for neighbor in neighbors {
648 links.add_link(neighbor, layer, node, max, vectors, metric);
649 }
650 entry_points = found;
651 }
652}
653
654#[inline(always)]
657fn prefetch(vector: &[f32]) {
658 for line in vector.chunks(16) {
659 let ptr = line.as_ptr();
660 #[cfg(target_arch = "aarch64")]
661 unsafe {
664 std::arch::asm!("prfm pldl1keep, [{0}]", in(reg) ptr, options(nostack, preserves_flags, readonly));
665 }
666 #[cfg(target_arch = "x86_64")]
667 unsafe {
669 std::arch::x86_64::_mm_prefetch(ptr.cast::<i8>(), std::arch::x86_64::_MM_HINT_T0);
670 }
671 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
672 let _ = ptr;
673 }
674}
675
676fn greedy<R: ReadLinks + ?Sized>(
677 links: &R,
678 query: &[f32],
679 mut best: Candidate,
680 layer: usize,
681 vectors: &Vectors,
682 metric: Metric,
683 counters: &mut Counters,
684) -> Candidate {
685 loop {
686 let current = best.id;
687 links.with_neighbors(current, layer, |neighbors| {
688 for &neighbor in neighbors {
689 prefetch(vectors.get(neighbor));
690 }
691 for &neighbor in neighbors {
692 counters.distance_computations += 1;
693 let dist = metric.distance(query, vectors.get(neighbor));
694 if dist < best.dist && links.is_ready(neighbor) {
697 best = Candidate { dist, id: neighbor };
698 }
699 }
700 });
701 if best.id == current {
702 return best;
703 }
704 }
705}
706
707#[allow(clippy::too_many_arguments)]
711fn search_layer<R: ReadLinks + ?Sized>(
712 links: &R,
713 query: &[f32],
714 entry_points: &[Candidate],
715 ef: usize,
716 layer: usize,
717 vectors: &Vectors,
718 metric: Metric,
719 accept: &dyn Fn(u32) -> bool,
720 counters: &mut Counters,
721 visited: &mut Visited,
722) -> Vec<Candidate> {
723 visited.reset(vectors.data.len() / vectors.dim.max(1));
724 let mut candidates = BinaryHeap::new();
725 let mut results: BinaryHeap<Candidate> = BinaryHeap::new();
726
727 for &ep in entry_points {
728 if !visited.insert(ep.id) {
729 continue;
730 }
731 counters.visited += 1;
732 candidates.push(Reverse(ep));
733 if accept(ep.id) {
734 results.push(ep);
735 if results.len() > ef {
736 results.pop();
737 }
738 }
739 }
740
741 while let Some(Reverse(current)) = candidates.pop() {
742 if results.len() >= ef && results.peek().is_some_and(|w| current.dist > w.dist) {
743 break;
744 }
745 let mut fresh = std::mem::take(&mut visited.fresh);
748 fresh.clear();
749 links.with_neighbors(current.id, layer, |neighbors| {
750 for &neighbor in neighbors {
751 if visited.insert(neighbor) {
752 fresh.push(neighbor);
753 }
754 }
755 });
756 for &neighbor in &fresh {
757 prefetch(vectors.get(neighbor));
758 }
759 for &neighbor in &fresh {
760 counters.visited += 1;
761 counters.distance_computations += 1;
762 let dist = metric.distance(query, vectors.get(neighbor));
763 let worst = results.peek().map_or(f32::INFINITY, |w| w.dist);
764 if results.len() < ef || dist < worst {
765 let candidate = Candidate { dist, id: neighbor };
766 candidates.push(Reverse(candidate));
767 if accept(neighbor) {
768 results.push(candidate);
769 if results.len() > ef {
770 results.pop();
771 }
772 }
773 }
774 }
775 visited.fresh = fresh;
776 }
777
778 let mut out = results.into_vec();
779 out.sort();
780 out
781}
782
783fn select_neighbors(
788 candidates: &[Candidate],
789 max: usize,
790 vectors: &Vectors,
791 metric: Metric,
792) -> Vec<u32> {
793 let mut selected: Vec<Candidate> = Vec::with_capacity(max);
794 let mut pruned = Vec::new();
795 for &candidate in candidates {
796 if selected.len() >= max {
797 break;
798 }
799 let vector = vectors.get(candidate.id);
800 let diverse = selected
801 .iter()
802 .all(|s| metric.distance(vector, vectors.get(s.id)) > candidate.dist);
803 if diverse {
804 selected.push(candidate);
805 } else {
806 pruned.push(candidate);
807 }
808 }
809 for candidate in pruned {
810 if selected.len() >= max {
811 break;
812 }
813 selected.push(candidate);
814 }
815 selected.into_iter().map(|c| c.id).collect()
816}
817
818fn push_and_prune(
819 list: &mut Vec<u32>,
820 node: u32,
821 new: u32,
822 max: usize,
823 vectors: &Vectors,
824 metric: Metric,
825) {
826 list.push(new);
827 prune(list, node, max, vectors, metric);
828}
829
830fn prune(list: &mut Vec<u32>, node: u32, max: usize, vectors: &Vectors, metric: Metric) {
832 if list.len() <= max {
833 return;
834 }
835 let base = vectors.get(node);
836 let mut candidates: Vec<Candidate> = list
837 .iter()
838 .map(|&n| Candidate {
839 dist: metric.distance(base, vectors.get(n)),
840 id: n,
841 })
842 .collect();
843 candidates.sort();
844 *list = select_neighbors(&candidates, max, vectors, metric);
845}
846
847#[cfg(test)]
848mod tests {
849 use super::*;
850
851 fn random_graph(n: usize, dim: usize) -> (Graph, Vectors) {
852 let mut rng = SplitMix64::new(9);
853 let mut vectors = Vectors::new(dim);
854 for _ in 0..n {
855 let v: Vec<f32> = (0..dim).map(|_| rng.next_f64() as f32).collect();
856 vectors.push(&v);
857 }
858 let mut graph = Graph::new(HnswParams::default(), 1);
859 graph.insert_batch(0..n as u32, &vectors, Metric::L2, 1);
860 (graph, vectors)
861 }
862
863 #[test]
864 fn repair_relinks_orphaned_nodes() {
865 let (mut graph, vectors) = random_graph(500, 8);
866 let orphan = (0..500u32).find(|&n| Some(n) != graph.entry).unwrap();
867 for layers in &mut graph.links {
868 for list in layers.iter_mut() {
869 list.retain(|&n| n != orphan);
870 }
871 }
872 assert!(!graph.reachable()[orphan as usize]);
873 assert_eq!(graph.repair_unreachable(&vectors, Metric::L2), 1);
874 assert!(graph.reachable().iter().all(|&r| r));
875 assert_eq!(graph.repair_unreachable(&vectors, Metric::L2), 0);
876 }
877
878 #[test]
879 fn visited_set_resets_between_searches() {
880 let mut visited = Visited::default();
881 visited.reset(200);
882 assert!(visited.insert(3) && visited.insert(130));
883 assert!(!visited.insert(3));
884 visited.reset(200);
885 assert!(visited.insert(3) && visited.insert(130));
886 }
887}