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}