Skip to main content

kevy_vector/
hnsw.rs

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