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        let mut visited: HashMap<u32, ()> = HashMap::new();
196        let mut result: BinaryHeap<Far> = BinaryHeap::new(); // max-heap of best ef
197        let mut frontier: BinaryHeap<std::cmp::Reverse<Far>> = BinaryHeap::new();
198        let d0 = self.params.distance.eval(&self.nodes[start as usize].vec, tv);
199        visited.insert(start, ());
200        result.push(Far(d0, start));
201        frontier.push(std::cmp::Reverse(Far(d0, start)));
202        while let Some(std::cmp::Reverse(Far(d, node))) = frontier.pop() {
203            if result.len() >= ef
204                && let Some(worst) = result.peek()
205                && d > worst.0
206            {
207                break;
208            }
209            if layer < self.nodes[node as usize].links.len() {
210                for &n in &self.nodes[node as usize].links[layer] {
211                    if visited.insert(n, ()).is_some() {
212                        continue;
213                    }
214                    let dn = self.params.distance.eval(&self.nodes[n as usize].vec, tv);
215                    if result.len() < ef || dn < result.peek().expect("nonempty").0 {
216                        result.push(Far(dn, n));
217                        if result.len() > ef {
218                            result.pop();
219                        }
220                        frontier.push(std::cmp::Reverse(Far(dn, n)));
221                    }
222                }
223            }
224        }
225        let mut out: Vec<(f32, u32)> = result.into_iter().map(|Far(d, n)| (d, n)).collect();
226        out.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
227        out
228    }
229
230    /// Malkov Algorithm 4 (diversity heuristic): walk candidates by
231    /// ascending distance; keep one only if it's closer to the node
232    /// than to every already-kept neighbor. This preserves BRIDGE
233    /// links to otherwise-isolated regions (an outlier's closest
234    /// in-graph node keeps its back-edge — plain closest-K pruning
235    /// disconnects it).
236    fn select_diverse(&self, sorted: &[(f32, u32)], cap: usize) -> Vec<u32> {
237        let mut kept: Vec<u32> = Vec::with_capacity(cap);
238        for &(d, c) in sorted {
239            if kept.len() == cap {
240                break;
241            }
242            let cv = &self.nodes[c as usize].vec;
243            let diverse = kept.iter().all(|&s| {
244                d < self.params.distance.eval(&self.nodes[s as usize].vec, cv)
245            });
246            if diverse {
247                kept.push(c);
248            }
249        }
250        // backfill with the nearest skipped candidates if under cap
251        if kept.len() < cap {
252            for &(_, c) in sorted {
253                if kept.len() == cap {
254                    break;
255                }
256                if !kept.contains(&c) {
257                    kept.push(c);
258                }
259            }
260        }
261        kept
262    }
263
264    fn shrink(&mut self, node: u32, layer: usize, cap: usize) {
265        if self.nodes[node as usize].links[layer].len() <= cap {
266            return;
267        }
268        let nv = &self.nodes[node as usize].vec;
269        let mut scored: Vec<(f32, u32)> = self.nodes[node as usize].links[layer]
270            .iter()
271            .map(|&n| (self.params.distance.eval(&self.nodes[n as usize].vec, nv), n))
272            .collect();
273        scored.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
274        scored.dedup_by_key(|e| e.1);
275        let kept = self.select_diverse(&scored, cap);
276        self.nodes[node as usize].links[layer] = kept;
277    }
278
279    /// k nearest LIVING vectors to `query` (raw form; prepared here).
280    /// `ef` = query beam width (0 → the max(4k, 100) default); larger
281    /// beams trade latency for recall — the canonical HNSW knob.
282    pub fn knn(&self, query: &[f32], k: usize, ef: usize) -> Vec<(Vec<u8>, f32)> {
283        let Some(entry) = self.entry else { return Vec::new() };
284        if query.len() != self.dim {
285            return Vec::new();
286        }
287        let mut q = query.to_vec();
288        self.params.distance.prepare(&mut q);
289        let mut cur = entry;
290        let top = self.nodes[cur as usize].links.len().saturating_sub(1);
291        for layer in (1..=top).rev() {
292            loop {
293                let cv = &self.nodes[cur as usize].vec;
294                let mut best = self.params.distance.eval(cv, &q);
295                let mut next = cur;
296                if layer < self.nodes[cur as usize].links.len() {
297                    for &n in &self.nodes[cur as usize].links[layer] {
298                        let d = self.params.distance.eval(&self.nodes[n as usize].vec, &q);
299                        if d < best {
300                            best = d;
301                            next = n;
302                        }
303                    }
304                }
305                if next == cur {
306                    break;
307                }
308                cur = next;
309            }
310        }
311        // Recall grows with beam width (measured on a dense 20k
312        // cluster @128d: ef 64 → 0.67 recall@10, 100 → 0.77); the
313        // default floor suits easy corpora, hard ones pass EF.
314        let ef = if ef == 0 { (k * 4).max(100) } else { ef.max(k) };
315        let found = self.search_layer_vec(cur, &q, 0, ef);
316        found
317            .into_iter()
318            .filter(|&(_, n)| !self.nodes[n as usize].dead)
319            .take(k)
320            .map(|(d, n)| (self.nodes[n as usize].key.clone(), d))
321            .collect()
322    }
323
324    /// Membership (living only).
325    pub fn contains(&self, key: &[u8]) -> bool {
326        self.by_key.contains_key(key)
327    }
328
329    /// Counters (RFC D6).
330    pub fn stats(&self) -> VectorStats {
331        let links: u64 = self.nodes.iter().map(|n| n.links.iter().map(Vec::len).sum::<usize>() as u64).sum();
332        let tombstones = self.nodes.len() as u64 - self.live;
333        let bytes_vec = (self.dim * 4) as u64;
334        let approx_bytes: u64 = self.nodes.len() as u64 * (bytes_vec + 40)
335            + links * 8
336            + self.live * 32;
337        VectorStats {
338            vectors: self.live,
339            tombstones,
340            links,
341            approx_bytes,
342            rebuild_recommended: !self.nodes.is_empty() && tombstones * 10 > self.nodes.len() as u64 * 3,
343        }
344    }
345
346    /// Bounded rebuild: re-insert every living vector into a fresh
347    /// graph (drops tombstones and their edges) — RFC D5.
348    pub fn rebuild(&mut self) {
349        let mut fresh = Hnsw::new(self.dim, self.params);
350        fresh.seed = self.seed;
351        for node in &self.nodes {
352            if !node.dead {
353                fresh.insert_prepared(node.key.clone(), node.vec.clone());
354            }
355        }
356        *self = fresh;
357    }
358}
359
360#[cfg(test)]
361mod tests {
362    use super::*;
363
364    fn grid(n: usize) -> Hnsw {
365        // 2-d grid points — L2 neighbors are unambiguous
366        let mut h = Hnsw::new(2, HnswParams { distance: Distance::L2, ..Default::default() });
367        for i in 0..n {
368            let (x, y) = ((i % 32) as f32, (i / 32) as f32);
369            h.apply(format!("p{i:04}").as_bytes(), Some(vec![x, y]));
370        }
371        h
372    }
373
374    #[test]
375    fn knn_exact_on_grid() {
376        let h = grid(1024);
377        // nearest to (5.1, 7.05) is p(7*32+5)=p0229, then p0197/p0230…
378        let hits = h.knn(&[5.1, 7.05], 3, 0);
379        assert_eq!(hits[0].0, b"p0229".to_vec(), "{hits:?}");
380        assert_eq!(hits.len(), 3);
381        assert!(hits[0].1 <= hits[1].1);
382    }
383
384    #[test]
385    fn tombstone_and_replace() {
386        let mut h = grid(256);
387        h.apply(b"p0000", None);
388        assert!(!h.contains(b"p0000"));
389        let hits = h.knn(&[0.0, 0.0], 1, 0);
390        assert_ne!(hits[0].0, b"p0000".to_vec(), "dead filtered");
391        // replace moves the point
392        h.apply(b"p0001", Some(vec![100.0, 100.0]));
393        let hits = h.knn(&[100.0, 100.0], 1, 0);
394        assert_eq!(hits[0].0, b"p0001".to_vec());
395        let st = h.stats();
396        assert_eq!(st.vectors, 255);
397        assert_eq!(st.tombstones, 2, "one delete + one replace");
398    }
399
400    #[test]
401    fn recall_on_random_vectors() {
402        // deterministic pseudo-random 64-d vectors; HNSW top-10 vs
403        // brute force ground truth, recall must be ≥ 0.9
404        let mut seed = 42u64;
405        let mut rnd = move || {
406            seed = seed.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
407            ((seed >> 11) as f64 / (1u64 << 53) as f64) as f32 - 0.5
408        };
409        let mut h = Hnsw::new(64, HnswParams::default());
410        let mut all: Vec<(Vec<u8>, Vec<f32>)> = Vec::new();
411        for i in 0..2000 {
412            let v: Vec<f32> = (0..64).map(|_| rnd()).collect();
413            let key = format!("v{i:04}").into_bytes();
414            h.apply(&key, Some(v.clone()));
415            all.push((key, v));
416        }
417        let mut hit = 0usize;
418        let mut total = 0usize;
419        for qi in 0..20 {
420            let q: Vec<f32> = (0..64).map(|_| rnd()).collect();
421            let got: Vec<Vec<u8>> = h.knn(&q, 10, 0).into_iter().map(|(k, _)| k).collect();
422            // brute force with the same metric incl. normalization
423            let mut qq = q.clone();
424            Distance::Cosine.prepare(&mut qq);
425            let mut truth: Vec<(f32, &[u8])> = all
426                .iter()
427                .map(|(k, v)| {
428                    let mut vv = v.clone();
429                    Distance::Cosine.prepare(&mut vv);
430                    (Distance::Cosine.eval(&vv, &qq), k.as_slice())
431                })
432                .collect();
433            truth.sort_by(|a, b| a.0.total_cmp(&b.0));
434            let want: Vec<&[u8]> = truth[..10].iter().map(|(_, k)| *k).collect();
435            for w in &want {
436                total += 1;
437                if got.iter().any(|g| g == w) {
438                    hit += 1;
439                }
440            }
441            let _ = qi;
442        }
443        let recall = hit as f64 / total as f64;
444        assert!(recall >= 0.9, "recall {recall}");
445    }
446
447    #[test]
448    fn rebuild_drops_tombstones_preserves_answers() {
449        let mut h = grid(512);
450        for i in 0..200 {
451            h.apply(format!("p{i:04}").as_bytes(), None);
452        }
453        assert!(h.stats().rebuild_recommended);
454        let before = h.knn(&[20.0, 10.0], 5, 0);
455        h.rebuild();
456        let st = h.stats();
457        assert_eq!(st.tombstones, 0);
458        assert_eq!(st.vectors, 312);
459        let after = h.knn(&[20.0, 10.0], 5, 0);
460        assert_eq!(
461            before.iter().map(|(k, _)| k).collect::<Vec<_>>(),
462            after.iter().map(|(k, _)| k).collect::<Vec<_>>()
463        );
464    }
465
466    #[test]
467    fn empty_and_dim_mismatch() {
468        let h = Hnsw::new(4, HnswParams::default());
469        assert!(h.knn(&[1.0, 2.0, 3.0, 4.0], 5, 0).is_empty());
470        let mut h = grid(16);
471        h.apply(b"bad", Some(vec![1.0, 2.0, 3.0])); // wrong dim ignored
472        assert!(!h.contains(b"bad"));
473        assert!(h.knn(&[1.0], 5, 0).is_empty(), "query dim mismatch");
474    }
475}