Skip to main content

kevy_vector/
hnsw.rs

1//! HNSW graph (RFC D2/D5): hierarchical layers, greedy descent +
2//! beam search on layer 0, tombstone deletes filtered at search
3//! time, bounded full rebuild by re-inserting the living.
4
5use std::collections::{BinaryHeap, HashMap};
6
7use crate::dist::Distance;
8
9/// Construction/search parameters (immutable once built — RFC D2).
10#[derive(Debug, Clone, Copy)]
11pub struct HnswParams {
12    /// Max bidirectional links per node per layer (layer 0 gets 2M).
13    pub m: usize,
14    /// Construction beam width.
15    pub ef_construction: usize,
16    /// Metric.
17    pub distance: Distance,
18}
19
20impl Default for HnswParams {
21    fn default() -> Self {
22        Self { m: 16, ef_construction: 200, distance: Distance::Cosine }
23    }
24}
25
26/// Sizing counters (RFC D6).
27#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
28pub struct VectorStats {
29    /// Living vectors.
30    pub vectors: u64,
31    /// Tombstoned nodes still in the graph.
32    pub tombstones: u64,
33    /// Total graph links.
34    pub links: u64,
35    /// Approximate heap bytes.
36    pub approx_bytes: u64,
37    /// 1 when tombstones exceed the rebuild threshold (30%).
38    pub rebuild_recommended: bool,
39}
40
41struct Node {
42    key: Vec<u8>,
43    vec: Vec<f32>,
44    /// links[layer] = neighbor node ids.
45    links: Vec<Vec<u32>>,
46    dead: bool,
47}
48
49/// One shard's ANN graph for one index.
50pub struct Hnsw {
51    params: HnswParams,
52    dim: usize,
53    nodes: Vec<Node>,
54    by_key: HashMap<Vec<u8>, u32>,
55    entry: Option<u32>,
56    live: u64,
57    /// Deterministic level generator (splitmix — no wall clock).
58    seed: u64,
59}
60
61/// Max-heap entry by distance (candidate pruning pops farthest).
62#[derive(PartialEq)]
63struct Far(f32, u32);
64impl Eq for Far {}
65impl PartialOrd for Far {
66    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
67        Some(self.cmp(other))
68    }
69}
70impl Ord for Far {
71    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
72        self.0.total_cmp(&other.0).then_with(|| self.1.cmp(&other.1))
73    }
74}
75
76impl Hnsw {
77    /// Empty graph for `dim`-dimensional vectors.
78    pub fn new(dim: usize, params: HnswParams) -> Self {
79        Self { params, dim, nodes: Vec::new(), by_key: HashMap::new(), entry: None, live: 0, seed: 0x9E3779B97F4A7C15 }
80    }
81
82    /// Declared dimensionality.
83    pub fn dim(&self) -> usize {
84        self.dim
85    }
86
87    /// Insert or replace `key`'s vector (`None` = remove). Replace =
88    /// tombstone old + insert new (RFC D5).
89    pub fn apply(&mut self, key: &[u8], vector: Option<Vec<f32>>) {
90        if let Some(&id) = self.by_key.get(key) {
91            let node = &mut self.nodes[id as usize];
92            if !node.dead {
93                node.dead = true;
94                self.live -= 1;
95            }
96            self.by_key.remove(key);
97            if self.entry == Some(id) {
98                self.entry = self.pick_entry();
99            }
100        }
101        let Some(mut v) = vector else { return };
102        if v.len() != self.dim {
103            return;
104        }
105        self.params.distance.prepare(&mut v);
106        self.insert_prepared(key.to_vec(), v);
107    }
108
109    fn pick_entry(&self) -> Option<u32> {
110        self.nodes
111            .iter()
112            .enumerate()
113            .filter(|(_, n)| !n.dead)
114            .max_by_key(|(_, n)| n.links.len())
115            .map(|(i, _)| i as u32)
116    }
117
118    fn rand_level(&mut self) -> usize {
119        // splitmix64 → uniform in (0,1) → geometric with 1/ln(M)
120        self.seed = self.seed.wrapping_add(0x9E3779B97F4A7C15);
121        let mut z = self.seed;
122        z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
123        z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
124        z ^= z >> 31;
125        let u = (z >> 11) as f64 / (1u64 << 53) as f64;
126        let ml = 1.0 / (self.params.m as f64).ln();
127        (-u.max(1e-12).ln() * ml).floor() as usize
128    }
129
130    fn insert_prepared(&mut self, key: Vec<u8>, v: Vec<f32>) {
131        let level = self.rand_level();
132        let id = self.nodes.len() as u32;
133        self.nodes.push(Node { key: key.clone(), vec: v, links: vec![Vec::new(); level + 1], dead: false });
134        self.by_key.insert(key, id);
135        self.live += 1;
136        let Some(mut cur) = self.entry else {
137            self.entry = Some(id);
138            return;
139        };
140        let top = (self.nodes[cur as usize].links.len() - 1) as i32;
141        // greedy descent above the node's level
142        for layer in ((level as i32 + 1)..=top).rev() {
143            cur = self.greedy_at(cur, id, layer as usize);
144        }
145        // beam insert on the node's layers
146        for layer in (0..=level.min(top.max(0) as usize)).rev() {
147            let found = self.search_layer(cur, id, layer, self.params.ef_construction, true);
148            let cap = if layer == 0 { self.params.m * 2 } else { self.params.m };
149            let chosen = self.select_diverse(&found, cap);
150            for &n in &chosen {
151                self.nodes[id as usize].links[layer].push(n);
152                self.nodes[n as usize].links[layer].push(id);
153                self.shrink(n, layer, cap);
154            }
155            if let Some(&(_, first)) = found.first() {
156                cur = first;
157            }
158        }
159        // a new top-level node becomes the entry
160        if level as i32 > top {
161            self.entry = Some(id);
162        }
163    }
164
165    fn greedy_at(&self, mut cur: u32, target: u32, layer: usize) -> u32 {
166        let tv = &self.nodes[target as usize].vec;
167        let mut best = self.params.distance.eval(&self.nodes[cur as usize].vec, tv);
168        loop {
169            let mut improved = false;
170            if layer < self.nodes[cur as usize].links.len() {
171                for &n in &self.nodes[cur as usize].links[layer] {
172                    let d = self.params.distance.eval(&self.nodes[n as usize].vec, tv);
173                    if d < best {
174                        best = d;
175                        cur = n;
176                        improved = true;
177                    }
178                }
179            }
180            if !improved {
181                return cur;
182            }
183        }
184    }
185
186    /// Beam search at one layer. `include_dead` keeps tombstones as
187    /// ROUTING waypoints (their links still connect the graph);
188    /// results always include them so the caller can filter.
189    fn search_layer(&self, start: u32, target: u32, layer: usize, ef: usize, _for_insert: bool) -> Vec<(f32, u32)> {
190        let tv = &self.nodes[target as usize].vec;
191        self.search_layer_vec(start, tv, layer, ef)
192    }
193
194    fn search_layer_vec(&self, start: u32, tv: &[f32], layer: usize, ef: usize) -> Vec<(f32, u32)> {
195        // The visited set is the beam search's hottest structure —
196        // perf-record put 63% of the EF16 KNN shape inside this fn,
197        // with the std HashMap's SipHash showing as a distinct cost.
198        // Classic hnswlib answer: an epoch-stamped visited pool —
199        // membership is ONE u32 array read, reset is `epoch += 1`,
200        // and the thread_local reuses the allocation across queries
201        // (thread-per-core: shards never share a search).
202        thread_local! {
203            static VISITED: std::cell::RefCell<(Vec<u32>, u32)> =
204                const { std::cell::RefCell::new((Vec::new(), 0)) };
205        }
206        VISITED.with(|cell| {
207            let (stamps, epoch) = &mut *cell.borrow_mut();
208            if stamps.len() < self.nodes.len() {
209                stamps.resize(self.nodes.len(), 0);
210            }
211            *epoch = epoch.wrapping_add(1);
212            if *epoch == 0 {
213                stamps.fill(0);
214                *epoch = 1;
215            }
216            let epoch = *epoch;
217            let mut result: BinaryHeap<Far> = BinaryHeap::with_capacity(ef + 1);
218            let mut frontier: BinaryHeap<std::cmp::Reverse<Far>> =
219                BinaryHeap::with_capacity(ef * 2);
220            let d0 = self.params.distance.eval(&self.nodes[start as usize].vec, tv);
221            stamps[start as usize] = epoch;
222            result.push(Far(d0, start));
223            frontier.push(std::cmp::Reverse(Far(d0, start)));
224            while let Some(std::cmp::Reverse(Far(d, node))) = frontier.pop() {
225                if result.len() >= ef
226                    && let Some(worst) = result.peek()
227                    && d > worst.0
228                {
229                    break;
230                }
231                if layer < self.nodes[node as usize].links.len() {
232                    for &n in &self.nodes[node as usize].links[layer] {
233                        if stamps[n as usize] == epoch {
234                            continue;
235                        }
236                        stamps[n as usize] = epoch;
237                        let dn = self.params.distance.eval(&self.nodes[n as usize].vec, tv);
238                        if result.len() < ef || dn < result.peek().expect("nonempty").0 {
239                            result.push(Far(dn, n));
240                            if result.len() > ef {
241                                result.pop();
242                            }
243                            frontier.push(std::cmp::Reverse(Far(dn, n)));
244                        }
245                    }
246                }
247            }
248            let mut out: Vec<(f32, u32)> = result.into_iter().map(|Far(d, n)| (d, n)).collect();
249            out.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
250            out
251        })
252    }
253
254    /// Malkov Algorithm 4 (diversity heuristic): walk candidates by
255    /// ascending distance; keep one only if it's closer to the node
256    /// than to every already-kept neighbor. This preserves BRIDGE
257    /// links to otherwise-isolated regions (an outlier's closest
258    /// in-graph node keeps its back-edge — plain closest-K pruning
259    /// disconnects it).
260    fn select_diverse(&self, sorted: &[(f32, u32)], cap: usize) -> Vec<u32> {
261        let mut kept: Vec<u32> = Vec::with_capacity(cap);
262        for &(d, c) in sorted {
263            if kept.len() == cap {
264                break;
265            }
266            let cv = &self.nodes[c as usize].vec;
267            let diverse = kept.iter().all(|&s| {
268                d < self.params.distance.eval(&self.nodes[s as usize].vec, cv)
269            });
270            if diverse {
271                kept.push(c);
272            }
273        }
274        // backfill with the nearest skipped candidates if under cap
275        if kept.len() < cap {
276            for &(_, c) in sorted {
277                if kept.len() == cap {
278                    break;
279                }
280                if !kept.contains(&c) {
281                    kept.push(c);
282                }
283            }
284        }
285        kept
286    }
287
288    fn shrink(&mut self, node: u32, layer: usize, cap: usize) {
289        if self.nodes[node as usize].links[layer].len() <= cap {
290            return;
291        }
292        let nv = &self.nodes[node as usize].vec;
293        let mut scored: Vec<(f32, u32)> = self.nodes[node as usize].links[layer]
294            .iter()
295            .map(|&n| (self.params.distance.eval(&self.nodes[n as usize].vec, nv), n))
296            .collect();
297        scored.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
298        scored.dedup_by_key(|e| e.1);
299        let kept = self.select_diverse(&scored, cap);
300        self.nodes[node as usize].links[layer] = kept;
301    }
302
303    /// k nearest LIVING vectors to `query` (raw form; prepared here).
304    /// `ef` = query beam width (0 → the max(4k, 100) default); larger
305    /// beams trade latency for recall — the canonical HNSW knob.
306    pub fn knn(&self, query: &[f32], k: usize, ef: usize) -> Vec<(Vec<u8>, f32)> {
307        let Some(entry) = self.entry else { return Vec::new() };
308        if query.len() != self.dim {
309            return Vec::new();
310        }
311        let mut q = query.to_vec();
312        self.params.distance.prepare(&mut q);
313        let mut cur = entry;
314        let top = self.nodes[cur as usize].links.len().saturating_sub(1);
315        for layer in (1..=top).rev() {
316            loop {
317                let cv = &self.nodes[cur as usize].vec;
318                let mut best = self.params.distance.eval(cv, &q);
319                let mut next = cur;
320                if layer < self.nodes[cur as usize].links.len() {
321                    for &n in &self.nodes[cur as usize].links[layer] {
322                        let d = self.params.distance.eval(&self.nodes[n as usize].vec, &q);
323                        if d < best {
324                            best = d;
325                            next = n;
326                        }
327                    }
328                }
329                if next == cur {
330                    break;
331                }
332                cur = next;
333            }
334        }
335        // Recall grows with beam width (measured on a dense 20k
336        // cluster @128d: ef 64 → 0.67 recall@10, 100 → 0.77); the
337        // default floor suits easy corpora, hard ones pass EF.
338        let ef = if ef == 0 { (k * 4).max(100) } else { ef.max(k) };
339        let found = self.search_layer_vec(cur, &q, 0, ef);
340        found
341            .into_iter()
342            .filter(|&(_, n)| !self.nodes[n as usize].dead)
343            .take(k)
344            .map(|(d, n)| (self.nodes[n as usize].key.clone(), d))
345            .collect()
346    }
347
348    /// Membership (living only).
349    pub fn contains(&self, key: &[u8]) -> bool {
350        self.by_key.contains_key(key)
351    }
352
353    /// Counters (RFC D6).
354    pub fn stats(&self) -> VectorStats {
355        let links: u64 = self.nodes.iter().map(|n| n.links.iter().map(Vec::len).sum::<usize>() as u64).sum();
356        let tombstones = self.nodes.len() as u64 - self.live;
357        let bytes_vec = (self.dim * 4) as u64;
358        let approx_bytes: u64 = self.nodes.len() as u64 * (bytes_vec + 40)
359            + links * 8
360            + self.live * 32;
361        VectorStats {
362            vectors: self.live,
363            tombstones,
364            links,
365            approx_bytes,
366            rebuild_recommended: !self.nodes.is_empty() && tombstones * 10 > self.nodes.len() as u64 * 3,
367        }
368    }
369
370    /// Bounded rebuild: re-insert every living vector into a fresh
371    /// graph (drops tombstones and their edges) — RFC D5.
372    pub fn rebuild(&mut self) {
373        let mut fresh = Hnsw::new(self.dim, self.params);
374        fresh.seed = self.seed;
375        for node in &self.nodes {
376            if !node.dead {
377                fresh.insert_prepared(node.key.clone(), node.vec.clone());
378            }
379        }
380        *self = fresh;
381    }
382}
383
384#[cfg(test)]
385mod tests {
386    use super::*;
387
388    fn grid(n: usize) -> Hnsw {
389        // 2-d grid points — L2 neighbors are unambiguous
390        let mut h = Hnsw::new(2, HnswParams { distance: Distance::L2, ..Default::default() });
391        for i in 0..n {
392            let (x, y) = ((i % 32) as f32, (i / 32) as f32);
393            h.apply(format!("p{i:04}").as_bytes(), Some(vec![x, y]));
394        }
395        h
396    }
397
398    #[test]
399    fn knn_exact_on_grid() {
400        let h = grid(1024);
401        // nearest to (5.1, 7.05) is p(7*32+5)=p0229, then p0197/p0230…
402        let hits = h.knn(&[5.1, 7.05], 3, 0);
403        assert_eq!(hits[0].0, b"p0229".to_vec(), "{hits:?}");
404        assert_eq!(hits.len(), 3);
405        assert!(hits[0].1 <= hits[1].1);
406    }
407
408    #[test]
409    fn tombstone_and_replace() {
410        let mut h = grid(256);
411        h.apply(b"p0000", None);
412        assert!(!h.contains(b"p0000"));
413        let hits = h.knn(&[0.0, 0.0], 1, 0);
414        assert_ne!(hits[0].0, b"p0000".to_vec(), "dead filtered");
415        // replace moves the point
416        h.apply(b"p0001", Some(vec![100.0, 100.0]));
417        let hits = h.knn(&[100.0, 100.0], 1, 0);
418        assert_eq!(hits[0].0, b"p0001".to_vec());
419        let st = h.stats();
420        assert_eq!(st.vectors, 255);
421        assert_eq!(st.tombstones, 2, "one delete + one replace");
422    }
423
424    #[test]
425    fn recall_on_random_vectors() {
426        // deterministic pseudo-random 64-d vectors; HNSW top-10 vs
427        // brute force ground truth, recall must be ≥ 0.9
428        let mut seed = 42u64;
429        let mut rnd = move || {
430            seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
431            ((seed >> 11) as f64 / (1u64 << 53) as f64) as f32 - 0.5
432        };
433        let mut h = Hnsw::new(64, HnswParams::default());
434        let mut all: Vec<(Vec<u8>, Vec<f32>)> = Vec::new();
435        for i in 0..2000 {
436            let v: Vec<f32> = (0..64).map(|_| rnd()).collect();
437            let key = format!("v{i:04}").into_bytes();
438            h.apply(&key, Some(v.clone()));
439            all.push((key, v));
440        }
441        let mut hit = 0usize;
442        let mut total = 0usize;
443        for qi in 0..20 {
444            let q: Vec<f32> = (0..64).map(|_| rnd()).collect();
445            let got: Vec<Vec<u8>> = h.knn(&q, 10, 0).into_iter().map(|(k, _)| k).collect();
446            // brute force with the same metric incl. normalization
447            let mut qq = q.clone();
448            Distance::Cosine.prepare(&mut qq);
449            let mut truth: Vec<(f32, &[u8])> = all
450                .iter()
451                .map(|(k, v)| {
452                    let mut vv = v.clone();
453                    Distance::Cosine.prepare(&mut vv);
454                    (Distance::Cosine.eval(&vv, &qq), k.as_slice())
455                })
456                .collect();
457            truth.sort_by(|a, b| a.0.total_cmp(&b.0));
458            let want: Vec<&[u8]> = truth[..10].iter().map(|(_, k)| *k).collect();
459            for w in &want {
460                total += 1;
461                if got.iter().any(|g| g == w) {
462                    hit += 1;
463                }
464            }
465            let _ = qi;
466        }
467        let recall = hit as f64 / total as f64;
468        assert!(recall >= 0.9, "recall {recall}");
469    }
470
471    #[test]
472    fn rebuild_drops_tombstones_preserves_answers() {
473        let mut h = grid(512);
474        for i in 0..200 {
475            h.apply(format!("p{i:04}").as_bytes(), None);
476        }
477        assert!(h.stats().rebuild_recommended);
478        let before = h.knn(&[20.0, 10.0], 5, 0);
479        h.rebuild();
480        let st = h.stats();
481        assert_eq!(st.tombstones, 0);
482        assert_eq!(st.vectors, 312);
483        let after = h.knn(&[20.0, 10.0], 5, 0);
484        assert_eq!(
485            before.iter().map(|(k, _)| k).collect::<Vec<_>>(),
486            after.iter().map(|(k, _)| k).collect::<Vec<_>>()
487        );
488    }
489
490    #[test]
491    fn empty_and_dim_mismatch() {
492        let h = Hnsw::new(4, HnswParams::default());
493        assert!(h.knn(&[1.0, 2.0, 3.0, 4.0], 5, 0).is_empty());
494        let mut h = grid(16);
495        h.apply(b"bad", Some(vec![1.0, 2.0, 3.0])); // wrong dim ignored
496        assert!(!h.contains(b"bad"));
497        assert!(h.knn(&[1.0], 5, 0).is_empty(), "query dim mismatch");
498    }
499}