urna_runtime/ann/
search.rs1use std::collections::{BinaryHeap, HashSet};
7
8use super::visited::VisitSet;
9use super::{Candidate, HnswIndex, Node, dist_q};
10use crate::materialize::PackedVectors;
11
12impl HnswIndex {
13 pub fn attach_vectors(&mut self, vectors: Vec<f32>) {
17 debug_assert_eq!(vectors.len(), self.n * self.dim);
18 self.store = PackedVectors::F32(vectors);
19 }
20
21 pub(crate) fn attach_store(&mut self, store: PackedVectors) {
25 self.store = store;
26 }
27
28 pub fn level0_neighbors(&self, i: usize) -> &[u32] {
34 self.nodes
35 .get(i)
36 .and_then(|node| node.neighbors.first())
37 .map(|v| v.as_slice())
38 .unwrap_or(&[])
39 }
40
41 pub fn search(&self, q: &[f32], ef: usize) -> Vec<usize> {
46 if self.n == 0 {
47 return Vec::new();
48 }
49 if !self.store.is_attached() {
50 return Vec::new();
53 }
54 let mut curr = self.entry_point;
55 for layer in (1..=self.max_level).rev() {
56 curr = greedy_search(curr, q, layer, &self.nodes, &self.store, self.dim);
57 }
58 let mut visited: HashSet<u32> = HashSet::new();
59 let candidates = layer_search(
60 &[curr],
61 q,
62 0,
63 ef.max(self.ef_search),
64 &self.nodes,
65 &self.store,
66 self.dim,
67 u32::MAX,
68 &mut visited,
69 );
70 candidates.into_iter().map(|c| c.id as usize).collect()
71 }
72}
73
74pub(super) fn greedy_search(
76 entry: u32,
77 q: &[f32],
78 layer: u32,
79 nodes: &[Node],
80 store: &PackedVectors,
81 dim: usize,
82) -> u32 {
83 let mut scratch = store.scratch(dim);
84 let mut curr = entry;
85 let mut curr_dist = dist_q(store, q, curr as usize, dim, &mut scratch);
86 loop {
87 let layer_idx = layer as usize;
88 if layer_idx >= nodes[curr as usize].neighbors.len() {
89 return curr;
90 }
91 let nbrs = &nodes[curr as usize].neighbors[layer_idx];
92 let mut best = curr;
93 let mut best_dist = curr_dist;
94 for &nbr in nbrs {
95 let d = dist_q(store, q, nbr as usize, dim, &mut scratch);
96 if d < best_dist {
97 best = nbr;
98 best_dist = d;
99 }
100 }
101 if best == curr {
102 return curr;
103 }
104 curr = best;
105 curr_dist = best_dist;
106 }
107}
108
109#[allow(clippy::too_many_arguments)]
112pub(super) fn layer_search(
113 entries: &[u32],
114 q: &[f32],
115 layer: u32,
116 ef: usize,
117 nodes: &[Node],
118 store: &PackedVectors,
119 dim: usize,
120 skip_id: u32,
121 visited: &mut impl VisitSet,
122) -> Vec<Candidate> {
123 let mut scratch = store.scratch(dim);
124 let mut frontier: BinaryHeap<ByDistAsc> = BinaryHeap::new();
128 let mut result: BinaryHeap<ByDistDesc> = BinaryHeap::new();
129
130 for &e in entries {
131 if e == skip_id {
132 continue;
133 }
134 let d = dist_q(store, q, e as usize, dim, &mut scratch);
135 let c = Candidate { id: e, dist: d };
136 frontier.push(ByDistAsc(c));
137 result.push(ByDistDesc(c));
138 visited.insert(e);
139 }
140
141 while let Some(ByDistAsc(curr)) = frontier.pop() {
142 let worst_in_result = result.peek().map(|r| r.0.dist).unwrap_or(f32::INFINITY);
143 if curr.dist > worst_in_result && result.len() >= ef {
144 break;
145 }
146 let layer_idx = layer as usize;
147 if layer_idx >= nodes[curr.id as usize].neighbors.len() {
148 continue;
149 }
150 let nbrs = &nodes[curr.id as usize].neighbors[layer_idx];
151 for &nbr in nbrs {
152 if nbr == skip_id || !visited.insert(nbr) {
153 continue;
154 }
155 let d = dist_q(store, q, nbr as usize, dim, &mut scratch);
156 let worst = result.peek().map(|r| r.0.dist).unwrap_or(f32::INFINITY);
157 if result.len() < ef || d < worst {
158 let c = Candidate { id: nbr, dist: d };
159 frontier.push(ByDistAsc(c));
160 result.push(ByDistDesc(c));
161 if result.len() > ef {
162 result.pop();
163 }
164 }
165 }
166 }
167
168 let mut out: Vec<Candidate> = result.into_iter().map(|w| w.0).collect();
169 out.sort_by(|a, b| crate::order::cmp_dist_asc(a.dist, b.dist));
170 out
171}
172
173#[derive(Clone, Copy)]
176struct ByDistAsc(Candidate);
177#[derive(Clone, Copy)]
178struct ByDistDesc(Candidate);
179
180impl PartialEq for ByDistAsc {
181 fn eq(&self, other: &Self) -> bool {
182 self.0.dist == other.0.dist
183 }
184}
185impl Eq for ByDistAsc {}
186impl Ord for ByDistAsc {
187 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
188 other
190 .0
191 .dist
192 .partial_cmp(&self.0.dist)
193 .unwrap_or(std::cmp::Ordering::Equal)
194 }
195}
196impl PartialOrd for ByDistAsc {
197 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
198 Some(self.cmp(other))
199 }
200}
201
202impl PartialEq for ByDistDesc {
203 fn eq(&self, other: &Self) -> bool {
204 self.0.dist == other.0.dist
205 }
206}
207impl Eq for ByDistDesc {}
208impl Ord for ByDistDesc {
209 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
210 self.0
211 .dist
212 .partial_cmp(&other.0.dist)
213 .unwrap_or(std::cmp::Ordering::Equal)
214 }
215}
216impl PartialOrd for ByDistDesc {
217 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
218 Some(self.cmp(other))
219 }
220}