1use rayon::prelude::*;
30
31#[inline(always)]
39fn u8_dot_u32(a: &[u8], b: &[u8]) -> u32 {
40 #[cfg(target_arch = "aarch64")]
41 {
42 use std::arch::aarch64::*;
43 let n = a.len();
44 let chunks = n / 16;
45 let rem = n % 16;
46
47 let mut acc0: uint32x4_t;
48 let mut acc1: uint32x4_t;
49 let mut acc2: uint32x4_t;
50 let mut acc3: uint32x4_t;
51
52 unsafe {
53 acc0 = vdupq_n_u32(0);
54 acc1 = vdupq_n_u32(0);
55 acc2 = vdupq_n_u32(0);
56 acc3 = vdupq_n_u32(0);
57
58 for i in 0..chunks {
59 let ap = a.as_ptr().add(i * 16);
60 let bp = b.as_ptr().add(i * 16);
61
62 let va = vld1q_u8(ap);
63 let vb = vld1q_u8(bp);
64
65 let lo_u16 = vmull_u8(vget_low_u8(va), vget_low_u8(vb));
66 let hi_u16 = vmull_high_u8(va, vb);
67
68 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
69 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
70 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
71 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
72 }
73
74 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
75 let mut total = vaddvq_u32(sum4);
76
77 for i in (n - rem)..n {
78 total += a[i] as u32 * b[i] as u32;
79 }
80 total
81 }
82 }
83
84 #[cfg(not(target_arch = "aarch64"))]
85 {
86 a.chunks(8)
87 .zip(b.chunks(8))
88 .map(|(ac, bc)| {
89 ac.iter()
90 .zip(bc.iter())
91 .map(|(&x, &y)| (x as u32) * (y as u32))
92 .sum::<u32>()
93 })
94 .sum()
95 }
96}
97
98#[inline(always)]
104pub fn u8_l2sq_u32(a: &[u8], b: &[u8]) -> u32 {
105 assert_eq!(
106 a.len(),
107 b.len(),
108 "u8_l2sq_u32 inputs must have equal length"
109 );
110
111 #[cfg(target_arch = "aarch64")]
112 {
113 use std::arch::aarch64::*;
114 let n = a.len();
115 let chunks = n / 16;
116 let rem = n % 16;
117
118 let mut acc0: uint32x4_t;
119 let mut acc1: uint32x4_t;
120 let mut acc2: uint32x4_t;
121 let mut acc3: uint32x4_t;
122
123 unsafe {
124 acc0 = vdupq_n_u32(0);
125 acc1 = vdupq_n_u32(0);
126 acc2 = vdupq_n_u32(0);
127 acc3 = vdupq_n_u32(0);
128
129 for i in 0..chunks {
130 let ap = a.as_ptr().add(i * 16);
131 let bp = b.as_ptr().add(i * 16);
132
133 let va = vld1q_u8(ap);
134 let vb = vld1q_u8(bp);
135
136 let diff = vabdq_u8(va, vb);
137
138 let lo_u16 = vmull_u8(vget_low_u8(diff), vget_low_u8(diff));
139 let hi_u16 = vmull_high_u8(diff, diff);
140
141 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
142 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
143 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
144 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
145 }
146
147 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
148 let mut total = vaddvq_u32(sum4);
149
150 for i in (n - rem)..n {
151 let d = (a[i] as i32) - (b[i] as i32);
152 total += (d * d) as u32;
153 }
154 total
155 }
156 }
157
158 #[cfg(not(target_arch = "aarch64"))]
159 {
160 a.chunks(8)
161 .zip(b.chunks(8))
162 .map(|(ac, bc)| {
163 ac.iter()
164 .zip(bc.iter())
165 .map(|(&x, &y)| {
166 let d = (x as i32) - (y as i32);
167 (d * d) as u32
168 })
169 .sum::<u32>()
170 })
171 .sum()
172 }
173}
174
175#[derive(Debug, Clone, PartialEq)]
183pub enum QuantError {
184 EmptyCorpus,
186 ZeroDims,
188 FlatLengthNotDivisible { len: usize, dims: usize },
190 RaggedRow {
192 row: usize,
193 expected: usize,
194 got: usize,
195 },
196 EncodeLengthMismatch { expected: usize, got: usize },
199}
200
201impl std::fmt::Display for QuantError {
202 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
203 match self {
204 Self::EmptyCorpus => write!(f, "cannot train on empty corpus"),
205 Self::ZeroDims => write!(f, "dims must be > 0"),
206 Self::FlatLengthNotDivisible { len, dims } => write!(
207 f,
208 "flat vector length {len} is not a multiple of dims {dims}"
209 ),
210 Self::RaggedRow { row, expected, got } => write!(
211 f,
212 "row {row} has length {got}, expected {expected} (dims fixed by row 0)"
213 ),
214 Self::EncodeLengthMismatch { expected, got } => write!(
215 f,
216 "vector length {got} does not match codec dims {expected}"
217 ),
218 }
219 }
220}
221
222impl std::error::Error for QuantError {}
223
224fn flat_min_max(vectors: &[f32], dims: usize) -> Result<(Vec<f32>, Vec<f32>), QuantError> {
229 if dims == 0 {
230 return Err(QuantError::ZeroDims);
231 }
232 if vectors.is_empty() {
233 return Err(QuantError::EmptyCorpus);
234 }
235 if !vectors.len().is_multiple_of(dims) {
236 return Err(QuantError::FlatLengthNotDivisible {
237 len: vectors.len(),
238 dims,
239 });
240 }
241
242 let n = vectors.len() / dims;
243 let mut min = vec![f32::INFINITY; dims];
244 let mut max = vec![f32::NEG_INFINITY; dims];
245
246 for row in 0..n {
247 let v = &vectors[row * dims..(row + 1) * dims];
248 for (d, &x) in v.iter().enumerate() {
249 if x.is_finite() {
250 if x < min[d] {
251 min[d] = x;
252 }
253 if x > max[d] {
254 max[d] = x;
255 }
256 }
257 }
258 }
259 finalize_min_max(&mut min, &mut max);
260 Ok((min, max))
261}
262
263fn 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> {
418 let dims = self.min.len();
419 if v.len() != dims {
420 return Err(QuantError::EncodeLengthMismatch {
421 expected: dims,
422 got: v.len(),
423 });
424 }
425 Ok(self.encode_unchecked(v))
426 }
427
428 fn encode_unchecked(&self, v: &[f32]) -> EncodedVector {
430 let dims = self.min.len();
431 let mut codes = Vec::with_capacity(dims);
432 let mut soc_sum = 0.0f32;
433 let mut residual_dot_bias = 0.0f32;
434 let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
435
436 for (d, &x) in v.iter().enumerate() {
437 let s = self.scale[d];
438 let inv_s = if s > 1e-12 { 1.0 / s } else { 0.0 };
439 let raw = (x - self.min[d]) * inv_s;
440 let code = raw.round().clamp(0.0, 255.0) as u8;
441 codes.push(code);
442 soc_sum += s * self.min[d] * code as f32;
443 residual_dot_bias += self.scale_sq_residual[d] * code as f32;
444 }
445
446 EncodedVector {
447 codes,
448 norm,
449 soc_sum,
450 residual_dot_bias,
451 }
452 }
453
454 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<EncodedVector> {
459 self.try_encode_flat_par(vectors, dims)
460 .unwrap_or_else(|e| panic!("{e}"))
461 }
462
463 pub fn try_encode_flat_par(
469 &self,
470 vectors: &[f32],
471 dims: usize,
472 ) -> Result<Vec<EncodedVector>, QuantError> {
473 if dims == 0 {
474 return Err(QuantError::ZeroDims);
475 }
476 if !vectors.len().is_multiple_of(dims) {
477 return Err(QuantError::FlatLengthNotDivisible {
478 len: vectors.len(),
479 dims,
480 });
481 }
482 if dims != self.min.len() {
483 return Err(QuantError::EncodeLengthMismatch {
484 expected: self.min.len(),
485 got: dims,
486 });
487 }
488 let n = vectors.len() / dims;
489 Ok((0..n)
490 .into_par_iter()
491 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
492 .collect())
493 }
494
495 pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
501 self.try_encode_par(vectors)
502 .unwrap_or_else(|e| panic!("{e}"))
503 }
504
505 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 Ok(vectors
521 .par_iter()
522 .map(|v| self.encode_unchecked(v))
523 .collect())
524 }
525
526 #[inline]
535 pub fn approx_dot(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
536 let raw = u8_dot_u32(&a.codes, &b.codes) as f32;
537 let residual_hot: f32 = self
538 .scale_sq_residual
539 .iter()
540 .zip(a.codes.iter())
541 .zip(b.codes.iter())
542 .map(|((r, &ac), &bc)| r * (ac as f32) * (bc as f32))
543 .sum();
544 self.mean_scale_sq * raw + residual_hot + a.soc_sum + b.soc_sum + self.offset_sq_sum
545 }
546
547 #[inline]
551 pub fn approx_cosine_dist(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
552 let denom = a.norm * b.norm;
553 if !denom.is_finite() || denom <= 0.0 {
554 return 1.0;
555 }
556 let dot = self.approx_dot(a, b);
557 let cosine = (dot / denom).clamp(-1.0, 1.0);
558 1.0 - cosine
559 }
560
561 #[inline]
573 pub fn approx_l2_sq(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
574 let raw = u8_l2sq_u32(&a.codes, &b.codes) as f32;
575 let residual_hot: f32 = self
576 .scale_sq_residual
577 .iter()
578 .zip(a.codes.iter())
579 .zip(b.codes.iter())
580 .map(|((r, &ac), &bc)| {
581 let d = (ac as i32) - (bc as i32);
582 r * (d as f32) * (d as f32)
583 })
584 .sum();
585 self.mean_scale_sq * raw + residual_hot
586 }
587
588 pub fn dims(&self) -> usize {
590 self.min.len()
591 }
592}
593
594#[derive(Debug, Clone)]
616pub struct GsSq8Codec {
617 pub min: Vec<f32>,
619 pub gs: f32,
621 pub gs_sq: f32,
623 pub anisotropy_ratio: f32,
626}
627
628#[derive(Debug, Clone)]
630pub struct GsEncodedVector {
631 pub codes: Vec<u8>,
633}
634
635impl GsSq8Codec {
636 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
637 let dims = min.len();
638 let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
639 let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
640 let gs = if max_range > 1e-12 {
641 max_range / 255.0
642 } else {
643 1.0 / 255.0
644 };
645
646 let min_range_nonzero = ranges
647 .iter()
648 .cloned()
649 .filter(|&r| r > 1e-12)
650 .fold(f32::INFINITY, f32::min);
651 let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
652 max_range / min_range_nonzero
653 } else {
654 1.0
655 };
656
657 Self {
658 min,
659 gs,
660 gs_sq: gs * gs,
661 anisotropy_ratio,
662 }
663 }
664
665 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
670 Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
671 }
672
673 pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
676 let (min, max) = flat_min_max(vectors, dims)?;
677 Ok(Self::build_from_min_max(min, max))
678 }
679
680 pub fn train(vectors: &[Vec<f32>]) -> Self {
685 Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
686 }
687
688 pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
693 let (_dims, min, max) = row_min_max(vectors)?;
694 Ok(Self::build_from_min_max(min, max))
695 }
696
697 #[inline]
703 pub fn encode(&self, v: &[f32]) -> GsEncodedVector {
704 self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
705 }
706
707 pub fn try_encode(&self, v: &[f32]) -> Result<GsEncodedVector, QuantError> {
714 let dims = self.min.len();
715 if v.len() != dims {
716 return Err(QuantError::EncodeLengthMismatch {
717 expected: dims,
718 got: v.len(),
719 });
720 }
721 Ok(self.encode_unchecked(v))
722 }
723
724 #[inline]
726 fn encode_unchecked(&self, v: &[f32]) -> GsEncodedVector {
727 let inv_gs = if self.gs > 1e-12 { 1.0 / self.gs } else { 0.0 };
728 let codes = v
729 .iter()
730 .enumerate()
731 .map(|(d, &x)| ((x - self.min[d]) * inv_gs).round().clamp(0.0, 255.0) as u8)
732 .collect();
733 GsEncodedVector { codes }
734 }
735
736 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<GsEncodedVector> {
741 self.try_encode_flat_par(vectors, dims)
742 .unwrap_or_else(|e| panic!("{e}"))
743 }
744
745 pub fn try_encode_flat_par(
750 &self,
751 vectors: &[f32],
752 dims: usize,
753 ) -> Result<Vec<GsEncodedVector>, QuantError> {
754 if dims == 0 {
755 return Err(QuantError::ZeroDims);
756 }
757 if !vectors.len().is_multiple_of(dims) {
758 return Err(QuantError::FlatLengthNotDivisible {
759 len: vectors.len(),
760 dims,
761 });
762 }
763 if dims != self.min.len() {
764 return Err(QuantError::EncodeLengthMismatch {
765 expected: self.min.len(),
766 got: dims,
767 });
768 }
769 let n = vectors.len() / dims;
770 Ok((0..n)
771 .into_par_iter()
772 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
773 .collect())
774 }
775
776 #[inline]
785 pub fn l2_sq(&self, a: &GsEncodedVector, b: &GsEncodedVector) -> f32 {
786 self.gs_sq * u8_l2sq_u32(&a.codes, &b.codes) as f32
787 }
788
789 pub fn dims(&self) -> usize {
791 self.min.len()
792 }
793
794 #[inline]
800 pub fn is_in_distribution(&self, v: &[f32]) -> bool {
801 let max_code = 255.0 * self.gs;
802 v.iter()
803 .zip(self.min.iter())
804 .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
805 }
806}
807
808#[cfg(test)]
809mod tests {
810 use super::*;
811
812 fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
813 let mut h = seed;
814 (0..n)
815 .map(|_| {
816 (0..dims)
817 .map(|_| {
818 h = h
819 .wrapping_mul(0x6c62_272e_07bb_0142)
820 .wrapping_add(0x62b8_2175_62d9_6b1a);
821 let bits = (h >> 33) as u32;
822 (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
823 })
824 .collect()
825 })
826 .collect()
827 }
828
829 fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
830 a.iter().zip(b).map(|(x, y)| x * y).sum()
831 }
832
833 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
834 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
835 }
836
837 #[test]
840 fn sq8_try_train_ragged_rows_returns_error_not_panic() {
841 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
842 let err = Sq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
843 assert_eq!(
844 err,
845 QuantError::RaggedRow {
846 row: 1,
847 expected: 1,
848 got: 2
849 }
850 );
851 }
852
853 #[test]
854 fn gs_try_train_ragged_rows_returns_error_not_panic() {
855 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
856 let err = GsSq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
857 assert_eq!(
858 err,
859 QuantError::RaggedRow {
860 row: 1,
861 expected: 1,
862 got: 2
863 }
864 );
865 }
866
867 #[test]
868 fn sq8_try_train_empty_corpus_returns_error() {
869 let vecs: Vec<Vec<f32>> = vec![];
870 assert_eq!(
871 Sq8Codec::try_train(&vecs).unwrap_err(),
872 QuantError::EmptyCorpus
873 );
874 }
875
876 #[test]
877 fn sq8_try_train_flat_zero_dims_returns_error() {
878 assert_eq!(
879 Sq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
880 QuantError::ZeroDims
881 );
882 }
883
884 #[test]
885 fn gs_try_train_flat_zero_dims_returns_error() {
886 assert_eq!(
887 GsSq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
888 QuantError::ZeroDims
889 );
890 }
891
892 #[test]
893 fn sq8_try_train_flat_remainder_returns_error() {
894 let err = Sq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
896 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
897 }
898
899 #[test]
900 fn gs_try_train_flat_remainder_returns_error() {
901 let err = GsSq8Codec::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 sq8_try_encode_shorter_input_returns_error_not_malformed_vector() {
907 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
908 let err = codec.try_encode(&[0.0, 1.0]).unwrap_err();
909 assert_eq!(
910 err,
911 QuantError::EncodeLengthMismatch {
912 expected: 4,
913 got: 2
914 }
915 );
916 }
917
918 #[test]
919 fn sq8_try_encode_longer_input_returns_error() {
920 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
921 let err = codec.try_encode(&[0.0, 1.0, 2.0, 3.0, 4.0]).unwrap_err();
922 assert_eq!(
923 err,
924 QuantError::EncodeLengthMismatch {
925 expected: 4,
926 got: 5
927 }
928 );
929 }
930
931 #[test]
932 fn gs_try_encode_empty_input_returns_error_not_malformed_vector() {
933 let codec = GsSq8Codec::train_flat(&[1.0, 2.0, 3.0, 4.0], 1);
939 let err = codec.try_encode(&[]).unwrap_err();
940 assert_eq!(
941 err,
942 QuantError::EncodeLengthMismatch {
943 expected: 1,
944 got: 0
945 }
946 );
947 }
948
949 #[test]
950 fn sq8_try_encode_flat_par_zero_dims_returns_error() {
951 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
952 assert_eq!(
953 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
954 QuantError::ZeroDims
955 );
956 }
957
958 #[test]
959 fn gs_try_encode_flat_par_zero_dims_returns_error() {
960 let codec = GsSq8Codec::train(&rand_vecs(10, 4, 1));
961 assert_eq!(
962 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
963 QuantError::ZeroDims
964 );
965 }
966
967 #[test]
968 fn sq8_try_encode_flat_par_remainder_returns_error() {
969 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
970 let flat: Vec<f32> = (0..9).map(|i| i as f32).collect(); assert_eq!(
972 codec.try_encode_flat_par(&flat, 4).unwrap_err(),
973 QuantError::FlatLengthNotDivisible { len: 9, dims: 4 }
974 );
975 }
976
977 #[test]
978 fn sq8_try_encode_flat_par_dims_mismatch_returns_error() {
979 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
980 let flat: Vec<f32> = (0..6).map(|i| i as f32).collect();
981 let err = codec.try_encode_flat_par(&flat, 3).unwrap_err();
982 assert_eq!(
983 err,
984 QuantError::EncodeLengthMismatch {
985 expected: 4,
986 got: 3
987 }
988 );
989 }
990
991 #[test]
992 #[should_panic(expected = "row 1 has length 2, expected 1")]
993 fn sq8_train_still_panics_with_typed_message_on_ragged_rows() {
994 let _ = Sq8Codec::train(&[vec![0.0], vec![1.0, 2.0]]);
998 }
999
1000 #[test]
1001 fn sq8_try_train_and_encode_roundtrip_matches_panicking_api() {
1002 let vecs = rand_vecs(20, 8, 7);
1003 let a = Sq8Codec::train(&vecs);
1004 let b = Sq8Codec::try_train(&vecs).expect("valid corpus must train");
1005 assert_eq!(a.min, b.min);
1006 assert_eq!(a.scale, b.scale);
1007 let ea = a.encode(&vecs[0]);
1008 let eb = b.try_encode(&vecs[0]).expect("valid vector must encode");
1009 assert_eq!(ea.codes, eb.codes);
1010 }
1011
1012 #[test]
1015 fn encode_decode_roundtrip_is_bounded() {
1016 let vecs = rand_vecs(100, 32, 42);
1017 let codec = Sq8Codec::train(&vecs);
1018 for v in &vecs {
1019 let ev = codec.encode(v);
1020 assert_eq!(ev.codes.len(), v.len());
1021 for (d, &code) in ev.codes.iter().enumerate() {
1022 let decoded = code as f32 * codec.scale[d] + codec.min[d];
1023 let err = (decoded - v[d]).abs();
1024 assert!(
1025 err <= codec.scale[d] + 1e-5,
1026 "dim {d}: err={err} scale={}",
1027 codec.scale[d]
1028 );
1029 }
1030 }
1031 }
1032
1033 #[test]
1034 fn approx_dot_relative_error_bounded() {
1035 let vecs = rand_vecs(200, 64, 77);
1036 let codec = Sq8Codec::train(&vecs);
1037 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1038
1039 let mut max_rel_err = 0.0f32;
1040 for i in 0..vecs.len() {
1041 for j in (i + 1)..vecs.len().min(i + 10) {
1042 let true_dot = dot_f32(&vecs[i], &vecs[j]);
1043 let approx = codec.approx_dot(&encoded[i], &encoded[j]);
1044 let denom = true_dot.abs().max(1e-3);
1045 let rel = (approx - true_dot).abs() / denom;
1046 if rel > max_rel_err {
1047 max_rel_err = rel;
1048 }
1049 }
1050 }
1051 assert!(
1052 max_rel_err < 0.15,
1053 "max relative dot error {max_rel_err:.4} >= 0.15"
1054 );
1055 }
1056
1057 #[test]
1058 fn approx_l2_sq_relative_error_bounded() {
1059 let vecs = rand_vecs(200, 64, 88);
1060 let codec = Sq8Codec::train(&vecs);
1061 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1062
1063 let mut max_rel_err = 0.0f32;
1064 for i in 0..vecs.len() {
1065 for j in (i + 1)..vecs.len().min(i + 10) {
1066 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1067 let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
1068 let denom = true_l2.max(1e-6);
1069 let rel = (approx - true_l2).abs() / denom;
1070 if rel > max_rel_err {
1071 max_rel_err = rel;
1072 }
1073 }
1074 }
1075 assert!(
1076 max_rel_err < 0.15,
1077 "max relative L2² error {max_rel_err:.4} >= 0.15"
1078 );
1079 }
1080
1081 #[test]
1082 fn order_preservation_triplets_cosine() {
1083 let vecs = rand_vecs(300, 64, 99);
1084 let codec = Sq8Codec::train(&vecs);
1085 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1086
1087 let n = vecs.len();
1088 let mut agree = 0usize;
1089 let mut total = 0usize;
1090
1091 for anchor in 0..50 {
1092 let a = &vecs[anchor];
1093 let ea = &encoded[anchor];
1094 for b_idx in 0..n {
1095 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1096 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1097 let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1098 let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1099
1100 let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
1101 let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
1102 let dist_ab_true = 1.0 - cos_ab;
1103 let dist_ac_true = 1.0 - cos_ac;
1104
1105 let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
1106 let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
1107
1108 if (dist_ab_true - dist_ac_true).abs() < 0.01 {
1109 continue;
1110 }
1111
1112 let true_closer_b = dist_ab_true < dist_ac_true;
1113 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1114 if true_closer_b == approx_closer_b {
1115 agree += 1;
1116 }
1117 total += 1;
1118 }
1119 }
1120 }
1121
1122 let rate = agree as f64 / total.max(1) as f64;
1123 assert!(
1124 rate >= 0.95,
1125 "order preservation {rate:.3} < 0.95 ({agree}/{total})"
1126 );
1127 }
1128
1129 #[test]
1130 fn order_preservation_triplets_l2() {
1131 let vecs = rand_vecs(300, 64, 101);
1132 let codec = Sq8Codec::train(&vecs);
1133 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1134
1135 let n = vecs.len();
1136 let mut agree = 0usize;
1137 let mut total = 0usize;
1138
1139 for anchor in 0..50 {
1140 let a = &vecs[anchor];
1141 let ea = &encoded[anchor];
1142 for b_idx in 0..n {
1143 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1144 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1145 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1146
1147 let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
1148 let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
1149
1150 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1151 continue;
1152 }
1153
1154 let true_closer_b = dist_ab_true < dist_ac_true;
1155 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1156 if true_closer_b == approx_closer_b {
1157 agree += 1;
1158 }
1159 total += 1;
1160 }
1161 }
1162 }
1163
1164 let rate = agree as f64 / total.max(1) as f64;
1165 assert!(
1166 rate >= 0.95,
1167 "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1168 );
1169 }
1170
1171 #[test]
1172 fn train_flat_matches_train_rows() {
1173 let vecs = rand_vecs(50, 16, 123);
1174 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1175
1176 let codec_rows = Sq8Codec::train(&vecs);
1177 let codec_flat = Sq8Codec::train_flat(&flat, 16);
1178
1179 for d in 0..16 {
1180 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1181 assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
1182 }
1183 }
1184
1185 #[test]
1186 fn encode_par_matches_sequential() {
1187 let vecs = rand_vecs(50, 32, 555);
1188 let codec = Sq8Codec::train(&vecs);
1189
1190 let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1191 let par = codec.encode_par(&vecs);
1192
1193 assert_eq!(seq.len(), par.len());
1194 for (s, p) in seq.iter().zip(par.iter()) {
1195 assert_eq!(s.codes, p.codes);
1196 assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
1197 }
1198 }
1199
1200 #[test]
1201 fn sq8_try_encode_par_short_row_returns_error_not_panic() {
1202 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1203 let mut rows = rand_vecs(5, 4, 2);
1204 rows[3] = vec![0.0, 1.0];
1205 let err = codec.try_encode_par(&rows).unwrap_err();
1206 assert_eq!(
1207 err,
1208 QuantError::EncodeLengthMismatch {
1209 expected: 4,
1210 got: 2
1211 }
1212 );
1213 }
1214
1215 #[test]
1216 fn sq8_try_encode_par_long_row_returns_error_not_panic() {
1217 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1218 let mut rows = rand_vecs(5, 4, 2);
1219 rows[3] = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
1220 let err = codec.try_encode_par(&rows).unwrap_err();
1221 assert_eq!(
1222 err,
1223 QuantError::EncodeLengthMismatch {
1224 expected: 4,
1225 got: 6
1226 }
1227 );
1228 }
1229
1230 #[test]
1231 #[should_panic(expected = "vector length 2 does not match codec dims 4")]
1232 fn sq8_encode_par_still_panics_with_typed_message_on_short_row() {
1233 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1234 let mut rows = rand_vecs(5, 4, 2);
1235 rows[3] = vec![0.0, 1.0];
1236 let _ = codec.encode_par(&rows);
1237 }
1238
1239 #[test]
1240 fn u8_dot_u32_matches_scalar() {
1241 let a: Vec<u8> = (0u8..=255).take(384).collect();
1242 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1243 let scalar: u32 = a
1244 .iter()
1245 .zip(b.iter())
1246 .map(|(&x, &y)| x as u32 * y as u32)
1247 .sum();
1248 assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
1249 }
1250
1251 #[test]
1252 fn u8_helpers_tail_path_max_diff() {
1253 for len in [1usize, 7, 15, 17, 100, 383] {
1254 let a = vec![255u8; len];
1255 let b = vec![0u8; len];
1256 assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
1257 assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
1258 }
1259 }
1260
1261 #[test]
1262 fn u8_l2sq_u32_matches_scalar() {
1263 let a: Vec<u8> = (0u8..=255).take(384).collect();
1264 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1265 let scalar: u32 = a
1266 .iter()
1267 .zip(b.iter())
1268 .map(|(&x, &y)| {
1269 let d = (x as i32) - (y as i32);
1270 (d * d) as u32
1271 })
1272 .sum();
1273 assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
1274 }
1275
1276 #[test]
1277 #[should_panic(expected = "u8_l2sq_u32 inputs must have equal length")]
1278 fn u8_l2sq_u32_rejects_shorter_second_slice() {
1279 let a = [1u8; 16];
1280 let b = [2u8; 1];
1281
1282 let _ = u8_l2sq_u32(&a, &b);
1283 }
1284
1285 #[test]
1293 fn gs_l2_sq_anisotropic_ordering_preserved() {
1294 let corpus = vec![
1295 vec![0.0f32, 0.0f32], vec![1.0f32, 1.0f32], vec![1.0f32, 4001.0f32], ];
1299 let codec = GsSq8Codec::train(&corpus);
1300
1301 let enc_origin = codec.encode(&corpus[0]);
1302 let enc_near = codec.encode(&corpus[1]);
1303 let enc_far = codec.encode(&corpus[2]);
1304
1305 let d_near = codec.l2_sq(&enc_origin, &enc_near);
1306 let d_far = codec.l2_sq(&enc_origin, &enc_far);
1307
1308 assert!(
1309 d_near < d_far,
1310 "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
1311 (anisotropy_ratio={:.1})",
1312 codec.anisotropy_ratio
1313 );
1314 }
1315
1316 #[test]
1317 fn gs_l2_sq_isotropic_small_error() {
1318 let vecs = rand_vecs(200, 64, 202);
1319 let codec = GsSq8Codec::train(&vecs);
1320 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1321
1322 let mut max_rel = 0.0f32;
1323 for i in 0..vecs.len() {
1324 for j in (i + 1)..vecs.len().min(i + 10) {
1325 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1326 let approx = codec.l2_sq(&encoded[i], &encoded[j]);
1327 let denom = true_l2.max(1e-6);
1328 let rel = (approx - true_l2).abs() / denom;
1329 if rel > max_rel {
1330 max_rel = rel;
1331 }
1332 }
1333 }
1334 assert!(
1335 max_rel < 0.15,
1336 "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
1337 );
1338 }
1339
1340 #[test]
1341 fn gs_train_flat_matches_train_rows() {
1342 let vecs = rand_vecs(50, 16, 321);
1343 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1344
1345 let codec_rows = GsSq8Codec::train(&vecs);
1346 let codec_flat = GsSq8Codec::train_flat(&flat, 16);
1347
1348 assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
1349 for d in 0..16 {
1350 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1351 }
1352 }
1353
1354 #[test]
1355 fn gs_l2_sq_order_preservation_triplets() {
1356 let vecs = rand_vecs(300, 64, 303);
1357 let codec = GsSq8Codec::train(&vecs);
1358 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1359
1360 let n = vecs.len();
1361 let mut agree = 0usize;
1362 let mut total = 0usize;
1363
1364 for anchor in 0..50 {
1365 let a = &vecs[anchor];
1366 let ea = &encoded[anchor];
1367 for b_idx in 0..n {
1368 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1369 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1370 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1371
1372 let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
1373 let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
1374
1375 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1376 continue;
1377 }
1378
1379 let true_closer_b = dist_ab_true < dist_ac_true;
1380 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1381 if true_closer_b == approx_closer_b {
1382 agree += 1;
1383 }
1384 total += 1;
1385 }
1386 }
1387 }
1388
1389 let rate = agree as f64 / total.max(1) as f64;
1390 assert!(
1391 rate >= 0.95,
1392 "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1393 );
1394 }
1395}