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 pub fn dims(&self) -> usize {
790 self.min.len()
791 }
792
793 #[inline]
799 pub fn is_in_distribution(&self, v: &[f32]) -> bool {
800 let max_code = 255.0 * self.gs;
801 v.iter()
802 .zip(self.min.iter())
803 .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
804 }
805}
806
807#[cfg(test)]
808mod tests {
809 use super::*;
810
811 fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
812 let mut h = seed;
813 (0..n)
814 .map(|_| {
815 (0..dims)
816 .map(|_| {
817 h = h
818 .wrapping_mul(0x6c62_272e_07bb_0142)
819 .wrapping_add(0x62b8_2175_62d9_6b1a);
820 let bits = (h >> 33) as u32;
821 (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
822 })
823 .collect()
824 })
825 .collect()
826 }
827
828 fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
829 a.iter().zip(b).map(|(x, y)| x * y).sum()
830 }
831
832 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
833 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
834 }
835
836 #[test]
839 fn sq8_try_train_ragged_rows_returns_error_not_panic() {
840 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
841 let err = Sq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
842 assert_eq!(
843 err,
844 QuantError::RaggedRow {
845 row: 1,
846 expected: 1,
847 got: 2
848 }
849 );
850 }
851
852 #[test]
853 fn gs_try_train_ragged_rows_returns_error_not_panic() {
854 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
855 let err = GsSq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
856 assert_eq!(
857 err,
858 QuantError::RaggedRow {
859 row: 1,
860 expected: 1,
861 got: 2
862 }
863 );
864 }
865
866 #[test]
867 fn sq8_try_train_empty_corpus_returns_error() {
868 let vecs: Vec<Vec<f32>> = vec![];
869 assert_eq!(
870 Sq8Codec::try_train(&vecs).unwrap_err(),
871 QuantError::EmptyCorpus
872 );
873 }
874
875 #[test]
876 fn sq8_try_train_flat_zero_dims_returns_error() {
877 assert_eq!(
878 Sq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
879 QuantError::ZeroDims
880 );
881 }
882
883 #[test]
884 fn gs_try_train_flat_zero_dims_returns_error() {
885 assert_eq!(
886 GsSq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
887 QuantError::ZeroDims
888 );
889 }
890
891 #[test]
892 fn sq8_try_train_flat_remainder_returns_error() {
893 let err = Sq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
895 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
896 }
897
898 #[test]
899 fn gs_try_train_flat_remainder_returns_error() {
900 let err = GsSq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
901 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
902 }
903
904 #[test]
905 fn sq8_try_encode_shorter_input_returns_error_not_malformed_vector() {
906 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
907 let err = codec.try_encode(&[0.0, 1.0]).unwrap_err();
908 assert_eq!(
909 err,
910 QuantError::EncodeLengthMismatch {
911 expected: 4,
912 got: 2
913 }
914 );
915 }
916
917 #[test]
918 fn sq8_try_encode_longer_input_returns_error() {
919 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
920 let err = codec.try_encode(&[0.0, 1.0, 2.0, 3.0, 4.0]).unwrap_err();
921 assert_eq!(
922 err,
923 QuantError::EncodeLengthMismatch {
924 expected: 4,
925 got: 5
926 }
927 );
928 }
929
930 #[test]
931 fn gs_try_encode_empty_input_returns_error_not_malformed_vector() {
932 let codec = GsSq8Codec::train_flat(&[1.0, 2.0, 3.0, 4.0], 1);
938 let err = codec.try_encode(&[]).unwrap_err();
939 assert_eq!(
940 err,
941 QuantError::EncodeLengthMismatch {
942 expected: 1,
943 got: 0
944 }
945 );
946 }
947
948 #[test]
949 fn sq8_try_encode_flat_par_zero_dims_returns_error() {
950 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
951 assert_eq!(
952 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
953 QuantError::ZeroDims
954 );
955 }
956
957 #[test]
958 fn gs_try_encode_flat_par_zero_dims_returns_error() {
959 let codec = GsSq8Codec::train(&rand_vecs(10, 4, 1));
960 assert_eq!(
961 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
962 QuantError::ZeroDims
963 );
964 }
965
966 #[test]
967 fn sq8_try_encode_flat_par_remainder_returns_error() {
968 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
969 let flat: Vec<f32> = (0..9).map(|i| i as f32).collect(); assert_eq!(
971 codec.try_encode_flat_par(&flat, 4).unwrap_err(),
972 QuantError::FlatLengthNotDivisible { len: 9, dims: 4 }
973 );
974 }
975
976 #[test]
977 fn sq8_try_encode_flat_par_dims_mismatch_returns_error() {
978 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
979 let flat: Vec<f32> = (0..6).map(|i| i as f32).collect();
980 let err = codec.try_encode_flat_par(&flat, 3).unwrap_err();
981 assert_eq!(
982 err,
983 QuantError::EncodeLengthMismatch {
984 expected: 4,
985 got: 3
986 }
987 );
988 }
989
990 #[test]
991 #[should_panic(expected = "row 1 has length 2, expected 1")]
992 fn sq8_train_still_panics_with_typed_message_on_ragged_rows() {
993 let _ = Sq8Codec::train(&[vec![0.0], vec![1.0, 2.0]]);
997 }
998
999 #[test]
1000 fn sq8_try_train_and_encode_roundtrip_matches_panicking_api() {
1001 let vecs = rand_vecs(20, 8, 7);
1002 let a = Sq8Codec::train(&vecs);
1003 let b = Sq8Codec::try_train(&vecs).expect("valid corpus must train");
1004 assert_eq!(a.min, b.min);
1005 assert_eq!(a.scale, b.scale);
1006 let ea = a.encode(&vecs[0]);
1007 let eb = b.try_encode(&vecs[0]).expect("valid vector must encode");
1008 assert_eq!(ea.codes, eb.codes);
1009 }
1010
1011 #[test]
1014 fn encode_decode_roundtrip_is_bounded() {
1015 let vecs = rand_vecs(100, 32, 42);
1016 let codec = Sq8Codec::train(&vecs);
1017 for v in &vecs {
1018 let ev = codec.encode(v);
1019 assert_eq!(ev.codes.len(), v.len());
1020 for (d, &code) in ev.codes.iter().enumerate() {
1021 let decoded = code as f32 * codec.scale[d] + codec.min[d];
1022 let err = (decoded - v[d]).abs();
1023 assert!(
1024 err <= codec.scale[d] + 1e-5,
1025 "dim {d}: err={err} scale={}",
1026 codec.scale[d]
1027 );
1028 }
1029 }
1030 }
1031
1032 #[test]
1033 fn approx_dot_relative_error_bounded() {
1034 let vecs = rand_vecs(200, 64, 77);
1035 let codec = Sq8Codec::train(&vecs);
1036 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1037
1038 let mut max_rel_err = 0.0f32;
1039 for i in 0..vecs.len() {
1040 for j in (i + 1)..vecs.len().min(i + 10) {
1041 let true_dot = dot_f32(&vecs[i], &vecs[j]);
1042 let approx = codec.approx_dot(&encoded[i], &encoded[j]);
1043 let denom = true_dot.abs().max(1e-3);
1044 let rel = (approx - true_dot).abs() / denom;
1045 if rel > max_rel_err {
1046 max_rel_err = rel;
1047 }
1048 }
1049 }
1050 assert!(
1051 max_rel_err < 0.15,
1052 "max relative dot error {max_rel_err:.4} >= 0.15"
1053 );
1054 }
1055
1056 #[test]
1057 fn approx_l2_sq_relative_error_bounded() {
1058 let vecs = rand_vecs(200, 64, 88);
1059 let codec = Sq8Codec::train(&vecs);
1060 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1061
1062 let mut max_rel_err = 0.0f32;
1063 for i in 0..vecs.len() {
1064 for j in (i + 1)..vecs.len().min(i + 10) {
1065 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1066 let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
1067 let denom = true_l2.max(1e-6);
1068 let rel = (approx - true_l2).abs() / denom;
1069 if rel > max_rel_err {
1070 max_rel_err = rel;
1071 }
1072 }
1073 }
1074 assert!(
1075 max_rel_err < 0.15,
1076 "max relative L2² error {max_rel_err:.4} >= 0.15"
1077 );
1078 }
1079
1080 #[test]
1081 fn order_preservation_triplets_cosine() {
1082 let vecs = rand_vecs(300, 64, 99);
1083 let codec = Sq8Codec::train(&vecs);
1084 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1085
1086 let n = vecs.len();
1087 let mut agree = 0usize;
1088 let mut total = 0usize;
1089
1090 for anchor in 0..50 {
1091 let a = &vecs[anchor];
1092 let ea = &encoded[anchor];
1093 for b_idx in 0..n {
1094 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1095 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1096 let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1097 let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1098
1099 let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
1100 let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
1101 let dist_ab_true = 1.0 - cos_ab;
1102 let dist_ac_true = 1.0 - cos_ac;
1103
1104 let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
1105 let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
1106
1107 if (dist_ab_true - dist_ac_true).abs() < 0.01 {
1108 continue;
1109 }
1110
1111 let true_closer_b = dist_ab_true < dist_ac_true;
1112 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1113 if true_closer_b == approx_closer_b {
1114 agree += 1;
1115 }
1116 total += 1;
1117 }
1118 }
1119 }
1120
1121 let rate = agree as f64 / total.max(1) as f64;
1122 assert!(
1123 rate >= 0.95,
1124 "order preservation {rate:.3} < 0.95 ({agree}/{total})"
1125 );
1126 }
1127
1128 #[test]
1129 fn order_preservation_triplets_l2() {
1130 let vecs = rand_vecs(300, 64, 101);
1131 let codec = Sq8Codec::train(&vecs);
1132 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1133
1134 let n = vecs.len();
1135 let mut agree = 0usize;
1136 let mut total = 0usize;
1137
1138 for anchor in 0..50 {
1139 let a = &vecs[anchor];
1140 let ea = &encoded[anchor];
1141 for b_idx in 0..n {
1142 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1143 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1144 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1145
1146 let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
1147 let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
1148
1149 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1150 continue;
1151 }
1152
1153 let true_closer_b = dist_ab_true < dist_ac_true;
1154 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1155 if true_closer_b == approx_closer_b {
1156 agree += 1;
1157 }
1158 total += 1;
1159 }
1160 }
1161 }
1162
1163 let rate = agree as f64 / total.max(1) as f64;
1164 assert!(
1165 rate >= 0.95,
1166 "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1167 );
1168 }
1169
1170 #[test]
1171 fn train_flat_matches_train_rows() {
1172 let vecs = rand_vecs(50, 16, 123);
1173 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1174
1175 let codec_rows = Sq8Codec::train(&vecs);
1176 let codec_flat = Sq8Codec::train_flat(&flat, 16);
1177
1178 for d in 0..16 {
1179 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1180 assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
1181 }
1182 }
1183
1184 #[test]
1185 fn encode_par_matches_sequential() {
1186 let vecs = rand_vecs(50, 32, 555);
1187 let codec = Sq8Codec::train(&vecs);
1188
1189 let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1190 let par = codec.encode_par(&vecs);
1191
1192 assert_eq!(seq.len(), par.len());
1193 for (s, p) in seq.iter().zip(par.iter()) {
1194 assert_eq!(s.codes, p.codes);
1195 assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
1196 }
1197 }
1198
1199 #[test]
1200 fn sq8_try_encode_par_short_row_returns_error_not_panic() {
1201 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1202 let mut rows = rand_vecs(5, 4, 2);
1203 rows[3] = vec![0.0, 1.0];
1204 let err = codec.try_encode_par(&rows).unwrap_err();
1205 assert_eq!(
1206 err,
1207 QuantError::EncodeLengthMismatch {
1208 expected: 4,
1209 got: 2
1210 }
1211 );
1212 }
1213
1214 #[test]
1215 fn sq8_try_encode_par_long_row_returns_error_not_panic() {
1216 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1217 let mut rows = rand_vecs(5, 4, 2);
1218 rows[3] = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
1219 let err = codec.try_encode_par(&rows).unwrap_err();
1220 assert_eq!(
1221 err,
1222 QuantError::EncodeLengthMismatch {
1223 expected: 4,
1224 got: 6
1225 }
1226 );
1227 }
1228
1229 #[test]
1230 #[should_panic(expected = "vector length 2 does not match codec dims 4")]
1231 fn sq8_encode_par_still_panics_with_typed_message_on_short_row() {
1232 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1233 let mut rows = rand_vecs(5, 4, 2);
1234 rows[3] = vec![0.0, 1.0];
1235 let _ = codec.encode_par(&rows);
1236 }
1237
1238 #[test]
1239 fn u8_dot_u32_matches_scalar() {
1240 let a: Vec<u8> = (0u8..=255).take(384).collect();
1241 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1242 let scalar: u32 = a
1243 .iter()
1244 .zip(b.iter())
1245 .map(|(&x, &y)| x as u32 * y as u32)
1246 .sum();
1247 assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
1248 }
1249
1250 #[test]
1251 fn u8_helpers_tail_path_max_diff() {
1252 for len in [1usize, 7, 15, 17, 100, 383] {
1253 let a = vec![255u8; len];
1254 let b = vec![0u8; len];
1255 assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
1256 assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
1257 }
1258 }
1259
1260 #[test]
1261 fn u8_l2sq_u32_matches_scalar() {
1262 let a: Vec<u8> = (0u8..=255).take(384).collect();
1263 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1264 let scalar: u32 = a
1265 .iter()
1266 .zip(b.iter())
1267 .map(|(&x, &y)| {
1268 let d = (x as i32) - (y as i32);
1269 (d * d) as u32
1270 })
1271 .sum();
1272 assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
1273 }
1274
1275 #[test]
1276 #[should_panic(expected = "u8_l2sq_u32 inputs must have equal length")]
1277 fn u8_l2sq_u32_rejects_shorter_second_slice() {
1278 let a = [1u8; 16];
1279 let b = [2u8; 1];
1280
1281 let _ = u8_l2sq_u32(&a, &b);
1282 }
1283
1284 #[test]
1292 fn gs_l2_sq_anisotropic_ordering_preserved() {
1293 let corpus = vec![
1294 vec![0.0f32, 0.0f32], vec![1.0f32, 1.0f32], vec![1.0f32, 4001.0f32], ];
1298 let codec = GsSq8Codec::train(&corpus);
1299
1300 let enc_origin = codec.encode(&corpus[0]);
1301 let enc_near = codec.encode(&corpus[1]);
1302 let enc_far = codec.encode(&corpus[2]);
1303
1304 let d_near = codec.l2_sq(&enc_origin, &enc_near);
1305 let d_far = codec.l2_sq(&enc_origin, &enc_far);
1306
1307 assert!(
1308 d_near < d_far,
1309 "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
1310 (anisotropy_ratio={:.1})",
1311 codec.anisotropy_ratio
1312 );
1313 }
1314
1315 #[test]
1316 fn gs_l2_sq_isotropic_small_error() {
1317 let vecs = rand_vecs(200, 64, 202);
1318 let codec = GsSq8Codec::train(&vecs);
1319 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1320
1321 let mut max_rel = 0.0f32;
1322 for i in 0..vecs.len() {
1323 for j in (i + 1)..vecs.len().min(i + 10) {
1324 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1325 let approx = codec.l2_sq(&encoded[i], &encoded[j]);
1326 let denom = true_l2.max(1e-6);
1327 let rel = (approx - true_l2).abs() / denom;
1328 if rel > max_rel {
1329 max_rel = rel;
1330 }
1331 }
1332 }
1333 assert!(
1334 max_rel < 0.15,
1335 "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
1336 );
1337 }
1338
1339 #[test]
1340 fn gs_train_flat_matches_train_rows() {
1341 let vecs = rand_vecs(50, 16, 321);
1342 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1343
1344 let codec_rows = GsSq8Codec::train(&vecs);
1345 let codec_flat = GsSq8Codec::train_flat(&flat, 16);
1346
1347 assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
1348 for d in 0..16 {
1349 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1350 }
1351 }
1352
1353 #[test]
1354 fn gs_l2_sq_order_preservation_triplets() {
1355 let vecs = rand_vecs(300, 64, 303);
1356 let codec = GsSq8Codec::train(&vecs);
1357 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1358
1359 let n = vecs.len();
1360 let mut agree = 0usize;
1361 let mut total = 0usize;
1362
1363 for anchor in 0..50 {
1364 let a = &vecs[anchor];
1365 let ea = &encoded[anchor];
1366 for b_idx in 0..n {
1367 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1368 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1369 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1370
1371 let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
1372 let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
1373
1374 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1375 continue;
1376 }
1377
1378 let true_closer_b = dist_ab_true < dist_ac_true;
1379 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1380 if true_closer_b == approx_closer_b {
1381 agree += 1;
1382 }
1383 total += 1;
1384 }
1385 }
1386 }
1387
1388 let rate = agree as f64 / total.max(1) as f64;
1389 assert!(
1390 rate >= 0.95,
1391 "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1392 );
1393 }
1394}