1use alloc::vec::Vec;
34
35#[cfg(feature = "counters")]
36use core::cell::Cell;
37
38use plugmem_arena::{Arena, ArenaCfg, ChunkPool, ChunkPoolCfg, ListHandle, ShardMode, Slot, key};
39use xxhash_rust::xxh3::xxh3_64;
40
41use crate::error::Error;
42use crate::id::NONE_U32;
43use crate::index::vecpool::{VecPool, dot_i8};
44
45const MAX_LEVEL: usize = 16;
49
50const UPPER_SHARDS: usize = 64;
52
53const NEIGHBOR_BYTES: usize = core::mem::size_of::<u32>();
55
56const META_BYTES: usize = 2 * core::mem::size_of::<u32>();
58
59#[derive(Clone, Copy, Debug, PartialEq, Eq)]
62struct UpperSlot {
63 slot: u32,
65 level: u32,
67 handle: ListHandle,
69}
70
71impl Slot for UpperSlot {
72 const SIZE: usize = 20;
73 const KEY_LEN: usize = 8;
74
75 fn write(&self, out: &mut [u8]) {
76 key::write_u32(out, self.slot);
77 key::write_u32(&mut out[4..], self.level);
78 out[8..20].copy_from_slice(&self.handle.to_bytes());
79 }
80
81 fn read(bytes: &[u8]) -> Self {
82 Self {
83 slot: key::read_u32(bytes),
84 level: key::read_u32(&bytes[4..]),
85 handle: ListHandle::from_bytes(bytes[8..20].try_into().unwrap()),
86 }
87 }
88}
89
90#[derive(Debug, Default)]
92pub struct HnswScratch {
93 visited: Vec<u32>,
95 epoch: u32,
97 cand: Vec<(f32, u32)>,
100 found: Vec<(f32, u32)>,
102 nbrs: Vec<u32>,
104 sel: Vec<u32>,
106 pruned: Vec<u32>,
108 relink: Vec<(f32, u32)>,
110}
111
112#[inline]
115fn better(a: (f32, u32), b: (f32, u32)) -> core::cmp::Ordering {
116 a.0.total_cmp(&b.0).then(b.1.cmp(&a.1))
117}
118
119pub struct HnswGraph<'a> {
126 m: usize,
128 m0: usize,
130 max_bytes: usize,
135 level0: Vec<u32>,
137 upper: Arena<'a, UpperSlot>,
139 lists: ChunkPool<'a>,
141 entry: u32,
143 indexed: u32,
145 thresholds: [u64; MAX_LEVEL],
148 #[cfg(feature = "counters")]
151 dist_evals: Cell<u64>,
152}
153
154impl<'a> HnswGraph<'a> {
155 pub fn new(m: usize, m0: usize, max_bytes: usize) -> Result<Self, Error> {
158 let mut thresholds = [0u64; MAX_LEVEL];
159 let mut t = u64::MAX;
160 for slot in &mut thresholds {
161 t /= m as u64;
162 *slot = t;
163 }
164 Ok(Self {
165 m,
166 m0,
167 max_bytes,
168 level0: Vec::new(),
169 upper: Arena::new(
170 ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
171 )?,
172 lists: ChunkPool::new(ChunkPoolCfg::new().with_max_bytes(max_bytes)),
173 entry: NONE_U32,
174 indexed: 0,
175 thresholds,
176 #[cfg(feature = "counters")]
177 dist_evals: Cell::new(0),
178 })
179 }
180
181 fn level0_len(&self, nodes: u32) -> Result<usize, Error> {
189 let slots = u64::from(nodes) * self.m0 as u64;
190 let bytes = slots * NEIGHBOR_BYTES as u64;
191 if bytes > self.max_bytes as u64 {
192 return Err(Error::CapacityExceeded {
193 what: "hnsw level-0 neighbours",
194 });
195 }
196 Ok(slots as usize)
199 }
200
201 pub fn indexed(&self) -> u32 {
203 self.indexed
204 }
205
206 fn level_of(&self, fact: u32) -> usize {
209 let h = xxh3_64(&fact.to_le_bytes());
210 self.thresholds.iter().take_while(|&&t| h < t).count()
211 }
212
213 #[inline]
215 fn sim_q(&self, pool: &VecPool<'_>, q: (f32, &[u8]), slot: u32) -> f32 {
216 #[cfg(feature = "counters")]
217 self.dist_evals.set(self.dist_evals.get() + 1);
218 let (s, qb) = pool.quant(slot as usize);
219 q.0 * s * dot_i8(q.1, qb) as f32
220 }
221
222 #[inline]
224 fn block(&self, slot: u32) -> &[u32] {
225 let at = slot as usize * self.m0;
226 &self.level0[at..at + self.m0]
227 }
228
229 fn neighbors_into(&self, slot: u32, level: usize, out: &mut Vec<u32>) {
231 out.clear();
232 if level == 0 {
233 out.extend(
234 self.block(slot)
235 .iter()
236 .copied()
237 .take_while(|&n| n != NONE_U32),
238 );
239 return;
240 }
241 let mut kb = [0u8; 8];
242 key::write_u32(&mut kb, slot);
243 key::write_u32(&mut kb[4..], level as u32);
244 let Some(entry) = self.upper.get(&kb) else {
245 return;
246 };
247 for chunk in self.lists.iter(&entry.handle) {
248 for raw in chunk.chunks_exact(4) {
249 out.push(u32::from_le_bytes(raw.try_into().unwrap()));
250 }
251 }
252 }
253
254 #[inline]
256 fn visit(scratch: &mut HnswScratch, slot: u32) -> bool {
257 let at = slot as usize;
258 if scratch.visited[at] == scratch.epoch {
259 return true;
260 }
261 scratch.visited[at] = scratch.epoch;
262 false
263 }
264
265 fn search_layer(
269 &self,
270 pool: &VecPool<'_>,
271 q: (f32, &[u8]),
272 level: usize,
273 ep: u32,
274 ef: usize,
275 scratch: &mut HnswScratch,
276 ) {
277 scratch.epoch = scratch.epoch.wrapping_add(1);
278 if scratch.epoch == 0 {
279 scratch.visited.fill(u32::MAX);
281 scratch.epoch = 1;
282 }
283 scratch
284 .visited
285 .resize(self.indexed as usize, scratch.epoch.wrapping_sub(1));
286 scratch.cand.clear();
287 scratch.found.clear();
288
289 Self::visit(scratch, ep);
290 let s = self.sim_q(pool, q, ep);
291 scratch.cand.push((s, ep));
292 scratch.found.push((s, ep));
293
294 while let Some(best) = scratch.cand.pop() {
295 if scratch.found.len() >= ef && better(best, scratch.found[0]).is_lt() {
298 break;
299 }
300 let nbrs = core::mem::take(&mut scratch.nbrs);
301 let mut nbrs = nbrs;
302 self.neighbors_into(best.1, level, &mut nbrs);
303 for &nb in &nbrs {
304 if Self::visit(scratch, nb) {
305 continue;
306 }
307 let s = self.sim_q(pool, q, nb);
308 let entry = (s, nb);
309 if scratch.found.len() < ef || better(entry, scratch.found[0]).is_gt() {
310 let at = scratch.found.partition_point(|&e| better(e, entry).is_lt());
311 scratch.found.insert(at, entry);
312 if scratch.found.len() > ef {
313 scratch.found.remove(0);
314 }
315 let at = scratch.cand.partition_point(|&e| better(e, entry).is_lt());
316 scratch.cand.insert(at, entry);
317 }
318 }
319 scratch.nbrs = nbrs;
320 }
321 }
322
323 fn select_neighbors(&self, pool: &VecPool<'_>, cap: usize, scratch: &mut HnswScratch) {
328 scratch.sel.clear();
329 scratch.pruned.clear();
330 for i in (0..scratch.found.len()).rev() {
331 let (sim, cand) = scratch.found[i];
332 if scratch.sel.len() >= cap {
333 break;
334 }
335 let dominated = scratch.sel.iter().any(|&kept| {
336 #[cfg(feature = "counters")]
337 self.dist_evals.set(self.dist_evals.get() + 1);
338 pool.sim(cand, kept) > sim
339 });
340 if dominated {
341 scratch.pruned.push(cand);
342 } else {
343 scratch.sel.push(cand);
344 }
345 }
346 for &p in scratch.pruned.iter() {
347 if scratch.sel.len() >= cap {
348 break;
349 }
350 scratch.sel.push(p);
351 }
352 }
353
354 fn write_list(&mut self, slot: u32, level: usize, sel: &[u32]) -> Result<(), Error> {
357 if level == 0 {
358 let at = slot as usize * self.m0;
359 let block = &mut self.level0[at..at + self.m0];
360 block.fill(NONE_U32);
361 block[..sel.len()].copy_from_slice(sel);
362 return Ok(());
363 }
364 let mut kb = [0u8; 8];
365 key::write_u32(&mut kb, slot);
366 key::write_u32(&mut kb[4..], level as u32);
367 let mut handle = match self.upper.get(&kb) {
368 Some(entry) => {
369 let mut h = entry.handle;
370 self.lists.free(&mut h);
371 h
372 }
373 None => ListHandle::EMPTY,
374 };
375 for &n in sel {
376 self.lists.push(&mut handle, &n.to_le_bytes())?;
377 }
378 let updated = UpperSlot {
379 slot,
380 level: level as u32,
381 handle,
382 };
383 if self.upper.contains(&kb) {
384 let payload = self.upper.payload_mut(&kb).expect("checked above");
385 let mut full = [0u8; UpperSlot::SIZE];
386 updated.write(&mut full);
387 payload.copy_from_slice(&full[UpperSlot::KEY_LEN..]);
388 } else {
389 self.upper.insert(&updated)?;
390 }
391 Ok(())
392 }
393
394 fn add_link(
397 &mut self,
398 pool: &VecPool<'_>,
399 v: u32,
400 new: u32,
401 level: usize,
402 scratch: &mut HnswScratch,
403 ) -> Result<(), Error> {
404 let cap = if level == 0 { self.m0 } else { self.m };
405 let mut nbrs = core::mem::take(&mut scratch.nbrs);
406 self.neighbors_into(v, level, &mut nbrs);
407 if nbrs.len() < cap {
408 nbrs.push(new);
409 let sel = core::mem::take(&mut scratch.sel);
410 let mut sel = sel;
411 sel.clear();
412 sel.extend_from_slice(&nbrs);
413 let res = self.write_list(v, level, &sel);
414 scratch.sel = sel;
415 scratch.nbrs = nbrs;
416 return res;
417 }
418 let vq = pool.quant(v as usize);
420 scratch.relink.clear();
421 for &n in nbrs.iter().chain(core::iter::once(&new)) {
422 let s = self.sim_q(pool, (vq.0, vq.1), n);
423 scratch.relink.push((s, n));
424 }
425 scratch.nbrs = nbrs;
426 scratch.relink.sort_unstable_by(|a, b| better(*a, *b));
427 scratch.found.clear();
428 scratch.found.extend_from_slice(&scratch.relink);
429 self.select_neighbors(pool, cap, scratch);
430 let sel = core::mem::take(&mut scratch.sel);
431 let res = self.write_list(v, level, &sel);
432 scratch.sel = sel;
433 res
434 }
435
436 pub fn insert_bulk(
440 &mut self,
441 pool: &VecPool<'_>,
442 upto: u32,
443 ef_construction: usize,
444 scratch: &mut HnswScratch,
445 ) -> Result<(), Error> {
446 debug_assert!(upto as usize <= pool.len());
447 self.level0.resize(self.level0_len(upto)?, NONE_U32);
448 for slot in self.indexed..upto {
449 self.insert_one(pool, slot, ef_construction, scratch)?;
450 self.indexed = slot + 1;
453 }
454 Ok(())
455 }
456
457 pub(crate) fn to_owned(&self, max_bytes: usize) -> Result<HnswGraph<'static>, Error> {
461 let meta = self.dump_meta();
462 let level0 = self.dump_level0();
463 let [upper_meta, upper_pool, lists_meta, lists_pool] = self.dump_upper();
464 HnswGraph::from_parts(
465 self.m,
466 self.m0,
467 max_bytes,
468 &meta,
469 &level0,
470 &upper_meta,
471 &upper_pool,
472 &lists_meta,
473 &lists_pool,
474 )
475 }
476
477 fn insert_one(
478 &mut self,
479 pool: &VecPool<'_>,
480 slot: u32,
481 ef_construction: usize,
482 scratch: &mut HnswScratch,
483 ) -> Result<(), Error> {
484 let level = self.level_of(pool.slot_fact(slot as usize));
485 if self.entry == NONE_U32 {
486 self.entry = slot;
487 return Ok(());
488 }
489 let q = pool.quant(slot as usize);
490 let q = (q.0, q.1);
491 let top = self.level_of(pool.slot_fact(self.entry as usize));
492 let mut ep = self.entry;
493 let mut lev = top;
495 while lev > level {
496 self.search_layer(pool, q, lev, ep, 1, scratch);
497 ep = scratch.found.last().expect("entry is always found").1;
498 lev -= 1;
499 }
500 let mut lev = level.min(top);
502 loop {
503 self.search_layer(pool, q, lev, ep, ef_construction, scratch);
504 ep = scratch.found.last().expect("entry is always found").1;
505 let cap = if lev == 0 { self.m0 } else { self.m };
506 self.select_neighbors(pool, cap, scratch);
507 let sel = core::mem::take(&mut scratch.sel);
508 self.write_list(slot, lev, &sel)?;
509 for &nb in &sel {
510 self.add_link(pool, nb, slot, lev, scratch)?;
511 }
512 scratch.sel = sel;
513 if lev == 0 {
514 break;
515 }
516 lev -= 1;
517 }
518 if level > top {
519 self.entry = slot;
520 }
521 Ok(())
522 }
523
524 pub fn search(
534 &self,
535 pool: &VecPool<'_>,
536 query: &[f32],
537 ef: usize,
538 vec_scratch: &mut crate::index::vecpool::VecScratch,
539 scratch: &mut HnswScratch,
540 out: &mut Vec<(u32, f32)>,
541 ) -> Result<(), Error> {
542 pool.quantize_query(query, vec_scratch)?;
543 let q = pool.quantized(vec_scratch);
544 self.search_quantized(pool, q, ef, scratch, out);
545 Ok(())
546 }
547
548 pub(crate) fn search_quantized(
552 &self,
553 pool: &VecPool<'_>,
554 q: (f32, &[u8]),
555 ef: usize,
556 scratch: &mut HnswScratch,
557 out: &mut Vec<(u32, f32)>,
558 ) {
559 out.clear();
560 if self.entry == NONE_U32 {
561 return;
562 }
563 let mut ep = self.entry;
564 let top = self.level_of(pool.slot_fact(self.entry as usize));
565 for lev in (1..=top).rev() {
566 self.search_layer(pool, q, lev, ep, 1, scratch);
567 ep = scratch.found.last().expect("entry is always found").1;
568 }
569 self.search_layer(pool, q, 0, ep, ef.max(1), scratch);
570 for &(sim, slot) in scratch.found.iter().rev() {
571 out.push((slot, sim));
572 }
573 }
574
575 pub(crate) fn pool_bytes(&self) -> usize {
577 self.level0.len() * NEIGHBOR_BYTES + self.upper.pool_bytes() + self.lists.pool_bytes()
578 }
579
580 pub(crate) fn remapped(
586 &self,
587 map: &[u32],
588 new_pool: &VecPool<'_>,
589 max_bytes: usize,
590 ) -> Result<HnswGraph<'static>, Error> {
591 let mut g: HnswGraph<'static> = HnswGraph::new(self.m, self.m0, max_bytes)?;
594 let old_indexed = self.indexed as usize;
595 let new_indexed = map[..old_indexed]
596 .iter()
597 .filter(|&&m| m != NONE_U32)
598 .count() as u32;
599 g.level0 = alloc::vec![NONE_U32; g.level0_len(new_indexed)?];
600 g.indexed = new_indexed;
601 let mut nbrs = Vec::new();
602 let mut sel = Vec::new();
603 for old in 0..old_indexed as u32 {
604 let new = map[old as usize];
605 if new == NONE_U32 {
606 continue;
607 }
608 self.neighbors_into(old, 0, &mut nbrs);
609 sel.clear();
610 sel.extend(
611 nbrs.iter()
612 .map(|&n| map[n as usize])
613 .filter(|&n| n != NONE_U32),
614 );
615 g.write_list(new, 0, &sel)?;
616 let levels = g.level_of(new_pool.slot_fact(new as usize));
617 for level in 1..=levels {
618 self.neighbors_into(old, level, &mut nbrs);
619 if nbrs.is_empty() {
620 continue;
621 }
622 sel.clear();
623 sel.extend(
624 nbrs.iter()
625 .map(|&n| map[n as usize])
626 .filter(|&n| n != NONE_U32),
627 );
628 g.write_list(new, level, &sel)?;
629 }
630 }
631 g.entry = if self.entry != NONE_U32 && map[self.entry as usize] != NONE_U32 {
632 map[self.entry as usize]
633 } else {
634 let mut best = NONE_U32;
637 let mut best_level = 0usize;
638 for slot in 0..new_indexed {
639 let level = g.level_of(new_pool.slot_fact(slot as usize));
640 if best == NONE_U32 || level > best_level {
641 best = slot;
642 best_level = level;
643 }
644 }
645 best
646 };
647 Ok(g)
648 }
649
650 pub(crate) fn dump_meta(&self) -> Vec<u8> {
652 let mut out = Vec::with_capacity(META_BYTES);
653 out.extend_from_slice(&self.entry.to_le_bytes());
654 out.extend_from_slice(&self.indexed.to_le_bytes());
655 out
656 }
657
658 pub(crate) fn dump_level0(&self) -> Vec<u8> {
660 let mut out = Vec::with_capacity(self.level0.len() * NEIGHBOR_BYTES);
661 for &n in &self.level0 {
662 out.extend_from_slice(&n.to_le_bytes());
663 }
664 out
665 }
666
667 pub(crate) fn dump_upper(&self) -> [Vec<u8>; 4] {
669 let (mut am, mut ap) = (Vec::new(), Vec::new());
670 self.upper.dump_meta(&mut am);
671 self.upper.dump_pool(&mut ap);
672 let (mut cm, mut cp) = (Vec::new(), Vec::new());
673 self.lists.dump_meta(&mut cm);
674 self.lists.dump_pool(&mut cp);
675 [am, ap, cm, cp]
676 }
677
678 #[allow(clippy::too_many_arguments)]
682 pub(crate) fn from_parts(
683 m: usize,
684 m0: usize,
685 max_bytes: usize,
686 meta: &[u8],
687 level0: &[u8],
688 upper_meta: &[u8],
689 upper_pool: &[u8],
690 lists_meta: &[u8],
691 lists_pool: &[u8],
692 ) -> Result<Self, Error> {
693 let mut g = Self::new(m, m0, max_bytes)?;
694 if meta.len() != META_BYTES {
695 return Err(Error::Corrupt("hnsw meta section has a wrong length"));
696 }
697 g.entry = u32::from_le_bytes(meta[0..4].try_into().unwrap());
698 g.indexed = u32::from_le_bytes(meta[4..8].try_into().unwrap());
699 if level0.len() as u64 != u64::from(g.indexed) * m0 as u64 * NEIGHBOR_BYTES as u64 {
700 return Err(Error::Corrupt("hnsw level0 length mismatch"));
701 }
702 g.level0 = level0
703 .chunks_exact(NEIGHBOR_BYTES)
704 .map(|b| u32::from_le_bytes(b.try_into().unwrap()))
705 .collect();
706 g.upper = Arena::load(
707 ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
708 upper_meta,
709 upper_pool,
710 )?;
711 g.lists = ChunkPool::load(
712 ChunkPoolCfg::new().with_max_bytes(max_bytes),
713 lists_meta,
714 lists_pool,
715 )?;
716 Ok(g)
717 }
718
719 #[allow(clippy::too_many_arguments)]
725 pub(crate) fn from_parts_borrowed(
726 m: usize,
727 m0: usize,
728 max_bytes: usize,
729 meta: &[u8],
730 level0: &[u8],
731 upper_meta: &[u8],
732 upper_pool: &'a [u8],
733 lists_meta: &[u8],
734 lists_pool: &'a [u8],
735 ) -> Result<Self, Error> {
736 let mut g = Self::new(m, m0, max_bytes)?;
737 if meta.len() != META_BYTES {
738 return Err(Error::Corrupt("hnsw meta section has a wrong length"));
739 }
740 g.entry = u32::from_le_bytes(meta[0..4].try_into().unwrap());
741 g.indexed = u32::from_le_bytes(meta[4..8].try_into().unwrap());
742 if level0.len() as u64 != u64::from(g.indexed) * m0 as u64 * NEIGHBOR_BYTES as u64 {
743 return Err(Error::Corrupt("hnsw level0 length mismatch"));
744 }
745 g.level0 = level0
746 .chunks_exact(NEIGHBOR_BYTES)
747 .map(|b| u32::from_le_bytes(b.try_into().unwrap()))
748 .collect();
749 g.upper = Arena::load_borrowed(
750 ArenaCfg::new(UPPER_SHARDS, ShardMode::Uniform).with_max_bytes(max_bytes),
751 upper_meta,
752 upper_pool,
753 )?;
754 g.lists = ChunkPool::load_borrowed(
755 ChunkPoolCfg::new().with_max_bytes(max_bytes),
756 lists_meta,
757 lists_pool,
758 )?;
759 Ok(g)
760 }
761
762 pub(crate) fn validate(&self, pool: &VecPool<'_>) -> Result<(), Error> {
769 if self.indexed as usize > pool.len() {
770 return Err(Error::Corrupt("hnsw indexes more slots than the pool"));
771 }
772 if self.level0.len() != self.indexed as usize * self.m0 {
773 return Err(Error::Corrupt("hnsw level0 disagrees with indexed"));
774 }
775 if self.indexed == 0 {
776 if self.entry != NONE_U32 || !self.upper.is_empty() || self.lists.chunks() != 0 {
777 return Err(Error::Corrupt("hnsw empty graph carries state"));
778 }
779 return Ok(());
780 }
781 if self.entry >= self.indexed {
782 return Err(Error::Corrupt("hnsw entry out of range"));
783 }
784 for slot in 0..self.indexed {
785 let block = self.block(slot);
786 let mut ended = false;
787 for &n in block {
788 if n == NONE_U32 {
789 ended = true;
790 continue;
791 }
792 if ended {
793 return Err(Error::Corrupt("hnsw level0 padding is not canonical"));
794 }
795 if n >= self.indexed || n == slot {
796 return Err(Error::Corrupt("hnsw level0 neighbor out of range"));
797 }
798 }
799 }
800 let mut visited = alloc::vec![false; self.lists.chunks()];
801 for entry in self.upper.iter() {
802 if entry.slot >= self.indexed {
803 return Err(Error::Corrupt("hnsw upper handle out of range"));
804 }
805 let max_level = self.level_of(pool.slot_fact(entry.slot as usize));
806 if entry.level == 0 || entry.level as usize > max_level {
807 return Err(Error::Corrupt("hnsw upper level disagrees with the hash"));
808 }
809 self.lists.validate_chain(&entry.handle, &mut visited)?;
810 let mut count = 0u32;
811 for chunk in self.lists.iter(&entry.handle) {
812 if !chunk.len().is_multiple_of(4) {
813 return Err(Error::Corrupt("hnsw upper list is not a slot sequence"));
814 }
815 for raw in chunk.chunks_exact(4) {
816 let n = u32::from_le_bytes(raw.try_into().unwrap());
817 if n >= self.indexed || n == entry.slot {
818 return Err(Error::Corrupt("hnsw upper neighbor out of range"));
819 }
820 count += 1;
821 }
822 }
823 if count != entry.handle.len() || count as usize > self.m {
824 return Err(Error::Corrupt("hnsw upper list disagrees with its handle"));
825 }
826 }
827 if self.lists.orphan_count(&visited) != 0 {
828 return Err(Error::Corrupt("hnsw list pool has orphan chunks"));
829 }
830 Ok(())
831 }
832
833 #[cfg(feature = "counters")]
835 pub fn dist_evals(&self) -> u64 {
836 self.dist_evals.get()
837 }
838
839 #[cfg(feature = "counters")]
841 pub fn reset_dist_evals(&self) {
842 self.dist_evals.set(0);
843 }
844}
845
846#[cfg(test)]
847mod tests {
848 use super::*;
849 use crate::id::FactId;
850 use alloc::vec;
851
852 struct Lcg(u64);
854 impl Lcg {
855 fn next(&mut self) -> f32 {
856 self.0 = self
857 .0
858 .wrapping_mul(6_364_136_223_846_793_005)
859 .wrapping_add(1_442_695_040_888_963_407);
860 ((self.0 >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
861 }
862 }
863
864 fn cluster_pool(n: usize, dim: usize, clusters: usize, seed: u64) -> VecPool<'static> {
866 let mut rng = Lcg(seed);
867 let centers: Vec<Vec<f32>> = (0..clusters)
868 .map(|_| (0..dim).map(|_| rng.next()).collect())
869 .collect();
870 let mut pool = VecPool::new(dim, usize::MAX);
871 for i in 0..n {
872 let c = ¢ers[i % clusters];
873 let v: Vec<f32> = c.iter().map(|&x| x + rng.next() * 0.3).collect();
874 pool.push(FactId(i as u32), &v).unwrap();
875 }
876 pool
877 }
878
879 fn build(pool: &VecPool<'_>, m: usize, m0: usize) -> HnswGraph<'static> {
881 let mut g = HnswGraph::new(m, m0, usize::MAX).unwrap();
882 let mut scratch = HnswScratch::default();
883 g.insert_bulk(pool, pool.len() as u32, 200, &mut scratch)
884 .unwrap();
885 g
886 }
887
888 fn brute_force(pool: &VecPool<'_>, q: u32, k: usize) -> Vec<u32> {
890 let mut all: Vec<(f32, u32)> = (0..pool.len() as u32)
891 .map(|i| (pool.sim(q, i), i))
892 .collect();
893 all.sort_unstable_by(|a, b| better(*b, *a));
894 all.into_iter().take(k).map(|(_, s)| s).collect()
895 }
896
897 #[test]
900 #[cfg_attr(miri, ignore)] fn levels_are_geometric_and_pure() {
902 let g = HnswGraph::new(16, 32, usize::MAX).unwrap();
903 let n = 100_000u32;
904 let mut per_level = [0usize; 4];
905 for fact in 0..n {
906 let l = g.level_of(fact).min(3);
907 per_level[l] += 1;
908 }
909 let at_least_1: usize = per_level[1..].iter().sum();
911 assert!(
912 (4_000..9_000).contains(&at_least_1),
913 "level>=1 count {at_least_1} is out of band"
914 );
915 let at_least_2: usize = per_level[2..].iter().sum();
916 assert!(
917 (150..800).contains(&at_least_2),
918 "level>=2 count {at_least_2} is out of band"
919 );
920 assert_eq!(g.level_of(42), g.level_of(42));
922 }
923
924 #[test]
927 #[cfg_attr(miri, ignore)] fn recall_against_brute_force() {
929 let dim = 32;
930 let pool = cluster_pool(2_000, dim, 64, 0xA11CE);
931 let g = build(&pool, 16, 32);
932 let mut scratch = HnswScratch::default();
933 let mut out = Vec::new();
934 let mut hits = 0usize;
935 let mut total = 0usize;
936 for q in (0..2_000u32).step_by(97) {
937 let truth = brute_force(&pool, q, 10);
938 let (scale, qb) = pool.quant(q as usize);
939 g.search_quantized(&pool, (scale, qb), 64, &mut scratch, &mut out);
940 let got: Vec<u32> = out.iter().take(10).map(|&(s, _)| s).collect();
941 hits += truth.iter().filter(|t| got.contains(t)).count();
942 total += truth.len();
943 }
944 let recall = hits as f64 / total as f64;
945 assert!(recall >= 0.9, "recall@10 {recall} below the 0.9 gate");
946 }
947
948 #[test]
951 #[cfg_attr(miri, ignore)] fn build_is_deterministic() {
953 let pool = cluster_pool(600, 24, 16, 7);
954 let a = build(&pool, 8, 16);
955 let b = build(&pool, 8, 16);
956 assert_eq!(a.level0, b.level0);
957 assert_eq!(a.entry, b.entry);
958 assert_eq!(a.indexed, b.indexed);
959 let (mut am, mut bm) = (Vec::new(), Vec::new());
960 a.upper.dump_meta(&mut am);
961 b.upper.dump_meta(&mut bm);
962 assert_eq!(am, bm);
963 let (mut ap, mut bp) = (Vec::new(), Vec::new());
964 a.lists.dump_pool(&mut ap);
965 b.lists.dump_pool(&mut bp);
966 assert_eq!(ap, bp);
967 }
968
969 #[test]
971 #[cfg_attr(miri, ignore)] fn degree_caps_hold() {
973 let pool = cluster_pool(800, 16, 8, 3);
974 let g = build(&pool, 6, 12);
975 let mut nbrs = Vec::new();
976 for slot in 0..g.indexed() {
977 g.neighbors_into(slot, 0, &mut nbrs);
978 assert!(nbrs.len() <= 12);
979 assert!(!nbrs.contains(&slot));
981 let mut sorted = nbrs.clone();
982 sorted.sort_unstable();
983 sorted.dedup();
984 assert_eq!(sorted.len(), nbrs.len());
985 assert!(nbrs.iter().all(|&n| n < g.indexed()));
986 for level in 1..=g.level_of(pool.slot_fact(slot as usize)) {
987 g.neighbors_into(slot, level, &mut nbrs);
988 assert!(nbrs.len() <= 6, "level {level} degree overflow");
989 }
990 }
991 }
992
993 #[test]
997 #[cfg_attr(miri, ignore)] fn remap_survives_a_dead_entry() {
999 let pool = cluster_pool(300, 16, 8, 21);
1000 let g = build(&pool, 6, 12);
1001 let entry = g.entry;
1002 let mut map = alloc::vec![NONE_U32; 300];
1005 let mut new_pool = VecPool::new(16, usize::MAX);
1006 let mut next = 0u32;
1007 for old in 0..300u32 {
1008 if old == entry || old % 7 == 0 {
1009 continue;
1010 }
1011 map[old as usize] = next;
1012 new_pool.copy_slot(&pool, old);
1013 next += 1;
1014 }
1015 let remapped = g.remapped(&map, &new_pool, usize::MAX).unwrap();
1016 assert_eq!(remapped.indexed(), next);
1017 assert_ne!(remapped.entry, NONE_U32, "a survivor takes the entry");
1018 remapped.validate(&new_pool).unwrap();
1019 let probe_old = 1u32; let probe_old = if probe_old == entry { 2 } else { probe_old };
1022 let probe_new = map[probe_old as usize];
1023 let (scale, qb) = new_pool.quant(probe_new as usize);
1024 let mut scratch = HnswScratch::default();
1025 let mut out = Vec::new();
1026 remapped.search_quantized(&new_pool, (scale, qb), 32, &mut scratch, &mut out);
1027 assert_eq!(out[0].0, probe_new);
1028 }
1029
1030 #[test]
1034 fn validate_rejects_malformed_graphs() {
1035 let pool = cluster_pool(50, 8, 4, 5);
1036 let g = build(&pool, 4, 8);
1037 let dump = (g.dump_meta(), g.dump_level0(), g.dump_upper());
1038 let load = |meta: &[u8], level0: &[u8]| {
1039 HnswGraph::from_parts(
1040 4,
1041 8,
1042 usize::MAX,
1043 meta,
1044 level0,
1045 &dump.2[0],
1046 &dump.2[1],
1047 &dump.2[2],
1048 &dump.2[3],
1049 )
1050 };
1051 load(&dump.0, &dump.1).unwrap().validate(&pool).unwrap();
1053 let mut meta = dump.0.clone();
1055 meta[0..4].copy_from_slice(&999u32.to_le_bytes());
1056 assert!(load(&meta, &dump.1).unwrap().validate(&pool).is_err());
1057 let mut level0 = dump.1.clone();
1059 level0[0..4].copy_from_slice(&500u32.to_le_bytes());
1060 assert!(load(&dump.0, &level0).unwrap().validate(&pool).is_err());
1061 let mut level0 = dump.1.clone();
1063 level0[0..4].copy_from_slice(&NONE_U32.to_le_bytes());
1064 level0[4..8].copy_from_slice(&1u32.to_le_bytes());
1065 assert!(load(&dump.0, &level0).unwrap().validate(&pool).is_err());
1066 assert!(load(&dump.0, &dump.1[..dump.1.len() - 4]).is_err());
1068 let mut meta = dump.0.clone();
1070 meta[4..8].copy_from_slice(&0u32.to_le_bytes());
1071 assert!(load(&meta, &[]).unwrap().validate(&pool).is_err());
1072 }
1073
1074 #[test]
1076 fn tiny_graphs() {
1077 let pool = cluster_pool(1, 8, 1, 1);
1078 let mut g = HnswGraph::new(4, 8, usize::MAX).unwrap();
1079 let mut scratch = HnswScratch::default();
1080 let mut out = vec![(0u32, 0.0f32)];
1081 let (scale, qb) = pool.quant(0);
1082 g.search_quantized(&pool, (scale, qb), 8, &mut scratch, &mut out);
1083 assert!(out.is_empty(), "an empty graph must answer empty");
1084 g.insert_bulk(&pool, 1, 50, &mut scratch).unwrap();
1085 g.search_quantized(&pool, (scale, qb), 8, &mut scratch, &mut out);
1086 assert_eq!(out.len(), 1);
1087 assert_eq!(out[0].0, 0);
1088 }
1089}