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 let mut visited: HashMap<u32, ()> = HashMap::new();
196 let mut result: BinaryHeap<Far> = BinaryHeap::new(); 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 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 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 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 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 pub fn contains(&self, key: &[u8]) -> bool {
326 self.by_key.contains_key(key)
327 }
328
329 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 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 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 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 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 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 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])); assert!(!h.contains(b"bad"));
473 assert!(h.knn(&[1.0], 5, 0).is_empty(), "query dim mismatch");
474 }
475}