Skip to main content

urna_runtime/ann/
mod.rs

1//! Pure-Rust HNSW index for approximate nearest neighbor search.
2//!
3//! This is a minimal, self-contained HNSW implementation tailored to
4//! `.urna`'s contract:
5//!
6//! - Vectors are L2-normalized → distance is `1 - cosine` (smaller = closer).
7//! - Index lives in section `0x07` and is bit-equal across rebuilds.
8//! - Search returns a candidate set; the runtime reranks with the exact
9//!   dot product against the embeddings section so the final score is
10//!   the real cosine.
11//!
12//! On-disk layout (`encoding=raw`, payload version 1):
13//!
14//! ```text
15//!   u32 LE  payload_version = 1
16//!   u32 LE  m                 - out-degree at non-zero levels
17//!   u32 LE  m_max0            - out-degree at level 0 (typically 2*m)
18//!   u32 LE  ef_construction
19//!   u32 LE  entry_point       - node id of the entry vertex
20//!   u32 LE  max_level         - highest layer with any node (0-based)
21//!   u32 LE  n_nodes           - equal to header.n_embeddings
22//!   for each node i in 0..n_nodes:
23//!       u32 LE  level_i       - top layer this node lives in
24//!       for layer in 0..=level_i:
25//!           u32 LE  k_i_l     - neighbor count at this layer
26//!           u32 LE * k_i_l    - neighbor ids
27//! ```
28//!
29//! Construction uses HNSW (Malkov & Yashunin, 2018) with a deterministic
30//! level distribution so the same input produces the same graph.
31//! Neighbor selection lives in `select_neighbors`; today it uses the
32//! `_simple` variant (top-m by distance), Phase 2 swaps in the
33//! Algorithm 4 heuristic for higher recall.
34
35mod build;
36mod codec;
37mod search;
38pub mod select_neighbors;
39mod visited;
40
41use crate::materialize::PackedVectors;
42
43/// on-disk payload version for the hnsw section (`0x07`). v1 stored every
44/// neighbour id as a raw u32; v2 bitpacks the level/count/neighbour columns
45/// with `intpack` (order-preserving, so the graph and its recall are
46/// unchanged). the reader still accepts v1 files. the section is optional
47/// and excluded from content_hash, so this bump is additive within v1.
48pub const HNSW_PAYLOAD_VERSION: u32 = 2;
49
50/// Default neighbor count at non-zero layers. 16 is a common HNSW sweet
51/// spot for ~1M points; for smaller corpora the recall-vs-size curve is
52/// flat enough that the default is fine.
53pub const DEFAULT_M: usize = 16;
54/// Default candidate-list size during construction. Larger = better
55/// recall, slower build. 400 is our chosen production default -
56/// empirically gives recall@10 ≥ 0.95 at typical corpus sizes
57/// (n ≤ 100k, dim ≤ 768) when paired with `ef_search ≥ 400`. Lower
58/// values save build time but require larger `ef_search` to match.
59pub const DEFAULT_EF_CONSTRUCTION: usize = 400;
60
61#[derive(Clone, Debug)]
62pub(super) struct Node {
63    /// Top layer this node lives in (0-based).
64    pub level: u32,
65    /// `neighbors[layer][i]` is the i-th neighbor id at `layer`. Index 0
66    /// is the densest layer (level 0).
67    pub neighbors: Vec<Vec<u32>>,
68}
69
70/// A built HNSW index. Reads borrow from the on-disk payload at open
71/// time; the graph is owned (small relative to embeddings).
72pub struct HnswIndex {
73    pub m: usize,
74    pub m_max0: usize,
75    pub ef_construction: usize,
76    pub entry_point: u32,
77    pub max_level: u32,
78    pub(super) nodes: Vec<Node>,
79    /// The vectors used at search time, kept in their on-disk packing
80    /// (int8 stays int8 + scales, f16 stays f16) and decoded one row at a
81    /// time. The graph stays dtype-independent (f16/i8 runtimes get the
82    /// same recall curve) without the old `n*dim*4` f32 snapshot: the
83    /// resident footprint is the packed size, not 4x it.
84    pub(super) store: PackedVectors,
85    pub(super) dim: usize,
86    pub(super) n: usize,
87    /// the search beam floor, `ef_construction` at decode time: a caller's
88    /// `ef` widens the beam above it (`max(ef, ef_search)`) and never
89    /// narrows it, so the recall measured at build stays the floor too.
90    pub ef_search: usize,
91}
92
93#[derive(Clone, Copy, Debug, PartialEq)]
94pub(super) struct Candidate {
95    pub id: u32,
96    /// `1 - cosine`. Smaller = closer.
97    pub dist: f32,
98}
99
100impl Eq for Candidate {}
101
102/// `1 - dot` over two l2-normalized rows, with eight independent
103/// accumulators combined in a fixed tree. plain rust on purpose: the
104/// compiler vectorizes the eight lanes on every target (sse2 on x86_64,
105/// neon on aarch64) and, without fma contraction, the arithmetic is the
106/// same on all of them, so a graph built on one machine is byte-identical
107/// to the same build on another. the simd kernels in `crate::simd` are
108/// NOT used here: their reduction trees differ per backend, which would
109/// make the graph bytes depend on the cpu that built them.
110#[inline]
111pub(super) fn cosine_dist(a: &[f32], b: &[f32]) -> f32 {
112    let mut acc = [0.0f32; 8];
113    let ca = a.chunks_exact(8);
114    let cb = b.chunks_exact(8);
115    let (ra, rb) = (ca.remainder(), cb.remainder());
116    for (x, y) in ca.zip(cb) {
117        for j in 0..8 {
118            acc[j] += x[j] * y[j];
119        }
120    }
121    let mut tail = 0.0f32;
122    for (x, y) in ra.iter().zip(rb.iter()) {
123        tail += x * y;
124    }
125    let dot =
126        ((acc[0] + acc[1]) + (acc[2] + acc[3])) + ((acc[4] + acc[5]) + (acc[6] + acc[7])) + tail;
127    1.0 - dot
128}
129
130/// Distance between an f32 query and stored row `i`, decoding `i` through
131/// `store` into `scratch` first. `scratch` is empty for the f32 store.
132#[inline]
133pub(super) fn dist_q(
134    store: &PackedVectors,
135    q: &[f32],
136    i: usize,
137    dim: usize,
138    scratch: &mut [f32],
139) -> f32 {
140    cosine_dist(q, store.row(i, dim, scratch))
141}
142
143/// Distance between two stored rows `a` and `b`, each decoded through
144/// `store` into its own scratch buffer (`sa`, `sb`).
145#[inline]
146pub(super) fn dist_rr(
147    store: &PackedVectors,
148    a: usize,
149    b: usize,
150    dim: usize,
151    sa: &mut [f32],
152    sb: &mut [f32],
153) -> f32 {
154    cosine_dist(store.row(a, dim, sa), store.row(b, dim, sb))
155}
156
157#[cfg(test)]
158mod tests {
159    use super::*;
160    use crate::ann::build::LcgRng;
161    use std::collections::HashSet;
162
163    pub(super) fn random_vectors(n: usize, dim: usize, seed: u64) -> Vec<f32> {
164        let mut rng = LcgRng::new(seed);
165        let mut v = Vec::with_capacity(n * dim);
166        for _ in 0..(n * dim) {
167            v.push((rng.next_f64() as f32) - 0.5);
168        }
169        // L2-normalize each row.
170        for i in 0..n {
171            let row = &mut v[i * dim..(i + 1) * dim];
172            let norm: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
173            if norm > 0.0 {
174                for x in row.iter_mut() {
175                    *x /= norm;
176                }
177            }
178        }
179        v
180    }
181
182    #[test]
183    fn small_index_recall_against_exact() {
184        // 200 random vectors, dim 32. Recall@10 vs exact should be very
185        // high - small enough that the graph is fully connected.
186        let n = 200;
187        let dim = 32;
188        let vecs = random_vectors(n, dim, 0xDEAD_BEEF);
189        let idx = HnswIndex::build(vecs.clone(), n, dim, 8, 50, 42);
190
191        let q = random_vectors(1, dim, 0xCAFEBABE);
192        // Exact top-10.
193        let mut exact: Vec<(usize, f32)> = (0..n)
194            .map(|i| {
195                let row = &vecs[i * dim..(i + 1) * dim];
196                let mut s = 0.0f32;
197                for j in 0..dim {
198                    s += q[j] * row[j];
199                }
200                (i, s)
201            })
202            .collect();
203        exact.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
204        let exact_top: HashSet<usize> = exact.iter().take(10).map(|p| p.0).collect();
205        let approx = idx.search(&q, 50);
206        let approx_set: HashSet<usize> = approx.into_iter().take(10).collect();
207        let overlap = exact_top.intersection(&approx_set).count();
208        assert!(
209            overlap >= 7,
210            "recall@10 too low: {} of 10 (expected >= 7)",
211            overlap
212        );
213    }
214
215    #[test]
216    fn serialize_roundtrip() {
217        let n = 50;
218        let dim = 16;
219        let vecs = random_vectors(n, dim, 7);
220        let idx = HnswIndex::build(vecs.clone(), n, dim, 8, 30, 42);
221        let bytes = idx.to_bytes();
222        let mut decoded = HnswIndex::from_bytes(&bytes, n, dim).unwrap();
223        decoded.attach_vectors(vecs);
224        // Same query should produce a similar candidate set.
225        let q: Vec<f32> = vec![1.0 / (dim as f32).sqrt(); dim];
226        let a = idx.search(&q, 20);
227        let b = decoded.search(&q, 20);
228        let a_set: HashSet<usize> = a.into_iter().collect();
229        let b_set: HashSet<usize> = b.into_iter().collect();
230        assert_eq!(a_set, b_set, "candidate sets must match after roundtrip");
231    }
232
233    #[test]
234    fn deterministic_build_same_seed() {
235        let n = 30;
236        let dim = 8;
237        let vecs = random_vectors(n, dim, 0xABCD);
238        let a = HnswIndex::build(vecs.clone(), n, dim, 4, 20, 123);
239        let b = HnswIndex::build(vecs, n, dim, 4, 20, 123);
240        assert_eq!(a.to_bytes(), b.to_bytes(), "same seed => same graph");
241    }
242}