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