1use std::collections::{BinaryHeap, HashMap};
6
7use crate::dist::Distance;
8
9#[derive(Debug, Clone, Copy)]
11pub struct HnswParams {
12 pub m: usize,
14 pub ef_construction: usize,
16 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#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
28pub struct VectorStats {
29 pub vectors: u64,
31 pub tombstones: u64,
33 pub links: u64,
35 pub approx_bytes: u64,
37 pub rebuild_recommended: bool,
39}
40
41struct Node {
42 key: Vec<u8>,
43 vec: Vec<f32>,
44 links: Vec<Vec<u32>>,
46 dead: bool,
47}
48
49pub 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 seed: u64,
59}
60
61#[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 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 pub fn dim(&self) -> usize {
84 self.dim
85 }
86
87 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 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 for layer in ((level as i32 + 1)..=top).rev() {
143 cur = self.greedy_at(cur, id, layer as usize);
144 }
145 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 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 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 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 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 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 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 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 pub fn contains(&self, key: &[u8]) -> bool {
350 self.by_key.contains_key(key)
351 }
352
353 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 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 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 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 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 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 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])); assert!(!h.contains(b"bad"));
497 assert!(h.knn(&[1.0], 5, 0).is_empty(), "query dim mismatch");
498 }
499}