Skip to main content

urna_runtime/ann/
search.rs

1//! HNSW search: greedy descent on upper layers, ef-bounded beam search
2//! on layer 0. Returns candidates sorted ascending by distance; the
3//! runtime reranks with the exact dot product to produce the final
4//! cosine score.
5
6use std::collections::{BinaryHeap, HashSet};
7
8use super::visited::VisitSet;
9use super::{Candidate, HnswIndex, Node, dist_q};
10use crate::materialize::PackedVectors;
11
12impl HnswIndex {
13    /// Attach f32 vectors as the search store. Kept for callers that
14    /// already hold an f32 buffer (build, tests); the runtime open path
15    /// uses `attach_store` to avoid an f32 expansion.
16    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    /// Attach a packed vector store at open time. Keeps int8/f16 rows in
22    /// their on-disk packing so the resident footprint is the packed size,
23    /// not the old `n*dim*4` f32 snapshot.
24    pub(crate) fn attach_store(&mut self, store: PackedVectors) {
25        self.store = store;
26    }
27
28    /// Level-0 (densest layer) out-neighbors of node `i`, or an empty slice
29    /// if `i` is out of range. Read-only accessor the graph build path uses to
30    /// derive top-m semantic edges from the already-built hnsw graph without
31    /// re-running an O(n^2) exact knn. The neighbor SET is the build-order
32    /// hnsw adjacency; the graph builder sorts canonically before writing.
33    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    /// Search for the `ef` closest candidates to `q`. Returns ids only -
42    /// the runtime reranks with the exact dot product to produce the
43    /// final cosine score. the beam is `max(ef, self.ef_search)`: `ef`
44    /// widens the search above the file's floor and never narrows it.
45    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            // Index is loaded but vectors haven't been attached. Return
51            // an empty candidate set; runtime falls back to exact.
52            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
74/// Greedy descent on a single layer until no neighbor is closer.
75pub(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/// Search a single layer with a candidate list of size `ef`. Returns
110/// the best `ef` candidates sorted ascending by distance.
111#[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    // BinaryHeap orderings:
125    //   `frontier` - min-heap by distance (closest first to expand).
126    //   `result`   - max-heap by distance (so we can prune the farthest).
127    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// `BinaryHeap` is a max-heap; we want closest-first / farthest-first
174// orderings. Define orderings explicitly.
175#[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        // Reverse so BinaryHeap (max-heap) pops smallest distance.
189        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}