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 keys: Vec<Vec<u8>>,
48 vec: Vec<f32>,
49 links: Vec<Vec<u32>>,
51 dead: bool,
52}
53
54pub struct Hnsw {
56 params: HnswParams,
57 dim: usize,
58 nodes: Vec<Node>,
59 by_key: HashMap<Vec<u8>, u32>,
60 by_vec: HashMap<Vec<u32>, u32>,
65 entry: Option<u32>,
66 live: u64,
68 links_total: u64,
72 tombstones: u64,
75 seed: u64,
77}
78
79fn vec_bits(v: &[f32]) -> Vec<u32> {
81 v.iter().map(|x| x.to_bits()).collect()
82}
83
84#[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 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 pub fn dim(&self) -> usize {
118 self.dim
119 }
120
121 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 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 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 for layer in ((level as i32 + 1)..=top).rev() {
198 cur = self.greedy_at(cur, id, layer as usize);
199 }
200 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 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 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 fn search_layer_vec(&self, start: u32, tv: &[f32], layer: usize, ef: usize) -> Vec<(f32, u32)> {
252 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 fn select_diverse(&self, sorted: &[(f32, u32)], cap: usize, node_vec: &[f32]) -> Vec<u32> {
333 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 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 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 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 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 pub fn contains(&self, key: &[u8]) -> bool {
450 self.by_key.contains_key(key)
451 }
452
453 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;