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    /// Every LIVING key whose vector is exactly this one (duplicate
43    /// vectors under different keys collapse onto ONE graph node —
44    /// fuzz finding 2026-07-10: one-node-per-key duplicate clusters
45    /// larger than the link cap disconnect from the graph because
46    /// every co-located edge ties in the diversity prune).
47    keys: Vec<Vec<u8>>,
48    vec: Vec<f32>,
49    /// links[layer] = neighbor node ids.
50    links: Vec<Vec<u32>>,
51    dead: bool,
52}
53
54/// One shard's ANN graph for one index.
55pub struct Hnsw {
56    params: HnswParams,
57    dim: usize,
58    nodes: Vec<Node>,
59    by_key: HashMap<Vec<u8>, u32>,
60    /// Prepared-vector bits → living node holding that exact vector
61    /// (the duplicate-collapse index; bitwise equality, so -0.0/0.0
62    /// stay distinct nodes — harmless, the tie-keeping prune covers
63    /// sub-cap co-located pairs).
64    by_vec: HashMap<Vec<u32>, u32>,
65    entry: Option<u32>,
66    /// Living KEYS (≥ living nodes when duplicates are collapsed).
67    live: u64,
68    /// Deterministic level generator (splitmix — no wall clock).
69    seed: u64,
70}
71
72/// Bitwise identity of a prepared vector (`by_vec` map key).
73fn vec_bits(v: &[f32]) -> Vec<u32> {
74    v.iter().map(|x| x.to_bits()).collect()
75}
76
77/// Max-heap entry by distance (candidate pruning pops farthest).
78#[derive(PartialEq)]
79struct Far(f32, u32);
80impl Eq for Far {}
81impl PartialOrd for Far {
82    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
83        Some(self.cmp(other))
84    }
85}
86impl Ord for Far {
87    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
88        self.0.total_cmp(&other.0).then_with(|| self.1.cmp(&other.1))
89    }
90}
91
92impl Hnsw {
93    /// Empty graph for `dim`-dimensional vectors.
94    pub fn new(dim: usize, params: HnswParams) -> Self {
95        Self {
96            params,
97            dim,
98            nodes: Vec::new(),
99            by_key: HashMap::new(),
100            by_vec: HashMap::new(),
101            entry: None,
102            live: 0,
103            seed: 0x9E37_79B9_7F4A_7C15,
104        }
105    }
106
107    /// Declared dimensionality.
108    pub fn dim(&self) -> usize {
109        self.dim
110    }
111
112    /// Insert or replace `key`'s vector (`None` = remove). Replace =
113    /// detach old key (tombstone the node once keyless) + insert new
114    /// (RFC D5). Keys sharing one exact vector share one graph node.
115    pub fn apply(&mut self, key: &[u8], vector: Option<Vec<f32>>) {
116        if let Some(id) = self.by_key.remove(key) {
117            let node = &mut self.nodes[id as usize];
118            node.keys.retain(|k| k != key);
119            self.live -= 1;
120            if node.keys.is_empty() {
121                node.dead = true;
122                self.by_vec.remove(&vec_bits(&node.vec));
123                if self.entry == Some(id) {
124                    self.entry = self.pick_entry();
125                }
126            }
127        }
128        let Some(mut v) = vector else { return };
129        if v.len() != self.dim {
130            return;
131        }
132        self.params.distance.prepare(&mut v);
133        self.add_key(key.to_vec(), v);
134    }
135
136    /// Attach a (key, PREPARED vector) pair: onto the living node
137    /// already holding that exact vector, or as a fresh graph node.
138    fn add_key(&mut self, key: Vec<u8>, v: Vec<f32>) {
139        if let Some(&id) = self.by_vec.get(&vec_bits(&v)) {
140            self.nodes[id as usize].keys.push(key.clone());
141            self.by_key.insert(key, id);
142            self.live += 1;
143            return;
144        }
145        self.insert_prepared(key, v);
146    }
147
148    fn pick_entry(&self) -> Option<u32> {
149        self.nodes
150            .iter()
151            .enumerate()
152            .filter(|(_, n)| !n.dead)
153            .max_by_key(|(_, n)| n.links.len())
154            .map(|(i, _)| i as u32)
155    }
156
157    fn rand_level(&mut self) -> usize {
158        // splitmix64 → uniform in (0,1) → geometric with 1/ln(M)
159        self.seed = self.seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
160        let mut z = self.seed;
161        z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
162        z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
163        z ^= z >> 31;
164        let u = (z >> 11) as f64 / (1u64 << 53) as f64;
165        let ml = 1.0 / (self.params.m as f64).ln();
166        (-u.max(1e-12).ln() * ml).floor() as usize
167    }
168
169    fn insert_prepared(&mut self, key: Vec<u8>, v: Vec<f32>) {
170        let level = self.rand_level();
171        let id = self.nodes.len() as u32;
172        self.by_vec.insert(vec_bits(&v), id);
173        self.nodes.push(Node {
174            keys: vec![key.clone()],
175            vec: v,
176            links: vec![Vec::new(); level + 1],
177            dead: false,
178        });
179        self.by_key.insert(key, id);
180        self.live += 1;
181        let Some(mut cur) = self.entry else {
182            self.entry = Some(id);
183            return;
184        };
185        let top = (self.nodes[cur as usize].links.len() - 1) as i32;
186        // greedy descent above the node's level
187        for layer in ((level as i32 + 1)..=top).rev() {
188            cur = self.greedy_at(cur, id, layer as usize);
189        }
190        // beam insert on the node's layers
191        for layer in (0..=level.min(top.max(0) as usize)).rev() {
192            let found = self.search_layer(cur, id, layer, self.params.ef_construction, true);
193            let cap = if layer == 0 { self.params.m * 2 } else { self.params.m };
194            let chosen = self.select_diverse(&found, cap, &self.nodes[id as usize].vec);
195            for &n in &chosen {
196                self.nodes[id as usize].links[layer].push(n);
197                self.nodes[n as usize].links[layer].push(id);
198                self.shrink(n, layer, cap);
199            }
200            if let Some(&(_, first)) = found.first() {
201                cur = first;
202            }
203        }
204        // a new top-level node becomes the entry
205        if level as i32 > top {
206            self.entry = Some(id);
207        }
208    }
209
210    fn greedy_at(&self, mut cur: u32, target: u32, layer: usize) -> u32 {
211        let tv = &self.nodes[target as usize].vec;
212        let mut best = self.params.distance.eval(&self.nodes[cur as usize].vec, tv);
213        loop {
214            let mut improved = false;
215            if layer < self.nodes[cur as usize].links.len() {
216                for &n in &self.nodes[cur as usize].links[layer] {
217                    let d = self.params.distance.eval(&self.nodes[n as usize].vec, tv);
218                    if d < best {
219                        best = d;
220                        cur = n;
221                        improved = true;
222                    }
223                }
224            }
225            if !improved {
226                return cur;
227            }
228        }
229    }
230
231    /// Beam search at one layer. `include_dead` keeps tombstones as
232    /// ROUTING waypoints (their links still connect the graph);
233    /// results always include them so the caller can filter.
234    fn search_layer(&self, start: u32, target: u32, layer: usize, ef: usize, _for_insert: bool) -> Vec<(f32, u32)> {
235        let tv = &self.nodes[target as usize].vec;
236        self.search_layer_vec(start, tv, layer, ef)
237    }
238
239    // LOC-WAIVER: per-query beam-search hot body (63% of EF16 KNN self-time; see comment below).
240    fn search_layer_vec(&self, start: u32, tv: &[f32], layer: usize, ef: usize) -> Vec<(f32, u32)> {
241        // The visited set is the beam search's hottest structure —
242        // perf-record put 63% of the EF16 KNN shape inside this fn,
243        // with the std HashMap's SipHash showing as a distinct cost.
244        // Classic hnswlib answer: an epoch-stamped visited pool —
245        // membership is ONE u32 array read, reset is `epoch += 1`,
246        // and the thread_local reuses the allocation across queries
247        // (thread-per-core: shards never share a search).
248        thread_local! {
249            static VISITED: std::cell::RefCell<(Vec<u32>, u32)> =
250                const { std::cell::RefCell::new((Vec::new(), 0)) };
251        }
252        VISITED.with(|cell| {
253            let (stamps, epoch) = &mut *cell.borrow_mut();
254            if stamps.len() < self.nodes.len() {
255                stamps.resize(self.nodes.len(), 0);
256            }
257            *epoch = epoch.wrapping_add(1);
258            if *epoch == 0 {
259                stamps.fill(0);
260                *epoch = 1;
261            }
262            let epoch = *epoch;
263            let mut result: BinaryHeap<Far> = BinaryHeap::with_capacity(ef + 1);
264            let mut frontier: BinaryHeap<std::cmp::Reverse<Far>> =
265                BinaryHeap::with_capacity(ef * 2);
266            let d0 = self.params.distance.eval(&self.nodes[start as usize].vec, tv);
267            stamps[start as usize] = epoch;
268            result.push(Far(d0, start));
269            frontier.push(std::cmp::Reverse(Far(d0, start)));
270            while let Some(std::cmp::Reverse(Far(d, node))) = frontier.pop() {
271                if result.len() >= ef
272                    && let Some(worst) = result.peek()
273                    && d > worst.0
274                {
275                    break;
276                }
277                if layer < self.nodes[node as usize].links.len() {
278                    for &n in &self.nodes[node as usize].links[layer] {
279                        if stamps[n as usize] == epoch {
280                            continue;
281                        }
282                        stamps[n as usize] = epoch;
283                        let dn = self.params.distance.eval(&self.nodes[n as usize].vec, tv);
284                        if result.len() < ef || dn < result.peek().expect("nonempty").0 {
285                            result.push(Far(dn, n));
286                            if result.len() > ef {
287                                result.pop();
288                            }
289                            frontier.push(std::cmp::Reverse(Far(dn, n)));
290                        }
291                    }
292                }
293            }
294            let mut out: Vec<(f32, u32)> = result.into_iter().map(|Far(d, n)| (d, n)).collect();
295            out.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
296            out
297        })
298    }
299
300    /// Malkov Algorithm 4 (diversity heuristic): walk candidates by
301    /// ascending distance; keep one unless an already-kept neighbor is
302    /// STRICTLY closer to it than the node is. This preserves BRIDGE
303    /// links to otherwise-isolated regions (an outlier's closest
304    /// in-graph node keeps its back-edge — plain closest-K pruning
305    /// disconnects it).
306    ///
307    /// Duplicate handling (fuzz finding 2026-07-10, recall@10 = 0.8
308    /// under an exhaustive beam — duplicate vectors under different
309    /// keys are legal in production):
310    ///
311    /// * ties are kept (`<=`, matching hnswlib): with a strict `<`,
312    ///   any candidate tying a kept neighbor — always the case once a
313    ///   kept neighbor duplicates the node — lost, degenerating the
314    ///   prune to closest-K, which drops bridges;
315    /// * candidates co-located WITH the node collapse to ONE
316    ///   representative edge (they tie everything, so without the cap
317    ///   a duplicate cluster larger than `cap` fills every slot and
318    ///   the cluster's bridges to the rest of the graph are all
319    ///   pruned — the cluster becomes an island); the backfill also
320    ///   prefers non-co-located candidates for the same reason.
321    fn select_diverse(&self, sorted: &[(f32, u32)], cap: usize, node_vec: &[f32]) -> Vec<u32> {
322        // Co-location = vector equality, NOT distance 0 (ip distance
323        // of co-located vectors is -|v|², and 0 for orthogonal ones).
324        let co = |c: u32| self.nodes[c as usize].vec == node_vec;
325        let mut kept: Vec<u32> = Vec::with_capacity(cap);
326        let mut have_twin = false;
327        for &(d, c) in sorted {
328            if kept.len() == cap {
329                break;
330            }
331            if co(c) {
332                if !have_twin {
333                    have_twin = true;
334                    kept.push(c);
335                }
336                continue;
337            }
338            let cv = &self.nodes[c as usize].vec;
339            let diverse = kept.iter().all(|&s| {
340                d <= self.params.distance.eval(&self.nodes[s as usize].vec, cv)
341            });
342            if diverse {
343                kept.push(c);
344            }
345        }
346        // Backfill with the nearest skipped candidates if under cap —
347        // non-co-located first (bridges), co-located twins last.
348        for pass in [false, true] {
349            for &(_, c) in sorted {
350                if kept.len() == cap {
351                    return kept;
352                }
353                if (pass || !co(c)) && !kept.contains(&c) {
354                    kept.push(c);
355                }
356            }
357        }
358        kept
359    }
360
361    fn shrink(&mut self, node: u32, layer: usize, cap: usize) {
362        if self.nodes[node as usize].links[layer].len() <= cap {
363            return;
364        }
365        let nv = &self.nodes[node as usize].vec;
366        let mut scored: Vec<(f32, u32)> = self.nodes[node as usize].links[layer]
367            .iter()
368            .map(|&n| (self.params.distance.eval(&self.nodes[n as usize].vec, nv), n))
369            .collect();
370        scored.sort_by(|a, b| a.0.total_cmp(&b.0).then_with(|| a.1.cmp(&b.1)));
371        scored.dedup_by_key(|e| e.1);
372        let kept = self.select_diverse(&scored, cap, &self.nodes[node as usize].vec);
373        self.nodes[node as usize].links[layer] = kept;
374    }
375
376    /// k nearest LIVING vectors to `query` (raw form; prepared here).
377    /// `ef` = query beam width (0 → the max(4k, 100) default); larger
378    /// beams trade latency for recall — the canonical HNSW knob.
379    pub fn knn(&self, query: &[f32], k: usize, ef: usize) -> Vec<(Vec<u8>, f32)> {
380        let Some(entry) = self.entry else { return Vec::new() };
381        if query.len() != self.dim {
382            return Vec::new();
383        }
384        let mut q = query.to_vec();
385        self.params.distance.prepare(&mut q);
386        let mut cur = entry;
387        let top = self.nodes[cur as usize].links.len().saturating_sub(1);
388        for layer in (1..=top).rev() {
389            loop {
390                let cv = &self.nodes[cur as usize].vec;
391                let mut best = self.params.distance.eval(cv, &q);
392                let mut next = cur;
393                if layer < self.nodes[cur as usize].links.len() {
394                    for &n in &self.nodes[cur as usize].links[layer] {
395                        let d = self.params.distance.eval(&self.nodes[n as usize].vec, &q);
396                        if d < best {
397                            best = d;
398                            next = n;
399                        }
400                    }
401                }
402                if next == cur {
403                    break;
404                }
405                cur = next;
406            }
407        }
408        // Recall grows with beam width (measured on a dense 20k
409        // cluster @128d: ef 64 → 0.67 recall@10, 100 → 0.77); the
410        // default floor suits easy corpora, hard ones pass EF.
411        let ef = if ef == 0 { (k * 4).max(100) } else { ef.max(k) };
412        let found = self.search_layer_vec(cur, &q, 0, ef);
413        self.expand_living(found, k)
414    }
415
416    /// Expand collapsed duplicates: one graph node answers for every
417    /// living key sharing its vector (all at the node's distance).
418    fn expand_living(&self, found: Vec<(f32, u32)>, k: usize) -> Vec<(Vec<u8>, f32)> {
419        let mut out: Vec<(Vec<u8>, f32)> = Vec::with_capacity(k);
420        for (d, n) in found {
421            let node = &self.nodes[n as usize];
422            if node.dead {
423                continue;
424            }
425            for key in &node.keys {
426                if out.len() == k {
427                    return out;
428                }
429                out.push((key.clone(), d));
430            }
431        }
432        out
433    }
434
435    /// Membership (living only).
436    pub fn contains(&self, key: &[u8]) -> bool {
437        self.by_key.contains_key(key)
438    }
439
440    /// Counters (RFC D6).
441    pub fn stats(&self) -> VectorStats {
442        let links: u64 = self.nodes.iter().map(|n| n.links.iter().map(Vec::len).sum::<usize>() as u64).sum();
443        let tombstones = self.nodes.iter().filter(|n| n.dead).count() as u64;
444        let bytes_vec = (self.dim * 4) as u64;
445        let approx_bytes: u64 = self.nodes.len() as u64 * (bytes_vec + 40)
446            + links * 8
447            + self.live * 32;
448        VectorStats {
449            vectors: self.live,
450            tombstones,
451            links,
452            approx_bytes,
453            rebuild_recommended: !self.nodes.is_empty() && tombstones * 10 > self.nodes.len() as u64 * 3,
454        }
455    }
456
457    /// Bounded rebuild: re-insert every living (key, vector) pair into
458    /// a fresh graph (drops tombstones and their edges) — RFC D5.
459    /// Vectors are already prepared; `add_key` re-collapses duplicates.
460    pub fn rebuild(&mut self) {
461        let mut fresh = Hnsw::new(self.dim, self.params);
462        fresh.seed = self.seed;
463        for node in &self.nodes {
464            if !node.dead {
465                for key in &node.keys {
466                    fresh.add_key(key.clone(), node.vec.clone());
467                }
468            }
469        }
470        *self = fresh;
471    }
472}
473
474#[cfg(test)]
475#[path = "hnsw_tests.rs"]
476mod tests;