1use alloc::vec::Vec;
33
34#[cfg(feature = "counters")]
35use core::cell::Cell;
36
37use crate::error::Error;
38use crate::id::FactId;
39
40const FACT_BYTES: usize = core::mem::size_of::<u32>();
42
43const SCALE_BYTES: usize = core::mem::size_of::<f32>();
45
46const HEAD: usize = FACT_BYTES + SCALE_BYTES;
48
49const SIG_WORD_BYTES: usize = core::mem::size_of::<u64>();
52
53#[derive(Debug, Default)]
56pub struct VecScratch {
57 cand: Vec<(u32, u32)>,
59 top: Vec<(f32, u32)>,
61 query: Vec<u8>,
63}
64
65impl VecScratch {
66 pub fn new() -> Self {
68 Self::default()
69 }
70}
71
72#[derive(Debug)]
83pub struct VecPool<'a> {
84 base: &'a [u8],
86 tail: Vec<u8>,
88 dim: usize,
89 max_bytes: usize,
90 #[cfg(feature = "counters")]
92 dots: Cell<u64>,
93}
94
95impl<'a> VecPool<'a> {
96 #[inline]
98 fn words(dim: usize) -> usize {
99 dim.div_ceil(64)
100 }
101
102 pub fn new(dim: usize, max_bytes: usize) -> Self {
105 Self {
106 base: &[],
107 tail: Vec::new(),
108 dim,
109 max_bytes,
110 #[cfg(feature = "counters")]
111 dots: Cell::new(0),
112 }
113 }
114
115 #[inline]
117 pub fn stride(&self) -> usize {
118 HEAD + Self::words(self.dim) * SIG_WORD_BYTES + self.dim
119 }
120
121 #[inline]
123 fn pool_len(&self) -> usize {
124 self.base.len() + self.tail.len()
125 }
126
127 #[inline]
132 pub(crate) fn slot_bytes(&self, i: usize) -> &[u8] {
133 let stride = self.stride();
134 let start = i * stride;
135 let base_len = self.base.len();
136 if start < base_len {
137 &self.base[start..start + stride]
138 } else {
139 let at = start - base_len;
140 &self.tail[at..at + stride]
141 }
142 }
143
144 #[inline]
146 pub fn len(&self) -> usize {
147 let pool_len = self.pool_len();
148 if pool_len == 0 {
149 0
150 } else {
151 pool_len / self.stride()
152 }
153 }
154
155 pub fn is_empty(&self) -> bool {
157 self.pool_len() == 0
158 }
159
160 pub fn pool_bytes(&self) -> usize {
162 self.pool_len()
163 }
164
165 #[inline]
167 pub fn slot_fact(&self, i: usize) -> u32 {
168 let slot = self.slot_bytes(i);
169 u32::from_le_bytes(slot[..FACT_BYTES].try_into().unwrap())
170 }
171
172 #[inline]
174 fn slot_scale(&self, i: usize) -> f32 {
175 let slot = self.slot_bytes(i);
176 f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap())
177 }
178
179 #[inline]
183 pub(crate) fn quant(&self, i: usize) -> (f32, &[u8]) {
184 let stride = self.stride();
185 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
186 let slot = self.slot_bytes(i);
187 let scale = f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap());
188 (scale, &slot[q_off..stride])
189 }
190
191 #[inline]
195 pub(crate) fn sim(&self, a: u32, b: u32) -> f32 {
196 self.cosine_at(a as usize, b as usize)
197 }
198
199 fn encode_slot(&self, fact: u32, v: &[f32], out: &mut [u8]) -> Result<(), Error> {
207 if v.len() != self.dim {
208 return Err(Error::DimMismatch {
209 got: v.len(),
210 want: self.dim,
211 });
212 }
213 let mut norm_sq = 0.0f32;
214 for &x in v {
215 if !x.is_finite() {
216 return Err(Error::Invalid("vector must be finite"));
217 }
218 norm_sq += x * x;
219 }
220 let norm = libm::sqrtf(norm_sq);
223 if norm <= 0.0 {
224 return Err(Error::Invalid("vector must be nonzero"));
225 }
226 let inv_norm = 1.0 / norm;
227 let mut max_abs = 0.0f32;
228 for &x in v {
229 max_abs = max_abs.max(libm::fabsf(x * inv_norm));
230 }
231 let scale = max_abs / 127.0;
233 out[..FACT_BYTES].copy_from_slice(&fact.to_le_bytes());
234 out[FACT_BYTES..HEAD].copy_from_slice(&scale.to_le_bytes());
235 let words = Self::words(self.dim);
236 let q_off = HEAD + words * SIG_WORD_BYTES;
237 for (i, &x) in v.iter().enumerate() {
238 let qf = libm::roundf((x * inv_norm) / scale);
239 let qi = qf.clamp(-127.0, 127.0) as i32 as i8;
240 out[q_off + i] = qi as u8;
241 }
242 for w in 0..words {
244 let mut word = 0u64;
245 for b in 0..64 {
246 let i = w * 64 + b;
247 if i >= self.dim {
248 break;
249 }
250 if out[q_off + i] as i8 >= 0 {
251 word |= 1 << b;
252 }
253 }
254 out[HEAD + w * SIG_WORD_BYTES..HEAD + w * SIG_WORD_BYTES + SIG_WORD_BYTES]
255 .copy_from_slice(&word.to_le_bytes());
256 }
257 Ok(())
258 }
259
260 pub fn push(&mut self, fact: FactId, v: &[f32]) -> Result<u32, Error> {
268 let stride = self.stride();
269 let pool_len = self.pool_len();
270 if pool_len + stride > self.max_bytes {
271 return Err(Error::CapacityExceeded { what: "vectors" });
272 }
273 let index = u32::try_from(pool_len / stride).map_err(|_| Error::CapacityExceeded {
274 what: "vector slots",
275 })?;
276 let mut tail = core::mem::take(&mut self.tail);
280 let at = tail.len();
281 tail.resize(at + stride, 0);
282 let res = match self.encode_slot(fact.0, v, &mut tail[at..]) {
283 Ok(()) => Ok(index),
284 Err(e) => {
285 tail.truncate(at);
287 Err(e)
288 }
289 };
290 self.tail = tail;
291 res
292 }
293
294 pub(crate) fn quantized<'s>(&self, scratch: &'s VecScratch) -> (f32, &'s [u8]) {
298 let stride = self.stride();
299 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
300 debug_assert_eq!(scratch.query.len(), stride);
301 (
302 f32::from_le_bytes(scratch.query[FACT_BYTES..HEAD].try_into().unwrap()),
303 &scratch.query[q_off..stride],
304 )
305 }
306
307 pub fn quantize_query(&self, v: &[f32], scratch: &mut VecScratch) -> Result<(), Error> {
310 let stride = self.stride();
311 scratch.query.clear();
312 scratch.query.resize(stride, 0);
313 let mut buf = core::mem::take(&mut scratch.query);
314 let res = self.encode_slot(0, v, &mut buf);
315 scratch.query = buf;
316 res
317 }
318
319 pub(crate) fn copy_slot(&mut self, src: &VecPool<'_>, i: u32) -> u32 {
324 debug_assert_eq!(self.dim, src.dim, "copy_slot across differing dims");
325 let stride = self.stride();
326 let index = (self.pool_len() / stride) as u32;
327 self.tail.extend_from_slice(src.slot_bytes(i as usize));
328 index
329 }
330
331 fn cosine_at(&self, a: usize, b: usize) -> f32 {
334 let stride = self.stride();
335 let q_off = HEAD + Self::words(self.dim) * SIG_WORD_BYTES;
336 let (sa, sb) = (self.slot_bytes(a), self.slot_bytes(b));
337 let dot = dot_i8(&sa[q_off..stride], &sb[q_off..stride]);
338 self.slot_scale(a) * self.slot_scale(b) * dot as f32
339 }
340
341 pub fn cosine_slots(&self, a: u32, b: u32) -> f32 {
344 let n = self.len();
345 if a as usize >= n || b as usize >= n {
346 return 0.0;
347 }
348 self.cosine_at(a as usize, b as usize)
349 }
350
351 pub fn search(
356 &self,
357 query: &[f32],
358 k: usize,
359 admit: &mut dyn FnMut(FactId) -> bool,
360 scratch: &mut VecScratch,
361 out: &mut Vec<(FactId, f32)>,
362 ) -> Result<(), Error> {
363 out.clear();
364 let n = self.len();
365 if n == 0 || k == 0 {
366 return Ok(());
367 }
368 self.quantize_query(query, scratch)?;
369 let stride = self.stride();
370 let words = Self::words(self.dim);
371 let q_off = HEAD + words * SIG_WORD_BYTES;
372
373 let VecScratch {
375 cand, top, query, ..
376 } = scratch;
377 let q_sig = &query[HEAD..HEAD + words * SIG_WORD_BYTES];
378 cand.clear();
379 cand.reserve(n);
380 for i in 0..n {
381 let slot = self.slot_bytes(i);
382 let s_sig = &slot[HEAD..HEAD + words * SIG_WORD_BYTES];
383 let mut ham = 0u32;
384 for w in 0..words {
385 let a = u64::from_le_bytes(
386 q_sig[w * SIG_WORD_BYTES..w * SIG_WORD_BYTES + SIG_WORD_BYTES]
387 .try_into()
388 .unwrap(),
389 );
390 let b = u64::from_le_bytes(
391 s_sig[w * SIG_WORD_BYTES..w * SIG_WORD_BYTES + SIG_WORD_BYTES]
392 .try_into()
393 .unwrap(),
394 );
395 ham += (a ^ b).count_ones();
396 }
397 cand.push((ham, i as u32));
398 }
399 let c = (4 * k).max(64).min(n);
400 if cand.len() > c {
401 cand.select_nth_unstable(c - 1);
402 }
403
404 let q_scale = f32::from_le_bytes(query[FACT_BYTES..HEAD].try_into().unwrap());
406 let q_q = &query[q_off..q_off + self.dim];
407 top.clear();
408 #[cfg(feature = "counters")]
409 let mut dots = 0u64;
410 for &(_, slot) in cand[..c].iter() {
411 let sb = self.slot_bytes(slot as usize);
412 let fact = FactId(u32::from_le_bytes(sb[..FACT_BYTES].try_into().unwrap()));
413 if !admit(fact) {
414 continue;
415 }
416 let s_scale = f32::from_le_bytes(sb[FACT_BYTES..HEAD].try_into().unwrap());
417 let dot = dot_i8(q_q, &sb[q_off..stride]);
418 top.push((q_scale * s_scale * dot as f32, fact.0));
419 #[cfg(feature = "counters")]
420 {
421 dots += 1;
422 }
423 }
424 #[cfg(feature = "counters")]
425 self.dots.set(self.dots.get() + dots);
426 top.sort_unstable_by(|a, b| b.0.total_cmp(&a.0).then(a.1.cmp(&b.1)));
427 for &(score, id) in top.iter().take(k) {
428 out.push((FactId(id), score));
429 }
430 Ok(())
431 }
432
433 #[cfg(test)]
438 pub(crate) fn dump(&self) -> Vec<u8> {
439 let mut out = Vec::with_capacity(self.pool_len());
440 out.extend_from_slice(self.base);
441 out.extend_from_slice(&self.tail);
442 out
443 }
444
445 pub(crate) fn pieces(&self) -> [&[u8]; 2] {
450 [self.base, &self.tail]
451 }
452
453 pub(crate) fn from_parts(dim: usize, max_bytes: usize, bytes: &[u8]) -> Result<Self, Error> {
458 Self::frame_check(dim, max_bytes, bytes.len())?;
459 let mut pool = Self::new(dim, max_bytes);
460 pool.tail = bytes.to_vec();
461 Ok(pool)
462 }
463
464 pub(crate) fn from_parts_borrowed(
470 dim: usize,
471 max_bytes: usize,
472 bytes: &'a [u8],
473 ) -> Result<Self, Error> {
474 Self::frame_check(dim, max_bytes, bytes.len())?;
475 let mut pool = Self::new(dim, max_bytes);
476 pool.base = bytes;
477 Ok(pool)
478 }
479
480 fn frame_check(dim: usize, max_bytes: usize, len: usize) -> Result<(), Error> {
484 if len > max_bytes {
485 return Err(Error::Corrupt("vector pool exceeds the configured ceiling"));
486 }
487 if dim == 0 {
488 if len != 0 {
489 return Err(Error::Corrupt("vector pool present with dim 0"));
490 }
491 return Ok(());
492 }
493 let stride = HEAD + Self::words(dim) * SIG_WORD_BYTES + dim;
494 if !len.is_multiple_of(stride) {
495 return Err(Error::Corrupt("vector pool is not a whole number of slots"));
496 }
497 Ok(())
498 }
499
500 pub(crate) fn validate(&self) -> Result<(), Error> {
506 if self.dim == 0 {
507 return Ok(());
508 }
509 let words = Self::words(self.dim);
510 let q_off = HEAD + words * SIG_WORD_BYTES;
511 for i in 0..self.len() {
512 let slot = self.slot_bytes(i);
513 let scale = f32::from_le_bytes(slot[FACT_BYTES..HEAD].try_into().unwrap());
514 if !scale.is_finite() || scale < 0.0 {
515 return Err(Error::Corrupt(
516 "vector slot scale is not finite and non-negative",
517 ));
518 }
519 for w in 0..words {
520 let stored = u64::from_le_bytes(
521 slot[HEAD + w * SIG_WORD_BYTES..HEAD + w * SIG_WORD_BYTES + SIG_WORD_BYTES]
522 .try_into()
523 .unwrap(),
524 );
525 let mut expect = 0u64;
526 for b in 0..64 {
527 let j = w * 64 + b;
528 if j >= self.dim {
529 break;
530 }
531 if slot[q_off + j] as i8 >= 0 {
532 expect |= 1 << b;
533 }
534 }
535 if stored != expect {
536 return Err(Error::Corrupt(
537 "vector slot signature disagrees with its components",
538 ));
539 }
540 }
541 }
542 Ok(())
543 }
544
545 #[cfg(feature = "counters")]
547 pub fn dots(&self) -> u64 {
548 self.dots.get()
549 }
550
551 #[cfg(feature = "counters")]
553 pub fn reset_dots(&self) {
554 self.dots.set(0);
555 }
556}
557
558#[inline]
562pub(crate) fn dot_i8(a: &[u8], b: &[u8]) -> i32 {
563 let mut acc = 0i32;
564 for (&x, &y) in a.iter().zip(b.iter()) {
565 acc += i32::from(x as i8) * i32::from(y as i8);
566 }
567 acc
568}
569
570#[cfg(test)]
571mod tests {
572 use super::*;
573 use alloc::vec;
574
575 struct Lcg(u64);
578 impl Lcg {
579 fn next(&mut self) -> f32 {
580 self.0 = self
581 .0
582 .wrapping_mul(6_364_136_223_846_793_005)
583 .wrapping_add(1_442_695_040_888_963_407);
584 ((self.0 >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
585 }
586 fn vector(&mut self, dim: usize) -> Vec<f32> {
587 (0..dim).map(|_| self.next()).collect()
588 }
589 }
590
591 fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
593 let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
594 let na: f32 = libm::sqrtf(a.iter().map(|x| x * x).sum());
595 let nb: f32 = libm::sqrtf(b.iter().map(|x| x * x).sum());
596 dot / (na * nb)
597 }
598
599 #[test]
602 fn quantized_cosine_tracks_f32() {
603 let dim = 384;
604 let mut rng = Lcg(0x1234_5678);
605 let mut worst = 0.0f32;
606 for i in 0..200u32 {
607 let a = rng.vector(dim);
608 let b = rng.vector(dim);
609 let mut pool = VecPool::new(dim, usize::MAX);
610 pool.push(FactId(2 * i), &a).unwrap();
611 pool.push(FactId(2 * i + 1), &b).unwrap();
612 let q = pool.cosine_slots(0, 1);
613 let t = cosine_f32(&a, &b);
614 worst = worst.max(libm::fabsf(q - t));
615 }
616 assert!(
617 worst < 0.05,
618 "worst quantization error {worst} exceeds 0.05"
619 );
620 }
621
622 #[test]
626 fn golden_dim4() {
627 let dim = 4;
628 let mut pool = VecPool::new(dim, usize::MAX);
629 pool.push(FactId(0), &[1.0, 1.0, 0.0, 0.0]).unwrap();
630 pool.push(FactId(1), &[2.0, 2.0, 0.0, 0.0]).unwrap(); pool.push(FactId(2), &[0.0, 0.0, 1.0, 1.0]).unwrap(); assert!((pool.cosine_slots(0, 1) - 1.0).abs() < 1e-3);
634 assert!(pool.cosine_slots(0, 2).abs() < 1e-3);
636 pool.validate().unwrap();
638 let stride = pool.stride();
641 let sig = u64::from_le_bytes(pool.dump()[HEAD..HEAD + SIG_WORD_BYTES].try_into().unwrap());
642 assert_eq!(sig & 0b1111, 0b1111);
643 assert_eq!(pool.len(), 3);
644 assert_eq!(stride, HEAD + SIG_WORD_BYTES + 4);
646 }
647
648 #[test]
650 fn search_surfaces_the_nearest() {
651 let dim = 64;
652 let mut rng = Lcg(0xdead_beef);
653 let mut pool = VecPool::new(dim, usize::MAX);
654 let target = rng.vector(dim);
655 for i in 0..200u32 {
657 pool.push(FactId(i), &rng.vector(dim)).unwrap();
658 }
659 pool.push(FactId(500), &target).unwrap();
660 let mut scratch = VecScratch::new();
661 let mut out = Vec::new();
662 pool.search(&target, 5, &mut |_| true, &mut scratch, &mut out)
663 .unwrap();
664 assert_eq!(out[0].0, FactId(500), "exact match must rank first");
665 assert!(out[0].1 > 0.99, "self-cosine ≈ 1, got {}", out[0].1);
666 }
667
668 #[test]
670 fn degenerate_vectors_are_invalid() {
671 let mut pool = VecPool::new(3, usize::MAX);
672 assert_eq!(
673 pool.push(FactId(0), &[0.0, 0.0, 0.0]).unwrap_err(),
674 Error::Invalid("vector must be nonzero")
675 );
676 assert_eq!(
677 pool.push(FactId(0), &[1.0, f32::NAN, 0.0]).unwrap_err(),
678 Error::Invalid("vector must be finite")
679 );
680 assert!(matches!(
681 pool.push(FactId(0), &[1.0, 2.0]).unwrap_err(),
682 Error::DimMismatch { got: 2, want: 3 }
683 ));
684 assert_eq!(pool.len(), 0);
686 assert!(pool.is_empty());
687 }
688
689 #[test]
692 fn accessors_and_edges() {
693 let dim = 4;
694 let mut pool = VecPool::new(dim, usize::MAX);
695 assert!(pool.is_empty());
696 assert_eq!(pool.pool_bytes(), 0);
697 let mut scratch = VecScratch::new();
698 let mut out = vec![(FactId(9), 1.0)];
699 pool.search(&[1.0; 4], 5, &mut |_| true, &mut scratch, &mut out)
701 .unwrap();
702 assert!(out.is_empty());
703
704 pool.push(FactId(0), &[1.0, 0.0, 0.0, 0.0]).unwrap();
705 assert!(!pool.is_empty());
706 assert_eq!(pool.pool_bytes(), pool.stride());
707 pool.search(&[1.0; 4], 0, &mut |_| true, &mut scratch, &mut out)
709 .unwrap();
710 assert!(out.is_empty());
711 assert_eq!(pool.cosine_slots(0, 9), 0.0);
713
714 let mut tight = VecPool::new(dim, 4);
716 assert_eq!(
717 tight.push(FactId(0), &[1.0, 0.0, 0.0, 0.0]).unwrap_err(),
718 Error::CapacityExceeded { what: "vectors" }
719 );
720 }
721
722 #[test]
724 fn from_parts_frames_slots() {
725 let dim = 8;
726 let mut pool = VecPool::new(dim, usize::MAX);
727 pool.push(FactId(0), &vec![0.5; dim]).unwrap();
728 pool.push(FactId(1), &vec![-0.5; dim]).unwrap();
729 let bytes = pool.dump();
730 let rebuilt = VecPool::from_parts(dim, usize::MAX, &bytes).unwrap();
731 assert_eq!(rebuilt.len(), 2);
732 rebuilt.validate().unwrap();
733 assert!(VecPool::from_parts(dim, usize::MAX, &bytes[..bytes.len() - 1]).is_err());
735 assert!(VecPool::from_parts(0, usize::MAX, &bytes).is_err());
737 assert!(VecPool::from_parts(dim, bytes.len() - 1, &bytes).is_err());
739 }
740
741 #[test]
744 fn validate_rejects_malformed_slots() {
745 let dim = 8;
746 let mut pool = VecPool::new(dim, usize::MAX);
747 pool.push(FactId(0), &vec![0.5; dim]).unwrap();
748 let good = pool.dump();
749
750 let mut bad = good.clone();
752 bad[FACT_BYTES..HEAD].copy_from_slice(&f32::NAN.to_le_bytes());
753 assert!(
754 VecPool::from_parts(dim, usize::MAX, &bad)
755 .unwrap()
756 .validate()
757 .is_err()
758 );
759
760 let mut bad = good.clone();
763 let q_off = HEAD + VecPool::words(dim) * SIG_WORD_BYTES;
764 bad[q_off] = (-1i8) as u8; assert!(
766 VecPool::from_parts(dim, usize::MAX, &bad)
767 .unwrap()
768 .validate()
769 .is_err()
770 );
771 }
772
773 #[test]
778 fn overlay_appends_to_tail_and_reads_span_the_boundary() {
779 let dim = 16;
780 let mut rng = Lcg(0x0ace_1a75);
781 let (va, vb, vc) = (rng.vector(dim), rng.vector(dim), rng.vector(dim));
782
783 let mut owned = VecPool::new(dim, usize::MAX);
785 owned.push(FactId(10), &va).unwrap();
786 owned.push(FactId(11), &vb).unwrap();
787 owned.push(FactId(12), &vc).unwrap();
788
789 let mut seed = VecPool::new(dim, usize::MAX);
792 seed.push(FactId(10), &va).unwrap();
793 seed.push(FactId(11), &vb).unwrap();
794 let base = seed.dump();
795 let base_snapshot = base.clone();
796
797 let mut pool = VecPool::from_parts_borrowed(dim, usize::MAX, &base).unwrap();
800 assert_eq!(pool.len(), 2);
801 let idx = pool.push(FactId(12), &vc).unwrap();
802 assert_eq!(idx, 2);
803 assert_eq!(pool.len(), 3);
804
805 assert_eq!(pool.slot_fact(0), 10); assert_eq!(pool.slot_fact(2), 12); assert!((pool.cosine_slots(0, 2) - owned.cosine_slots(0, 2)).abs() < 1e-6);
811 pool.validate().unwrap();
812
813 let mut scratch = VecScratch::new();
815 let mut out = Vec::new();
816 pool.search(&vc, 1, &mut |_| true, &mut scratch, &mut out)
817 .unwrap();
818 assert_eq!(out[0].0, FactId(12));
819
820 assert_eq!(pool.dump(), owned.dump());
822 assert_eq!(base, base_snapshot);
823 }
824}