1#[cfg(feature = "parallel")]
33use rayon::prelude::*;
34
35#[inline(always)]
42fn u8_dot_u32(a: &[u8], b: &[u8]) -> u32 {
43 #[cfg(target_arch = "aarch64")]
44 {
45 use std::arch::aarch64::*;
46 let n = a.len();
47 let chunks = n / 16;
48 let rem = n % 16;
49
50 let mut acc0: uint32x4_t;
51 let mut acc1: uint32x4_t;
52 let mut acc2: uint32x4_t;
53 let mut acc3: uint32x4_t;
54
55 unsafe {
56 acc0 = vdupq_n_u32(0);
57 acc1 = vdupq_n_u32(0);
58 acc2 = vdupq_n_u32(0);
59 acc3 = vdupq_n_u32(0);
60
61 for i in 0..chunks {
62 let ap = a.as_ptr().add(i * 16);
63 let bp = b.as_ptr().add(i * 16);
64
65 let va = vld1q_u8(ap);
66 let vb = vld1q_u8(bp);
67
68 let lo_u16 = vmull_u8(vget_low_u8(va), vget_low_u8(vb));
69 let hi_u16 = vmull_high_u8(va, vb);
70
71 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
72 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
73 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
74 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
75 }
76
77 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
78 let mut total = vaddvq_u32(sum4);
79
80 for i in (n - rem)..n {
81 total += a[i] as u32 * b[i] as u32;
82 }
83 total
84 }
85 }
86
87 #[cfg(not(target_arch = "aarch64"))]
88 {
89 a.chunks(8)
90 .zip(b.chunks(8))
91 .map(|(ac, bc)| {
92 ac.iter()
93 .zip(bc.iter())
94 .map(|(&x, &y)| (x as u32) * (y as u32))
95 .sum::<u32>()
96 })
97 .sum()
98 }
99}
100
101#[inline(always)]
107pub fn u8_l2sq_u32(a: &[u8], b: &[u8]) -> u32 {
108 assert_eq!(
109 a.len(),
110 b.len(),
111 "u8_l2sq_u32 inputs must have equal length"
112 );
113
114 #[cfg(target_arch = "aarch64")]
115 {
116 use std::arch::aarch64::*;
117 let n = a.len();
118 let chunks = n / 16;
119 let rem = n % 16;
120
121 let mut acc0: uint32x4_t;
122 let mut acc1: uint32x4_t;
123 let mut acc2: uint32x4_t;
124 let mut acc3: uint32x4_t;
125
126 unsafe {
127 acc0 = vdupq_n_u32(0);
128 acc1 = vdupq_n_u32(0);
129 acc2 = vdupq_n_u32(0);
130 acc3 = vdupq_n_u32(0);
131
132 for i in 0..chunks {
133 let ap = a.as_ptr().add(i * 16);
134 let bp = b.as_ptr().add(i * 16);
135
136 let va = vld1q_u8(ap);
137 let vb = vld1q_u8(bp);
138
139 let diff = vabdq_u8(va, vb);
140
141 let lo_u16 = vmull_u8(vget_low_u8(diff), vget_low_u8(diff));
142 let hi_u16 = vmull_high_u8(diff, diff);
143
144 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
145 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
146 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
147 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
148 }
149
150 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
151 let mut total = vaddvq_u32(sum4);
152
153 for i in (n - rem)..n {
154 let d = (a[i] as i32) - (b[i] as i32);
155 total += (d * d) as u32;
156 }
157 total
158 }
159 }
160
161 #[cfg(not(target_arch = "aarch64"))]
162 {
163 a.chunks(8)
164 .zip(b.chunks(8))
165 .map(|(ac, bc)| {
166 ac.iter()
167 .zip(bc.iter())
168 .map(|(&x, &y)| {
169 let d = (x as i32) - (y as i32);
170 (d * d) as u32
171 })
172 .sum::<u32>()
173 })
174 .sum()
175 }
176}
177
178#[derive(Debug, Clone, PartialEq)]
186pub enum QuantError {
187 EmptyCorpus,
189 ZeroDims,
191 FlatLengthNotDivisible { len: usize, dims: usize },
193 RaggedRow {
195 row: usize,
196 expected: usize,
197 got: usize,
198 },
199 EncodeLengthMismatch { expected: usize, got: usize },
202}
203
204impl std::fmt::Display for QuantError {
205 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
206 match self {
207 Self::EmptyCorpus => write!(f, "cannot train on empty corpus"),
208 Self::ZeroDims => write!(f, "dims must be > 0"),
209 Self::FlatLengthNotDivisible { len, dims } => write!(
210 f,
211 "flat vector length {len} is not a multiple of dims {dims}"
212 ),
213 Self::RaggedRow { row, expected, got } => write!(
214 f,
215 "row {row} has length {got}, expected {expected} (dims fixed by row 0)"
216 ),
217 Self::EncodeLengthMismatch { expected, got } => write!(
218 f,
219 "vector length {got} does not match codec dims {expected}"
220 ),
221 }
222 }
223}
224
225impl std::error::Error for QuantError {}
226
227fn flat_min_max(vectors: &[f32], dims: usize) -> Result<(Vec<f32>, Vec<f32>), QuantError> {
231 if dims == 0 {
232 return Err(QuantError::ZeroDims);
233 }
234 if vectors.is_empty() {
235 return Err(QuantError::EmptyCorpus);
236 }
237 if !vectors.len().is_multiple_of(dims) {
238 return Err(QuantError::FlatLengthNotDivisible {
239 len: vectors.len(),
240 dims,
241 });
242 }
243
244 let n = vectors.len() / dims;
245 let mut min = vec![f32::INFINITY; dims];
246 let mut max = vec![f32::NEG_INFINITY; dims];
247
248 for row in 0..n {
249 let v = &vectors[row * dims..(row + 1) * dims];
250 for (d, &x) in v.iter().enumerate() {
251 if x.is_finite() {
252 if x < min[d] {
253 min[d] = x;
254 }
255 if x > max[d] {
256 max[d] = x;
257 }
258 }
259 }
260 }
261 finalize_min_max(&mut min, &mut max);
262 Ok((min, max))
263}
264
265fn row_min_max(vectors: &[Vec<f32>]) -> Result<(usize, Vec<f32>, Vec<f32>), QuantError> {
268 if vectors.is_empty() {
269 return Err(QuantError::EmptyCorpus);
270 }
271 let dims = vectors[0].len();
272 if dims == 0 {
273 return Err(QuantError::ZeroDims);
274 }
275
276 let mut min = vec![f32::INFINITY; dims];
277 let mut max = vec![f32::NEG_INFINITY; dims];
278
279 for (row, v) in vectors.iter().enumerate() {
280 if v.len() != dims {
281 return Err(QuantError::RaggedRow {
282 row,
283 expected: dims,
284 got: v.len(),
285 });
286 }
287 for (d, &x) in v.iter().enumerate() {
288 if x.is_finite() {
289 if x < min[d] {
290 min[d] = x;
291 }
292 if x > max[d] {
293 max[d] = x;
294 }
295 }
296 }
297 }
298 finalize_min_max(&mut min, &mut max);
299 Ok((dims, min, max))
300}
301
302fn finalize_min_max(min: &mut [f32], max: &mut [f32]) {
305 for d in 0..min.len() {
306 if !min[d].is_finite() {
307 min[d] = 0.0;
308 }
309 if !max[d].is_finite() || max[d] <= min[d] {
310 max[d] = min[d] + 1.0;
311 }
312 }
313}
314
315#[derive(Debug, Clone)]
325pub struct Sq8Codec {
326 pub min: Vec<f32>,
328 pub scale: Vec<f32>,
330 pub scale_sq: Vec<f32>,
332 pub mean_scale_sq: f32,
334 pub scale_sq_residual: Vec<f32>,
336 pub offset_sq_sum: f32,
338}
339
340#[derive(Debug, Clone)]
342pub struct EncodedVector {
343 pub codes: Vec<u8>,
345 pub norm: f32,
347 pub soc_sum: f32,
349 pub residual_dot_bias: f32,
351}
352
353impl Sq8Codec {
354 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
355 let dims = min.len();
356 let scale: Vec<f32> = (0..dims).map(|d| (max[d] - min[d]) / 255.0).collect();
357 let scale_sq: Vec<f32> = scale.iter().map(|s| s * s).collect();
358 let mean_scale_sq = scale_sq.iter().sum::<f32>() / dims as f32;
359 let scale_sq_residual: Vec<f32> = scale_sq.iter().map(|&ss| ss - mean_scale_sq).collect();
360 let offset_sq_sum: f32 = min.iter().map(|o| o * o).sum();
361
362 Self {
363 min,
364 scale,
365 scale_sq,
366 mean_scale_sq,
367 scale_sq_residual,
368 offset_sq_sum,
369 }
370 }
371
372 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
377 Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
378 }
379
380 pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
383 let (min, max) = flat_min_max(vectors, dims)?;
384 Ok(Self::build_from_min_max(min, max))
385 }
386
387 pub fn train(vectors: &[Vec<f32>]) -> Self {
392 Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
393 }
394
395 pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
400 let (_dims, min, max) = row_min_max(vectors)?;
401 Ok(Self::build_from_min_max(min, max))
402 }
403
404 pub fn encode(&self, v: &[f32]) -> EncodedVector {
410 self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
411 }
412
413 pub fn try_encode(&self, v: &[f32]) -> Result<EncodedVector, QuantError> {
417 let dims = self.min.len();
418 if v.len() != dims {
419 return Err(QuantError::EncodeLengthMismatch {
420 expected: dims,
421 got: v.len(),
422 });
423 }
424 Ok(self.encode_unchecked(v))
425 }
426
427 fn encode_unchecked(&self, v: &[f32]) -> EncodedVector {
429 let dims = self.min.len();
430 let mut codes = Vec::with_capacity(dims);
431 let mut soc_sum = 0.0f32;
432 let mut residual_dot_bias = 0.0f32;
433 let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
434
435 for (d, &x) in v.iter().enumerate() {
436 let s = self.scale[d];
437 let inv_s = if s > 1e-12 { 1.0 / s } else { 0.0 };
438 let raw = (x - self.min[d]) * inv_s;
439 let code = raw.round().clamp(0.0, 255.0) as u8;
440 codes.push(code);
441 soc_sum += s * self.min[d] * code as f32;
442 residual_dot_bias += self.scale_sq_residual[d] * code as f32;
443 }
444
445 EncodedVector {
446 codes,
447 norm,
448 soc_sum,
449 residual_dot_bias,
450 }
451 }
452
453 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<EncodedVector> {
458 self.try_encode_flat_par(vectors, dims)
459 .unwrap_or_else(|e| panic!("{e}"))
460 }
461
462 pub fn try_encode_flat_par(
466 &self,
467 vectors: &[f32],
468 dims: usize,
469 ) -> Result<Vec<EncodedVector>, QuantError> {
470 if dims == 0 {
471 return Err(QuantError::ZeroDims);
472 }
473 if !vectors.len().is_multiple_of(dims) {
474 return Err(QuantError::FlatLengthNotDivisible {
475 len: vectors.len(),
476 dims,
477 });
478 }
479 if dims != self.min.len() {
480 return Err(QuantError::EncodeLengthMismatch {
481 expected: self.min.len(),
482 got: dims,
483 });
484 }
485 let n = vectors.len() / dims;
486 #[cfg(feature = "parallel")]
487 let encoded = (0..n)
488 .into_par_iter()
489 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
490 .collect();
491 #[cfg(not(feature = "parallel"))]
492 let encoded = (0..n)
493 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
494 .collect();
495 Ok(encoded)
496 }
497
498 pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
504 self.try_encode_par(vectors)
505 .unwrap_or_else(|e| panic!("{e}"))
506 }
507
508 pub fn try_encode_par(&self, vectors: &[Vec<f32>]) -> Result<Vec<EncodedVector>, QuantError> {
511 let dims = self.min.len();
512 for v in vectors {
513 if v.len() != dims {
514 return Err(QuantError::EncodeLengthMismatch {
515 expected: dims,
516 got: v.len(),
517 });
518 }
519 }
520 #[cfg(feature = "parallel")]
521 let encoded = vectors
522 .par_iter()
523 .map(|v| self.encode_unchecked(v))
524 .collect();
525 #[cfg(not(feature = "parallel"))]
526 let encoded = vectors.iter().map(|v| self.encode_unchecked(v)).collect();
527 Ok(encoded)
528 }
529
530 #[inline]
539 pub fn approx_dot(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
540 let raw = u8_dot_u32(&a.codes, &b.codes) as f32;
541 let residual_hot: f32 = self
542 .scale_sq_residual
543 .iter()
544 .zip(a.codes.iter())
545 .zip(b.codes.iter())
546 .map(|((r, &ac), &bc)| r * (ac as f32) * (bc as f32))
547 .sum();
548 self.mean_scale_sq * raw + residual_hot + a.soc_sum + b.soc_sum + self.offset_sq_sum
549 }
550
551 #[inline]
555 pub fn approx_cosine_dist(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
556 let denom = a.norm * b.norm;
557 if !denom.is_finite() || denom <= 0.0 {
558 return 1.0;
559 }
560 let dot = self.approx_dot(a, b);
561 let cosine = (dot / denom).clamp(-1.0, 1.0);
562 1.0 - cosine
563 }
564
565 #[inline]
577 pub fn approx_l2_sq(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
578 let raw = u8_l2sq_u32(&a.codes, &b.codes) as f32;
579 let residual_hot: f32 = self
580 .scale_sq_residual
581 .iter()
582 .zip(a.codes.iter())
583 .zip(b.codes.iter())
584 .map(|((r, &ac), &bc)| {
585 let d = (ac as i32) - (bc as i32);
586 r * (d as f32) * (d as f32)
587 })
588 .sum();
589 self.mean_scale_sq * raw + residual_hot
590 }
591
592 pub fn dims(&self) -> usize {
594 self.min.len()
595 }
596}
597
598#[derive(Debug, Clone)]
612pub struct GsSq8Codec {
613 pub min: Vec<f32>,
615 pub gs: f32,
617 pub gs_sq: f32,
619 pub anisotropy_ratio: f32,
622}
623
624#[derive(Debug, Clone)]
626pub struct GsEncodedVector {
627 pub codes: Vec<u8>,
629}
630
631impl GsSq8Codec {
632 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
633 let dims = min.len();
634 let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
635 let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
636 let gs = if max_range > 1e-12 {
637 max_range / 255.0
638 } else {
639 1.0 / 255.0
640 };
641
642 let min_range_nonzero = ranges
643 .iter()
644 .cloned()
645 .filter(|&r| r > 1e-12)
646 .fold(f32::INFINITY, f32::min);
647 let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
648 max_range / min_range_nonzero
649 } else {
650 1.0
651 };
652
653 Self {
654 min,
655 gs,
656 gs_sq: gs * gs,
657 anisotropy_ratio,
658 }
659 }
660
661 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
666 Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
667 }
668
669 pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
672 let (min, max) = flat_min_max(vectors, dims)?;
673 Ok(Self::build_from_min_max(min, max))
674 }
675
676 pub fn train(vectors: &[Vec<f32>]) -> Self {
681 Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
682 }
683
684 pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
689 let (_dims, min, max) = row_min_max(vectors)?;
690 Ok(Self::build_from_min_max(min, max))
691 }
692
693 #[inline]
699 pub fn encode(&self, v: &[f32]) -> GsEncodedVector {
700 self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
701 }
702
703 pub fn try_encode(&self, v: &[f32]) -> Result<GsEncodedVector, QuantError> {
708 let dims = self.min.len();
709 if v.len() != dims {
710 return Err(QuantError::EncodeLengthMismatch {
711 expected: dims,
712 got: v.len(),
713 });
714 }
715 Ok(self.encode_unchecked(v))
716 }
717
718 #[inline]
720 fn encode_unchecked(&self, v: &[f32]) -> GsEncodedVector {
721 let inv_gs = if self.gs > 1e-12 { 1.0 / self.gs } else { 0.0 };
722 let codes = v
723 .iter()
724 .enumerate()
725 .map(|(d, &x)| ((x - self.min[d]) * inv_gs).round().clamp(0.0, 255.0) as u8)
726 .collect();
727 GsEncodedVector { codes }
728 }
729
730 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<GsEncodedVector> {
735 self.try_encode_flat_par(vectors, dims)
736 .unwrap_or_else(|e| panic!("{e}"))
737 }
738
739 pub fn try_encode_flat_par(
743 &self,
744 vectors: &[f32],
745 dims: usize,
746 ) -> Result<Vec<GsEncodedVector>, QuantError> {
747 if dims == 0 {
748 return Err(QuantError::ZeroDims);
749 }
750 if !vectors.len().is_multiple_of(dims) {
751 return Err(QuantError::FlatLengthNotDivisible {
752 len: vectors.len(),
753 dims,
754 });
755 }
756 if dims != self.min.len() {
757 return Err(QuantError::EncodeLengthMismatch {
758 expected: self.min.len(),
759 got: dims,
760 });
761 }
762 let n = vectors.len() / dims;
763 #[cfg(feature = "parallel")]
764 let encoded = (0..n)
765 .into_par_iter()
766 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
767 .collect();
768 #[cfg(not(feature = "parallel"))]
769 let encoded = (0..n)
770 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
771 .collect();
772 Ok(encoded)
773 }
774
775 #[inline]
784 pub fn l2_sq(&self, a: &GsEncodedVector, b: &GsEncodedVector) -> f32 {
785 self.gs_sq * u8_l2sq_u32(&a.codes, &b.codes) as f32
786 }
787
788 #[inline]
791 pub fn l2_sq_codes(&self, a: &[u8], b: &[u8]) -> f32 {
792 self.gs_sq * u8_l2sq_u32(a, b) as f32
793 }
794
795 pub fn dims(&self) -> usize {
797 self.min.len()
798 }
799
800 #[inline]
806 pub fn is_in_distribution(&self, v: &[f32]) -> bool {
807 let max_code = 255.0 * self.gs;
808 v.iter()
809 .zip(self.min.iter())
810 .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
811 }
812}
813
814#[cfg(test)]
815mod tests {
816 use super::*;
817
818 fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
819 let mut h = seed;
820 (0..n)
821 .map(|_| {
822 (0..dims)
823 .map(|_| {
824 h = h
825 .wrapping_mul(0x6c62_272e_07bb_0142)
826 .wrapping_add(0x62b8_2175_62d9_6b1a);
827 let bits = (h >> 33) as u32;
828 (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
829 })
830 .collect()
831 })
832 .collect()
833 }
834
835 fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
836 a.iter().zip(b).map(|(x, y)| x * y).sum()
837 }
838
839 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
840 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
841 }
842
843 #[test]
846 fn sq8_try_train_ragged_rows_returns_error_not_panic() {
847 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
848 let err = Sq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
849 assert_eq!(
850 err,
851 QuantError::RaggedRow {
852 row: 1,
853 expected: 1,
854 got: 2
855 }
856 );
857 }
858
859 #[test]
860 fn gs_try_train_ragged_rows_returns_error_not_panic() {
861 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
862 let err = GsSq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
863 assert_eq!(
864 err,
865 QuantError::RaggedRow {
866 row: 1,
867 expected: 1,
868 got: 2
869 }
870 );
871 }
872
873 #[test]
874 fn sq8_try_train_empty_corpus_returns_error() {
875 let vecs: Vec<Vec<f32>> = vec![];
876 assert_eq!(
877 Sq8Codec::try_train(&vecs).unwrap_err(),
878 QuantError::EmptyCorpus
879 );
880 }
881
882 #[test]
883 fn sq8_try_train_flat_zero_dims_returns_error() {
884 assert_eq!(
885 Sq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
886 QuantError::ZeroDims
887 );
888 }
889
890 #[test]
891 fn gs_try_train_flat_zero_dims_returns_error() {
892 assert_eq!(
893 GsSq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
894 QuantError::ZeroDims
895 );
896 }
897
898 #[test]
899 fn sq8_try_train_flat_remainder_returns_error() {
900 let err = Sq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
902 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
903 }
904
905 #[test]
906 fn gs_try_train_flat_remainder_returns_error() {
907 let err = GsSq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
908 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
909 }
910
911 #[test]
912 fn sq8_try_encode_shorter_input_returns_error_not_malformed_vector() {
913 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
914 let err = codec.try_encode(&[0.0, 1.0]).unwrap_err();
915 assert_eq!(
916 err,
917 QuantError::EncodeLengthMismatch {
918 expected: 4,
919 got: 2
920 }
921 );
922 }
923
924 #[test]
925 fn sq8_try_encode_longer_input_returns_error() {
926 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
927 let err = codec.try_encode(&[0.0, 1.0, 2.0, 3.0, 4.0]).unwrap_err();
928 assert_eq!(
929 err,
930 QuantError::EncodeLengthMismatch {
931 expected: 4,
932 got: 5
933 }
934 );
935 }
936
937 #[test]
938 fn gs_try_encode_empty_input_returns_error_not_malformed_vector() {
939 let codec = GsSq8Codec::train_flat(&[1.0, 2.0, 3.0, 4.0], 1);
945 let err = codec.try_encode(&[]).unwrap_err();
946 assert_eq!(
947 err,
948 QuantError::EncodeLengthMismatch {
949 expected: 1,
950 got: 0
951 }
952 );
953 }
954
955 #[test]
956 fn sq8_try_encode_flat_par_zero_dims_returns_error() {
957 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
958 assert_eq!(
959 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
960 QuantError::ZeroDims
961 );
962 }
963
964 #[test]
965 fn gs_try_encode_flat_par_zero_dims_returns_error() {
966 let codec = GsSq8Codec::train(&rand_vecs(10, 4, 1));
967 assert_eq!(
968 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
969 QuantError::ZeroDims
970 );
971 }
972
973 #[test]
974 fn sq8_try_encode_flat_par_remainder_returns_error() {
975 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
976 let flat: Vec<f32> = (0..9).map(|i| i as f32).collect(); assert_eq!(
978 codec.try_encode_flat_par(&flat, 4).unwrap_err(),
979 QuantError::FlatLengthNotDivisible { len: 9, dims: 4 }
980 );
981 }
982
983 #[test]
984 fn sq8_try_encode_flat_par_dims_mismatch_returns_error() {
985 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
986 let flat: Vec<f32> = (0..6).map(|i| i as f32).collect();
987 let err = codec.try_encode_flat_par(&flat, 3).unwrap_err();
988 assert_eq!(
989 err,
990 QuantError::EncodeLengthMismatch {
991 expected: 4,
992 got: 3
993 }
994 );
995 }
996
997 #[test]
998 #[should_panic(expected = "row 1 has length 2, expected 1")]
999 fn sq8_train_still_panics_with_typed_message_on_ragged_rows() {
1000 let _ = Sq8Codec::train(&[vec![0.0], vec![1.0, 2.0]]);
1004 }
1005
1006 #[test]
1007 fn sq8_try_train_and_encode_roundtrip_matches_panicking_api() {
1008 let vecs = rand_vecs(20, 8, 7);
1009 let a = Sq8Codec::train(&vecs);
1010 let b = Sq8Codec::try_train(&vecs).expect("valid corpus must train");
1011 assert_eq!(a.min, b.min);
1012 assert_eq!(a.scale, b.scale);
1013 let ea = a.encode(&vecs[0]);
1014 let eb = b.try_encode(&vecs[0]).expect("valid vector must encode");
1015 assert_eq!(ea.codes, eb.codes);
1016 }
1017
1018 #[test]
1021 fn encode_decode_roundtrip_is_bounded() {
1022 let vecs = rand_vecs(100, 32, 42);
1023 let codec = Sq8Codec::train(&vecs);
1024 for v in &vecs {
1025 let ev = codec.encode(v);
1026 assert_eq!(ev.codes.len(), v.len());
1027 for (d, &code) in ev.codes.iter().enumerate() {
1028 let decoded = code as f32 * codec.scale[d] + codec.min[d];
1029 let err = (decoded - v[d]).abs();
1030 assert!(
1031 err <= codec.scale[d] + 1e-5,
1032 "dim {d}: err={err} scale={}",
1033 codec.scale[d]
1034 );
1035 }
1036 }
1037 }
1038
1039 #[test]
1040 fn approx_dot_relative_error_bounded() {
1041 let vecs = rand_vecs(200, 64, 77);
1042 let codec = Sq8Codec::train(&vecs);
1043 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1044
1045 let mut max_rel_err = 0.0f32;
1046 for i in 0..vecs.len() {
1047 for j in (i + 1)..vecs.len().min(i + 10) {
1048 let true_dot = dot_f32(&vecs[i], &vecs[j]);
1049 let approx = codec.approx_dot(&encoded[i], &encoded[j]);
1050 let denom = true_dot.abs().max(1e-3);
1051 let rel = (approx - true_dot).abs() / denom;
1052 if rel > max_rel_err {
1053 max_rel_err = rel;
1054 }
1055 }
1056 }
1057 assert!(
1058 max_rel_err < 0.15,
1059 "max relative dot error {max_rel_err:.4} >= 0.15"
1060 );
1061 }
1062
1063 #[test]
1064 fn approx_l2_sq_relative_error_bounded() {
1065 let vecs = rand_vecs(200, 64, 88);
1066 let codec = Sq8Codec::train(&vecs);
1067 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1068
1069 let mut max_rel_err = 0.0f32;
1070 for i in 0..vecs.len() {
1071 for j in (i + 1)..vecs.len().min(i + 10) {
1072 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1073 let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
1074 let denom = true_l2.max(1e-6);
1075 let rel = (approx - true_l2).abs() / denom;
1076 if rel > max_rel_err {
1077 max_rel_err = rel;
1078 }
1079 }
1080 }
1081 assert!(
1082 max_rel_err < 0.15,
1083 "max relative L2² error {max_rel_err:.4} >= 0.15"
1084 );
1085 }
1086
1087 #[test]
1088 fn order_preservation_triplets_cosine() {
1089 let vecs = rand_vecs(300, 64, 99);
1090 let codec = Sq8Codec::train(&vecs);
1091 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1092
1093 let n = vecs.len();
1094 let mut agree = 0usize;
1095 let mut total = 0usize;
1096
1097 for anchor in 0..50 {
1098 let a = &vecs[anchor];
1099 let ea = &encoded[anchor];
1100 for b_idx in 0..n {
1101 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1102 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1103 let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1104 let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1105
1106 let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
1107 let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
1108 let dist_ab_true = 1.0 - cos_ab;
1109 let dist_ac_true = 1.0 - cos_ac;
1110
1111 let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
1112 let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
1113
1114 if (dist_ab_true - dist_ac_true).abs() < 0.01 {
1115 continue;
1116 }
1117
1118 let true_closer_b = dist_ab_true < dist_ac_true;
1119 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1120 if true_closer_b == approx_closer_b {
1121 agree += 1;
1122 }
1123 total += 1;
1124 }
1125 }
1126 }
1127
1128 let rate = agree as f64 / total.max(1) as f64;
1129 assert!(
1130 rate >= 0.95,
1131 "order preservation {rate:.3} < 0.95 ({agree}/{total})"
1132 );
1133 }
1134
1135 #[test]
1136 fn order_preservation_triplets_l2() {
1137 let vecs = rand_vecs(300, 64, 101);
1138 let codec = Sq8Codec::train(&vecs);
1139 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1140
1141 let n = vecs.len();
1142 let mut agree = 0usize;
1143 let mut total = 0usize;
1144
1145 for anchor in 0..50 {
1146 let a = &vecs[anchor];
1147 let ea = &encoded[anchor];
1148 for b_idx in 0..n {
1149 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1150 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1151 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1152
1153 let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
1154 let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
1155
1156 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1157 continue;
1158 }
1159
1160 let true_closer_b = dist_ab_true < dist_ac_true;
1161 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1162 if true_closer_b == approx_closer_b {
1163 agree += 1;
1164 }
1165 total += 1;
1166 }
1167 }
1168 }
1169
1170 let rate = agree as f64 / total.max(1) as f64;
1171 assert!(
1172 rate >= 0.95,
1173 "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1174 );
1175 }
1176
1177 #[test]
1178 fn train_flat_matches_train_rows() {
1179 let vecs = rand_vecs(50, 16, 123);
1180 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1181
1182 let codec_rows = Sq8Codec::train(&vecs);
1183 let codec_flat = Sq8Codec::train_flat(&flat, 16);
1184
1185 for d in 0..16 {
1186 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1187 assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
1188 }
1189 }
1190
1191 #[test]
1192 fn encode_par_matches_sequential() {
1193 let vecs = rand_vecs(50, 32, 555);
1194 let codec = Sq8Codec::train(&vecs);
1195
1196 let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1197 let par = codec.encode_par(&vecs);
1198
1199 assert_eq!(seq.len(), par.len());
1200 for (s, p) in seq.iter().zip(par.iter()) {
1201 assert_eq!(s.codes, p.codes);
1202 assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
1203 }
1204 }
1205
1206 #[test]
1207 fn sq8_try_encode_par_short_row_returns_error_not_panic() {
1208 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1209 let mut rows = rand_vecs(5, 4, 2);
1210 rows[3] = vec![0.0, 1.0];
1211 let err = codec.try_encode_par(&rows).unwrap_err();
1212 assert_eq!(
1213 err,
1214 QuantError::EncodeLengthMismatch {
1215 expected: 4,
1216 got: 2
1217 }
1218 );
1219 }
1220
1221 #[test]
1222 fn sq8_try_encode_par_long_row_returns_error_not_panic() {
1223 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1224 let mut rows = rand_vecs(5, 4, 2);
1225 rows[3] = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
1226 let err = codec.try_encode_par(&rows).unwrap_err();
1227 assert_eq!(
1228 err,
1229 QuantError::EncodeLengthMismatch {
1230 expected: 4,
1231 got: 6
1232 }
1233 );
1234 }
1235
1236 #[test]
1237 #[should_panic(expected = "vector length 2 does not match codec dims 4")]
1238 fn sq8_encode_par_still_panics_with_typed_message_on_short_row() {
1239 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1240 let mut rows = rand_vecs(5, 4, 2);
1241 rows[3] = vec![0.0, 1.0];
1242 let _ = codec.encode_par(&rows);
1243 }
1244
1245 #[test]
1246 fn u8_dot_u32_matches_scalar() {
1247 let a: Vec<u8> = (0u8..=255).take(384).collect();
1248 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1249 let scalar: u32 = a
1250 .iter()
1251 .zip(b.iter())
1252 .map(|(&x, &y)| x as u32 * y as u32)
1253 .sum();
1254 assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
1255 }
1256
1257 #[test]
1258 fn u8_helpers_tail_path_max_diff() {
1259 for len in [1usize, 7, 15, 17, 100, 383] {
1260 let a = vec![255u8; len];
1261 let b = vec![0u8; len];
1262 assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
1263 assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
1264 }
1265 }
1266
1267 #[test]
1268 fn u8_l2sq_u32_matches_scalar() {
1269 let a: Vec<u8> = (0u8..=255).take(384).collect();
1270 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1271 let scalar: u32 = a
1272 .iter()
1273 .zip(b.iter())
1274 .map(|(&x, &y)| {
1275 let d = (x as i32) - (y as i32);
1276 (d * d) as u32
1277 })
1278 .sum();
1279 assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
1280 }
1281
1282 #[test]
1283 #[should_panic(expected = "u8_l2sq_u32 inputs must have equal length")]
1284 fn u8_l2sq_u32_rejects_shorter_second_slice() {
1285 let a = [1u8; 16];
1286 let b = [2u8; 1];
1287
1288 let _ = u8_l2sq_u32(&a, &b);
1289 }
1290
1291 #[test]
1299 fn gs_l2_sq_anisotropic_ordering_preserved() {
1300 let corpus = vec![
1301 vec![0.0f32, 0.0f32], vec![1.0f32, 1.0f32], vec![1.0f32, 4001.0f32], ];
1305 let codec = GsSq8Codec::train(&corpus);
1306
1307 let enc_origin = codec.encode(&corpus[0]);
1308 let enc_near = codec.encode(&corpus[1]);
1309 let enc_far = codec.encode(&corpus[2]);
1310
1311 let d_near = codec.l2_sq(&enc_origin, &enc_near);
1312 let d_far = codec.l2_sq(&enc_origin, &enc_far);
1313
1314 assert!(
1315 d_near < d_far,
1316 "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
1317 (anisotropy_ratio={:.1})",
1318 codec.anisotropy_ratio
1319 );
1320 }
1321
1322 #[test]
1323 fn gs_l2_sq_isotropic_small_error() {
1324 let vecs = rand_vecs(200, 64, 202);
1325 let codec = GsSq8Codec::train(&vecs);
1326 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1327
1328 let mut max_rel = 0.0f32;
1329 for i in 0..vecs.len() {
1330 for j in (i + 1)..vecs.len().min(i + 10) {
1331 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1332 let approx = codec.l2_sq(&encoded[i], &encoded[j]);
1333 let denom = true_l2.max(1e-6);
1334 let rel = (approx - true_l2).abs() / denom;
1335 if rel > max_rel {
1336 max_rel = rel;
1337 }
1338 }
1339 }
1340 assert!(
1341 max_rel < 0.15,
1342 "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
1343 );
1344 }
1345
1346 #[test]
1347 fn gs_train_flat_matches_train_rows() {
1348 let vecs = rand_vecs(50, 16, 321);
1349 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1350
1351 let codec_rows = GsSq8Codec::train(&vecs);
1352 let codec_flat = GsSq8Codec::train_flat(&flat, 16);
1353
1354 assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
1355 for d in 0..16 {
1356 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1357 }
1358 }
1359
1360 #[test]
1361 fn gs_l2_sq_order_preservation_triplets() {
1362 let vecs = rand_vecs(300, 64, 303);
1363 let codec = GsSq8Codec::train(&vecs);
1364 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1365
1366 let n = vecs.len();
1367 let mut agree = 0usize;
1368 let mut total = 0usize;
1369
1370 for anchor in 0..50 {
1371 let a = &vecs[anchor];
1372 let ea = &encoded[anchor];
1373 for b_idx in 0..n {
1374 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1375 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1376 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1377
1378 let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
1379 let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
1380
1381 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1382 continue;
1383 }
1384
1385 let true_closer_b = dist_ab_true < dist_ac_true;
1386 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1387 if true_closer_b == approx_closer_b {
1388 agree += 1;
1389 }
1390 total += 1;
1391 }
1392 }
1393 }
1394
1395 let rate = agree as f64 / total.max(1) as f64;
1396 assert!(
1397 rate >= 0.95,
1398 "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1399 );
1400 }
1401}