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 seed: u64,
70}
71
72fn vec_bits(v: &[f32]) -> Vec<u32> {
74 v.iter().map(|x| x.to_bits()).collect()
75}
76
77#[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 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 pub fn dim(&self) -> usize {
109 self.dim
110 }
111
112 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 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 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 for layer in ((level as i32 + 1)..=top).rev() {
188 cur = self.greedy_at(cur, id, layer as usize);
189 }
190 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 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 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 fn search_layer_vec(&self, start: u32, tv: &[f32], layer: usize, ef: usize) -> Vec<(f32, u32)> {
241 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 fn select_diverse(&self, sorted: &[(f32, u32)], cap: usize, node_vec: &[f32]) -> Vec<u32> {
322 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 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 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 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 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 pub fn contains(&self, key: &[u8]) -> bool {
437 self.by_key.contains_key(key)
438 }
439
440 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 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;