1#[cfg(feature = "parallel")]
33use rayon::prelude::*;
34
35#[cfg(test)]
44#[inline(always)]
45fn u8_dot_u32(a: &[u8], b: &[u8]) -> u32 {
46 assert_eq!(a.len(), b.len(), "u8_dot_u32 inputs must have equal length");
47
48 #[cfg(target_arch = "aarch64")]
49 {
50 use std::arch::aarch64::*;
51 let n = a.len();
52 let chunks = n / 16;
53 let rem = n % 16;
54
55 let mut acc0: uint32x4_t;
56 let mut acc1: uint32x4_t;
57 let mut acc2: uint32x4_t;
58 let mut acc3: uint32x4_t;
59
60 unsafe {
61 acc0 = vdupq_n_u32(0);
62 acc1 = vdupq_n_u32(0);
63 acc2 = vdupq_n_u32(0);
64 acc3 = vdupq_n_u32(0);
65
66 for i in 0..chunks {
67 let ap = a.as_ptr().add(i * 16);
68 let bp = b.as_ptr().add(i * 16);
69
70 let va = vld1q_u8(ap);
71 let vb = vld1q_u8(bp);
72
73 let lo_u16 = vmull_u8(vget_low_u8(va), vget_low_u8(vb));
74 let hi_u16 = vmull_high_u8(va, vb);
75
76 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
77 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
78 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
79 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
80 }
81
82 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
83 let mut total = vaddvq_u32(sum4);
84
85 for i in (n - rem)..n {
86 total += a[i] as u32 * b[i] as u32;
87 }
88 total
89 }
90 }
91
92 #[cfg(not(target_arch = "aarch64"))]
93 {
94 a.chunks(8)
95 .zip(b.chunks(8))
96 .map(|(ac, bc)| {
97 ac.iter()
98 .zip(bc.iter())
99 .map(|(&x, &y)| (x as u32) * (y as u32))
100 .sum::<u32>()
101 })
102 .sum()
103 }
104}
105
106#[inline(always)]
112pub fn u8_l2sq_u32(a: &[u8], b: &[u8]) -> u32 {
113 assert_eq!(
114 a.len(),
115 b.len(),
116 "u8_l2sq_u32 inputs must have equal length"
117 );
118
119 #[cfg(target_arch = "aarch64")]
120 {
121 use std::arch::aarch64::*;
122 let n = a.len();
123 let chunks = n / 16;
124 let rem = n % 16;
125
126 let mut acc0: uint32x4_t;
127 let mut acc1: uint32x4_t;
128 let mut acc2: uint32x4_t;
129 let mut acc3: uint32x4_t;
130
131 unsafe {
132 acc0 = vdupq_n_u32(0);
133 acc1 = vdupq_n_u32(0);
134 acc2 = vdupq_n_u32(0);
135 acc3 = vdupq_n_u32(0);
136
137 for i in 0..chunks {
138 let ap = a.as_ptr().add(i * 16);
139 let bp = b.as_ptr().add(i * 16);
140
141 let va = vld1q_u8(ap);
142 let vb = vld1q_u8(bp);
143
144 let diff = vabdq_u8(va, vb);
145
146 let lo_u16 = vmull_u8(vget_low_u8(diff), vget_low_u8(diff));
147 let hi_u16 = vmull_high_u8(diff, diff);
148
149 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
150 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
151 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
152 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
153 }
154
155 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
156 let mut total = vaddvq_u32(sum4);
157
158 for i in (n - rem)..n {
159 let d = (a[i] as i32) - (b[i] as i32);
160 total += (d * d) as u32;
161 }
162 total
163 }
164 }
165
166 #[cfg(not(target_arch = "aarch64"))]
167 {
168 a.chunks(8)
169 .zip(b.chunks(8))
170 .map(|(ac, bc)| {
171 ac.iter()
172 .zip(bc.iter())
173 .map(|(&x, &y)| {
174 let d = (x as i32) - (y as i32);
175 (d * d) as u32
176 })
177 .sum::<u32>()
178 })
179 .sum()
180 }
181}
182
183#[derive(Debug, Clone, PartialEq)]
191pub enum QuantError {
192 EmptyCorpus,
194 ZeroDims,
196 FlatLengthNotDivisible { len: usize, dims: usize },
198 RaggedRow {
200 row: usize,
201 expected: usize,
202 got: usize,
203 },
204 EncodeLengthMismatch { expected: usize, got: usize },
207}
208
209impl std::fmt::Display for QuantError {
210 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
211 match self {
212 Self::EmptyCorpus => write!(f, "cannot train on empty corpus"),
213 Self::ZeroDims => write!(f, "dims must be > 0"),
214 Self::FlatLengthNotDivisible { len, dims } => write!(
215 f,
216 "flat vector length {len} is not a multiple of dims {dims}"
217 ),
218 Self::RaggedRow { row, expected, got } => write!(
219 f,
220 "row {row} has length {got}, expected {expected} (dims fixed by row 0)"
221 ),
222 Self::EncodeLengthMismatch { expected, got } => write!(
223 f,
224 "vector length {got} does not match codec dims {expected}"
225 ),
226 }
227 }
228}
229
230impl std::error::Error for QuantError {}
231
232fn flat_min_max(vectors: &[f32], dims: usize) -> Result<(Vec<f32>, Vec<f32>), QuantError> {
236 if dims == 0 {
237 return Err(QuantError::ZeroDims);
238 }
239 if vectors.is_empty() {
240 return Err(QuantError::EmptyCorpus);
241 }
242 if !vectors.len().is_multiple_of(dims) {
243 return Err(QuantError::FlatLengthNotDivisible {
244 len: vectors.len(),
245 dims,
246 });
247 }
248
249 let n = vectors.len() / dims;
250 let mut min = vec![f32::INFINITY; dims];
251 let mut max = vec![f32::NEG_INFINITY; dims];
252
253 for row in 0..n {
254 let v = &vectors[row * dims..(row + 1) * dims];
255 for (d, &x) in v.iter().enumerate() {
256 if x.is_finite() {
257 if x < min[d] {
258 min[d] = x;
259 }
260 if x > max[d] {
261 max[d] = x;
262 }
263 }
264 }
265 }
266 finalize_min_max(&mut min, &mut max);
267 Ok((min, max))
268}
269
270fn row_min_max(vectors: &[Vec<f32>]) -> Result<(usize, Vec<f32>, Vec<f32>), QuantError> {
273 if vectors.is_empty() {
274 return Err(QuantError::EmptyCorpus);
275 }
276 let dims = vectors[0].len();
277 if dims == 0 {
278 return Err(QuantError::ZeroDims);
279 }
280
281 let mut min = vec![f32::INFINITY; dims];
282 let mut max = vec![f32::NEG_INFINITY; dims];
283
284 for (row, v) in vectors.iter().enumerate() {
285 if v.len() != dims {
286 return Err(QuantError::RaggedRow {
287 row,
288 expected: dims,
289 got: v.len(),
290 });
291 }
292 for (d, &x) in v.iter().enumerate() {
293 if x.is_finite() {
294 if x < min[d] {
295 min[d] = x;
296 }
297 if x > max[d] {
298 max[d] = x;
299 }
300 }
301 }
302 }
303 finalize_min_max(&mut min, &mut max);
304 Ok((dims, min, max))
305}
306
307fn finalize_min_max(min: &mut [f32], max: &mut [f32]) {
310 for d in 0..min.len() {
311 if !min[d].is_finite() {
312 min[d] = 0.0;
313 }
314 if !max[d].is_finite() || max[d] <= min[d] {
315 max[d] = min[d] + 1.0;
316 }
317 }
318}
319
320#[derive(Debug, Clone)]
330pub struct Sq8Codec {
331 pub min: Vec<f32>,
333 pub scale: Vec<f32>,
335 pub scale_sq: Vec<f32>,
337 pub scale_sq_f64: Vec<f64>,
339 pub mean_scale_sq: f32,
341 pub scale_sq_residual: Vec<f32>,
343 pub offset_sq_sum: f32,
345 pub offset_sq_sum_f64: f64,
347}
348
349#[derive(Debug, Clone)]
351pub struct EncodedVector {
352 pub codes: Vec<u8>,
354 pub norm: f32,
356 pub soc_sum: f32,
358 pub soc_sum_f64: f64,
360 pub residual_dot_bias: f32,
362}
363
364impl Sq8Codec {
365 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
366 let dims = min.len();
367 let scale: Vec<f32> = (0..dims).map(|d| (max[d] - min[d]) / 255.0).collect();
368 let scale_sq: Vec<f32> = scale.iter().map(|s| s * s).collect();
369 let scale_sq_f64: Vec<f64> = scale.iter().map(|&s| f64::from(s) * f64::from(s)).collect();
370 let mean_scale_sq = scale_sq.iter().sum::<f32>() / dims as f32;
371 let scale_sq_residual: Vec<f32> = scale_sq.iter().map(|&ss| ss - mean_scale_sq).collect();
372 let offset_sq_sum: f32 = min.iter().map(|o| o * o).sum();
373 let offset_sq_sum_f64: f64 = min.iter().map(|&o| f64::from(o) * f64::from(o)).sum();
374
375 Self {
376 min,
377 scale,
378 scale_sq,
379 scale_sq_f64,
380 mean_scale_sq,
381 scale_sq_residual,
382 offset_sq_sum,
383 offset_sq_sum_f64,
384 }
385 }
386
387 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
392 Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
393 }
394
395 pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
398 let (min, max) = flat_min_max(vectors, dims)?;
399 Ok(Self::build_from_min_max(min, max))
400 }
401
402 pub fn train(vectors: &[Vec<f32>]) -> Self {
407 Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
408 }
409
410 pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
415 let (_dims, min, max) = row_min_max(vectors)?;
416 Ok(Self::build_from_min_max(min, max))
417 }
418
419 pub fn encode(&self, v: &[f32]) -> EncodedVector {
425 self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
426 }
427
428 pub fn try_encode(&self, v: &[f32]) -> Result<EncodedVector, QuantError> {
432 let dims = self.min.len();
433 if v.len() != dims {
434 return Err(QuantError::EncodeLengthMismatch {
435 expected: dims,
436 got: v.len(),
437 });
438 }
439 Ok(self.encode_unchecked(v))
440 }
441
442 fn encode_unchecked(&self, v: &[f32]) -> EncodedVector {
444 let dims = self.min.len();
445 let mut codes = Vec::with_capacity(dims);
446 let mut soc_sum = 0.0f32;
447 let mut soc_sum_f64 = 0.0f64;
448 let mut residual_dot_bias = 0.0f32;
449 let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
450
451 for (d, &x) in v.iter().enumerate() {
452 let s = self.scale[d];
453 let inv_s = if s > 1e-12 { 1.0 / s } else { 0.0 };
454 let raw = (x - self.min[d]) * inv_s;
455 let code = raw.round().clamp(0.0, 255.0) as u8;
456 codes.push(code);
457 soc_sum += s * self.min[d] * code as f32;
458 soc_sum_f64 += f64::from(s) * f64::from(self.min[d]) * f64::from(code);
459 residual_dot_bias += self.scale_sq_residual[d] * code as f32;
460 }
461
462 EncodedVector {
463 codes,
464 norm,
465 soc_sum,
466 soc_sum_f64,
467 residual_dot_bias,
468 }
469 }
470
471 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<EncodedVector> {
476 self.try_encode_flat_par(vectors, dims)
477 .unwrap_or_else(|e| panic!("{e}"))
478 }
479
480 pub fn try_encode_flat_par(
484 &self,
485 vectors: &[f32],
486 dims: usize,
487 ) -> Result<Vec<EncodedVector>, QuantError> {
488 if dims == 0 {
489 return Err(QuantError::ZeroDims);
490 }
491 if !vectors.len().is_multiple_of(dims) {
492 return Err(QuantError::FlatLengthNotDivisible {
493 len: vectors.len(),
494 dims,
495 });
496 }
497 if dims != self.min.len() {
498 return Err(QuantError::EncodeLengthMismatch {
499 expected: self.min.len(),
500 got: dims,
501 });
502 }
503 let n = vectors.len() / dims;
504 #[cfg(feature = "parallel")]
505 let encoded = (0..n)
506 .into_par_iter()
507 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
508 .collect();
509 #[cfg(not(feature = "parallel"))]
510 let encoded = (0..n)
511 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
512 .collect();
513 Ok(encoded)
514 }
515
516 pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
522 self.try_encode_par(vectors)
523 .unwrap_or_else(|e| panic!("{e}"))
524 }
525
526 pub fn try_encode_par(&self, vectors: &[Vec<f32>]) -> Result<Vec<EncodedVector>, QuantError> {
529 let dims = self.min.len();
530 for v in vectors {
531 if v.len() != dims {
532 return Err(QuantError::EncodeLengthMismatch {
533 expected: dims,
534 got: v.len(),
535 });
536 }
537 }
538 #[cfg(feature = "parallel")]
539 let encoded = vectors
540 .par_iter()
541 .map(|v| self.encode_unchecked(v))
542 .collect();
543 #[cfg(not(feature = "parallel"))]
544 let encoded = vectors.iter().map(|v| self.encode_unchecked(v)).collect();
545 Ok(encoded)
546 }
547
548 #[inline]
556 pub fn approx_dot(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
557 let dims = self.scale_sq.len();
558 assert_eq!(
559 a.codes.len(),
560 dims,
561 "approx_dot input codes must match codec dims"
562 );
563 assert_eq!(
564 b.codes.len(),
565 dims,
566 "approx_dot input codes must match codec dims"
567 );
568 let dot: f64 = self
569 .scale
570 .iter()
571 .zip(self.min.iter())
572 .zip(a.codes.iter())
573 .zip(b.codes.iter())
574 .map(|(((&scale, &min), &ac), &bc)| {
575 let scale = f64::from(scale);
576 let min = f64::from(min);
577 let a_value = scale.mul_add(f64::from(ac), min);
578 let b_value = scale.mul_add(f64::from(bc), min);
579 a_value * b_value
580 })
581 .sum();
582 dot as f32
583 }
584
585 #[inline]
589 pub fn approx_cosine_dist(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
590 let denom = a.norm * b.norm;
591 if !denom.is_finite() || denom <= 0.0 {
592 return 1.0;
593 }
594 let dot = self.approx_dot(a, b);
595 let cosine = (dot / denom).clamp(-1.0, 1.0);
596 1.0 - cosine
597 }
598
599 #[inline]
612 pub fn approx_l2_sq(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
613 let dims = self.scale_sq.len();
614 assert_eq!(
615 a.codes.len(),
616 dims,
617 "approx_l2_sq input codes must match codec dims"
618 );
619 assert_eq!(
620 b.codes.len(),
621 dims,
622 "approx_l2_sq input codes must match codec dims"
623 );
624 let weighted: f64 = self
625 .scale_sq
626 .iter()
627 .zip(a.codes.iter())
628 .zip(b.codes.iter())
629 .map(|((&weight, &ac), &bc)| {
630 let delta = f64::from(ac) - f64::from(bc);
631 f64::from(weight) * delta * delta
632 })
633 .sum();
634 weighted as f32
635 }
636
637 pub fn dims(&self) -> usize {
639 self.min.len()
640 }
641}
642
643#[derive(Debug, Clone)]
657pub struct GsSq8Codec {
658 pub min: Vec<f32>,
660 pub gs: f32,
662 pub gs_sq: f32,
664 pub anisotropy_ratio: f32,
667}
668
669#[derive(Debug, Clone)]
671pub struct GsEncodedVector {
672 pub codes: Vec<u8>,
674}
675
676impl GsSq8Codec {
677 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
678 let dims = min.len();
679 let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
680 let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
681 let gs = if max_range > 1e-12 {
682 max_range / 255.0
683 } else {
684 1.0 / 255.0
685 };
686
687 let min_range_nonzero = ranges
688 .iter()
689 .cloned()
690 .filter(|&r| r > 1e-12)
691 .fold(f32::INFINITY, f32::min);
692 let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
693 max_range / min_range_nonzero
694 } else {
695 1.0
696 };
697
698 Self {
699 min,
700 gs,
701 gs_sq: gs * gs,
702 anisotropy_ratio,
703 }
704 }
705
706 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
711 Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
712 }
713
714 pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
717 let (min, max) = flat_min_max(vectors, dims)?;
718 Ok(Self::build_from_min_max(min, max))
719 }
720
721 pub fn train(vectors: &[Vec<f32>]) -> Self {
726 Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
727 }
728
729 pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
734 let (_dims, min, max) = row_min_max(vectors)?;
735 Ok(Self::build_from_min_max(min, max))
736 }
737
738 #[inline]
744 pub fn encode(&self, v: &[f32]) -> GsEncodedVector {
745 self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
746 }
747
748 pub fn try_encode(&self, v: &[f32]) -> Result<GsEncodedVector, QuantError> {
753 let dims = self.min.len();
754 if v.len() != dims {
755 return Err(QuantError::EncodeLengthMismatch {
756 expected: dims,
757 got: v.len(),
758 });
759 }
760 Ok(self.encode_unchecked(v))
761 }
762
763 #[inline]
765 fn encode_unchecked(&self, v: &[f32]) -> GsEncodedVector {
766 let inv_gs = if self.gs > 1e-12 { 1.0 / self.gs } else { 0.0 };
767 let codes = v
768 .iter()
769 .enumerate()
770 .map(|(d, &x)| ((x - self.min[d]) * inv_gs).round().clamp(0.0, 255.0) as u8)
771 .collect();
772 GsEncodedVector { codes }
773 }
774
775 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<GsEncodedVector> {
780 self.try_encode_flat_par(vectors, dims)
781 .unwrap_or_else(|e| panic!("{e}"))
782 }
783
784 pub fn try_encode_flat_par(
788 &self,
789 vectors: &[f32],
790 dims: usize,
791 ) -> Result<Vec<GsEncodedVector>, QuantError> {
792 if dims == 0 {
793 return Err(QuantError::ZeroDims);
794 }
795 if !vectors.len().is_multiple_of(dims) {
796 return Err(QuantError::FlatLengthNotDivisible {
797 len: vectors.len(),
798 dims,
799 });
800 }
801 if dims != self.min.len() {
802 return Err(QuantError::EncodeLengthMismatch {
803 expected: self.min.len(),
804 got: dims,
805 });
806 }
807 let n = vectors.len() / dims;
808 #[cfg(feature = "parallel")]
809 let encoded = (0..n)
810 .into_par_iter()
811 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
812 .collect();
813 #[cfg(not(feature = "parallel"))]
814 let encoded = (0..n)
815 .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
816 .collect();
817 Ok(encoded)
818 }
819
820 #[inline]
829 pub fn l2_sq(&self, a: &GsEncodedVector, b: &GsEncodedVector) -> f32 {
830 self.gs_sq * u8_l2sq_u32(&a.codes, &b.codes) as f32
831 }
832
833 #[inline]
836 pub fn l2_sq_codes(&self, a: &[u8], b: &[u8]) -> f32 {
837 self.gs_sq * u8_l2sq_u32(a, b) as f32
838 }
839
840 pub fn dims(&self) -> usize {
842 self.min.len()
843 }
844
845 #[inline]
851 pub fn is_in_distribution(&self, v: &[f32]) -> bool {
852 let max_code = 255.0 * self.gs;
853 v.iter()
854 .zip(self.min.iter())
855 .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
856 }
857}
858
859#[cfg(test)]
860mod tests {
861 use super::*;
862
863 fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
864 let mut h = seed;
865 (0..n)
866 .map(|_| {
867 (0..dims)
868 .map(|_| {
869 h = h
870 .wrapping_mul(0x6c62_272e_07bb_0142)
871 .wrapping_add(0x62b8_2175_62d9_6b1a);
872 let bits = (h >> 33) as u32;
873 (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
874 })
875 .collect()
876 })
877 .collect()
878 }
879
880 fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
881 a.iter().zip(b).map(|(x, y)| x * y).sum()
882 }
883
884 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
885 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
886 }
887
888 #[test]
891 fn sq8_try_train_ragged_rows_returns_error_not_panic() {
892 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
893 let err = Sq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
894 assert_eq!(
895 err,
896 QuantError::RaggedRow {
897 row: 1,
898 expected: 1,
899 got: 2
900 }
901 );
902 }
903
904 #[test]
905 fn gs_try_train_ragged_rows_returns_error_not_panic() {
906 let vecs = vec![vec![0.0], vec![1.0, 2.0]];
907 let err = GsSq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
908 assert_eq!(
909 err,
910 QuantError::RaggedRow {
911 row: 1,
912 expected: 1,
913 got: 2
914 }
915 );
916 }
917
918 #[test]
919 fn sq8_try_train_empty_corpus_returns_error() {
920 let vecs: Vec<Vec<f32>> = vec![];
921 assert_eq!(
922 Sq8Codec::try_train(&vecs).unwrap_err(),
923 QuantError::EmptyCorpus
924 );
925 }
926
927 #[test]
928 fn sq8_try_train_flat_zero_dims_returns_error() {
929 assert_eq!(
930 Sq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
931 QuantError::ZeroDims
932 );
933 }
934
935 #[test]
936 fn gs_try_train_flat_zero_dims_returns_error() {
937 assert_eq!(
938 GsSq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
939 QuantError::ZeroDims
940 );
941 }
942
943 #[test]
944 fn sq8_try_train_flat_remainder_returns_error() {
945 let err = Sq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
947 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
948 }
949
950 #[test]
951 fn gs_try_train_flat_remainder_returns_error() {
952 let err = GsSq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
953 assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
954 }
955
956 #[test]
957 fn sq8_try_encode_shorter_input_returns_error_not_malformed_vector() {
958 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
959 let err = codec.try_encode(&[0.0, 1.0]).unwrap_err();
960 assert_eq!(
961 err,
962 QuantError::EncodeLengthMismatch {
963 expected: 4,
964 got: 2
965 }
966 );
967 }
968
969 #[test]
970 fn sq8_try_encode_longer_input_returns_error() {
971 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
972 let err = codec.try_encode(&[0.0, 1.0, 2.0, 3.0, 4.0]).unwrap_err();
973 assert_eq!(
974 err,
975 QuantError::EncodeLengthMismatch {
976 expected: 4,
977 got: 5
978 }
979 );
980 }
981
982 #[test]
983 fn gs_try_encode_empty_input_returns_error_not_malformed_vector() {
984 let codec = GsSq8Codec::train_flat(&[1.0, 2.0, 3.0, 4.0], 1);
990 let err = codec.try_encode(&[]).unwrap_err();
991 assert_eq!(
992 err,
993 QuantError::EncodeLengthMismatch {
994 expected: 1,
995 got: 0
996 }
997 );
998 }
999
1000 #[test]
1001 fn sq8_try_encode_flat_par_zero_dims_returns_error() {
1002 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1003 assert_eq!(
1004 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
1005 QuantError::ZeroDims
1006 );
1007 }
1008
1009 #[test]
1010 fn gs_try_encode_flat_par_zero_dims_returns_error() {
1011 let codec = GsSq8Codec::train(&rand_vecs(10, 4, 1));
1012 assert_eq!(
1013 codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
1014 QuantError::ZeroDims
1015 );
1016 }
1017
1018 #[test]
1019 fn sq8_try_encode_flat_par_remainder_returns_error() {
1020 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1021 let flat: Vec<f32> = (0..9).map(|i| i as f32).collect(); assert_eq!(
1023 codec.try_encode_flat_par(&flat, 4).unwrap_err(),
1024 QuantError::FlatLengthNotDivisible { len: 9, dims: 4 }
1025 );
1026 }
1027
1028 #[test]
1029 fn sq8_try_encode_flat_par_dims_mismatch_returns_error() {
1030 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1031 let flat: Vec<f32> = (0..6).map(|i| i as f32).collect();
1032 let err = codec.try_encode_flat_par(&flat, 3).unwrap_err();
1033 assert_eq!(
1034 err,
1035 QuantError::EncodeLengthMismatch {
1036 expected: 4,
1037 got: 3
1038 }
1039 );
1040 }
1041
1042 #[test]
1043 #[should_panic(expected = "row 1 has length 2, expected 1")]
1044 fn sq8_train_still_panics_with_typed_message_on_ragged_rows() {
1045 let _ = Sq8Codec::train(&[vec![0.0], vec![1.0, 2.0]]);
1049 }
1050
1051 #[test]
1052 fn sq8_try_train_and_encode_roundtrip_matches_panicking_api() {
1053 let vecs = rand_vecs(20, 8, 7);
1054 let a = Sq8Codec::train(&vecs);
1055 let b = Sq8Codec::try_train(&vecs).expect("valid corpus must train");
1056 assert_eq!(a.min, b.min);
1057 assert_eq!(a.scale, b.scale);
1058 let ea = a.encode(&vecs[0]);
1059 let eb = b.try_encode(&vecs[0]).expect("valid vector must encode");
1060 assert_eq!(ea.codes, eb.codes);
1061 }
1062
1063 #[test]
1066 fn encode_decode_roundtrip_is_bounded() {
1067 let vecs = rand_vecs(100, 32, 42);
1068 let codec = Sq8Codec::train(&vecs);
1069 for v in &vecs {
1070 let ev = codec.encode(v);
1071 assert_eq!(ev.codes.len(), v.len());
1072 for (d, &code) in ev.codes.iter().enumerate() {
1073 let decoded = code as f32 * codec.scale[d] + codec.min[d];
1074 let err = (decoded - v[d]).abs();
1075 assert!(
1076 err <= codec.scale[d] + 1e-5,
1077 "dim {d}: err={err} scale={}",
1078 codec.scale[d]
1079 );
1080 }
1081 }
1082 }
1083
1084 #[test]
1085 fn approx_dot_relative_error_bounded() {
1086 let vecs = rand_vecs(200, 64, 77);
1087 let codec = Sq8Codec::train(&vecs);
1088 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1089
1090 let mut max_rel_err = 0.0f32;
1091 for i in 0..vecs.len() {
1092 for j in (i + 1)..vecs.len().min(i + 10) {
1093 let true_dot = dot_f32(&vecs[i], &vecs[j]);
1094 let approx = codec.approx_dot(&encoded[i], &encoded[j]);
1095 let denom = true_dot.abs().max(1e-3);
1096 let rel = (approx - true_dot).abs() / denom;
1097 if rel > max_rel_err {
1098 max_rel_err = rel;
1099 }
1100 }
1101 }
1102 assert!(
1103 max_rel_err < 0.15,
1104 "max relative dot error {max_rel_err:.4} >= 0.15"
1105 );
1106 }
1107
1108 #[test]
1109 fn approx_l2_sq_relative_error_bounded() {
1110 let vecs = rand_vecs(200, 64, 88);
1111 let codec = Sq8Codec::train(&vecs);
1112 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1113
1114 let mut max_rel_err = 0.0f32;
1115 for i in 0..vecs.len() {
1116 for j in (i + 1)..vecs.len().min(i + 10) {
1117 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1118 let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
1119 let denom = true_l2.max(1e-6);
1120 let rel = (approx - true_l2).abs() / denom;
1121 if rel > max_rel_err {
1122 max_rel_err = rel;
1123 }
1124 }
1125 }
1126 assert!(
1127 max_rel_err < 0.15,
1128 "max relative L2² error {max_rel_err:.4} >= 0.15"
1129 );
1130 }
1131
1132 #[test]
1133 fn narrow_sq8_dimensions_survive_a_wide_training_dimension() {
1134 for dims in [2, 17, 384, 768, 1536, 3072] {
1135 let zero = vec![0.0_f32; dims];
1136 let mut upper = vec![255.0_f32 / 16384.0; dims];
1137 upper[0] = 255.0;
1138 let mut narrow = upper.clone();
1139 narrow[0] = 0.0;
1140 let codec = Sq8Codec::train(&[zero.clone(), upper]);
1141 let q = codec.encode(&zero);
1142 let b = codec.encode(&narrow);
1143 assert!(q.codes.iter().all(|&code| code == 0));
1144 assert_eq!(b.codes[0], 0);
1145 assert!(b.codes[1..].iter().all(|&code| code == 255));
1146
1147 let expected = ((dims - 1) as f64 * 65025.0 / 268435456.0) as f32;
1148 let distance = codec.approx_l2_sq(&q, &b);
1149 assert_eq!(distance, expected, "dims={dims}");
1150 assert!(distance > codec.approx_l2_sq(&q, &q), "dims={dims}");
1151 assert_eq!(codec.approx_dot(&b, &b), expected, "dims={dims}");
1152 if dims == 2 {
1153 assert!(codec.approx_cosine_dist(&b, &b) < 1e-6);
1154 }
1155 }
1156 }
1157
1158 #[test]
1159 fn order_preservation_triplets_cosine() {
1160 let vecs = rand_vecs(300, 64, 99);
1161 let codec = Sq8Codec::train(&vecs);
1162 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1163
1164 let n = vecs.len();
1165 let mut agree = 0usize;
1166 let mut total = 0usize;
1167
1168 for anchor in 0..50 {
1169 let a = &vecs[anchor];
1170 let ea = &encoded[anchor];
1171 for b_idx in 0..n {
1172 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1173 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1174 let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1175 let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1176
1177 let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
1178 let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
1179 let dist_ab_true = 1.0 - cos_ab;
1180 let dist_ac_true = 1.0 - cos_ac;
1181
1182 let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
1183 let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
1184
1185 if (dist_ab_true - dist_ac_true).abs() < 0.01 {
1186 continue;
1187 }
1188
1189 let true_closer_b = dist_ab_true < dist_ac_true;
1190 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1191 if true_closer_b == approx_closer_b {
1192 agree += 1;
1193 }
1194 total += 1;
1195 }
1196 }
1197 }
1198
1199 let rate = agree as f64 / total.max(1) as f64;
1200 assert!(
1201 rate >= 0.95,
1202 "order preservation {rate:.3} < 0.95 ({agree}/{total})"
1203 );
1204 }
1205
1206 #[test]
1207 fn order_preservation_triplets_l2() {
1208 let vecs = rand_vecs(300, 64, 101);
1209 let codec = Sq8Codec::train(&vecs);
1210 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1211
1212 let n = vecs.len();
1213 let mut agree = 0usize;
1214 let mut total = 0usize;
1215
1216 for anchor in 0..50 {
1217 let a = &vecs[anchor];
1218 let ea = &encoded[anchor];
1219 for b_idx in 0..n {
1220 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1221 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1222 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1223
1224 let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
1225 let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
1226
1227 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1228 continue;
1229 }
1230
1231 let true_closer_b = dist_ab_true < dist_ac_true;
1232 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1233 if true_closer_b == approx_closer_b {
1234 agree += 1;
1235 }
1236 total += 1;
1237 }
1238 }
1239 }
1240
1241 let rate = agree as f64 / total.max(1) as f64;
1242 assert!(
1243 rate >= 0.95,
1244 "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1245 );
1246 }
1247
1248 #[test]
1249 fn train_flat_matches_train_rows() {
1250 let vecs = rand_vecs(50, 16, 123);
1251 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1252
1253 let codec_rows = Sq8Codec::train(&vecs);
1254 let codec_flat = Sq8Codec::train_flat(&flat, 16);
1255
1256 for d in 0..16 {
1257 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1258 assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
1259 }
1260 }
1261
1262 #[test]
1263 fn encode_par_matches_sequential() {
1264 let vecs = rand_vecs(50, 32, 555);
1265 let codec = Sq8Codec::train(&vecs);
1266
1267 let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1268 let par = codec.encode_par(&vecs);
1269
1270 assert_eq!(seq.len(), par.len());
1271 for (s, p) in seq.iter().zip(par.iter()) {
1272 assert_eq!(s.codes, p.codes);
1273 assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
1274 }
1275 }
1276
1277 #[test]
1278 fn sq8_try_encode_par_short_row_returns_error_not_panic() {
1279 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1280 let mut rows = rand_vecs(5, 4, 2);
1281 rows[3] = vec![0.0, 1.0];
1282 let err = codec.try_encode_par(&rows).unwrap_err();
1283 assert_eq!(
1284 err,
1285 QuantError::EncodeLengthMismatch {
1286 expected: 4,
1287 got: 2
1288 }
1289 );
1290 }
1291
1292 #[test]
1293 fn sq8_try_encode_par_long_row_returns_error_not_panic() {
1294 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1295 let mut rows = rand_vecs(5, 4, 2);
1296 rows[3] = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
1297 let err = codec.try_encode_par(&rows).unwrap_err();
1298 assert_eq!(
1299 err,
1300 QuantError::EncodeLengthMismatch {
1301 expected: 4,
1302 got: 6
1303 }
1304 );
1305 }
1306
1307 #[test]
1308 #[should_panic(expected = "vector length 2 does not match codec dims 4")]
1309 fn sq8_encode_par_still_panics_with_typed_message_on_short_row() {
1310 let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1311 let mut rows = rand_vecs(5, 4, 2);
1312 rows[3] = vec![0.0, 1.0];
1313 let _ = codec.encode_par(&rows);
1314 }
1315
1316 #[test]
1317 fn u8_dot_u32_matches_scalar() {
1318 let a: Vec<u8> = (0u8..=255).take(384).collect();
1319 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1320 let scalar: u32 = a
1321 .iter()
1322 .zip(b.iter())
1323 .map(|(&x, &y)| x as u32 * y as u32)
1324 .sum();
1325 assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
1326 }
1327
1328 #[test]
1329 #[should_panic(expected = "u8_dot_u32 inputs must have equal length")]
1330 fn u8_dot_u32_rejects_shorter_second_slice() {
1331 let a = [1u8; 16];
1332 let b = [2u8; 1];
1333
1334 let _ = u8_dot_u32(&a, &b);
1335 }
1336
1337 #[test]
1338 #[should_panic(expected = "u8_dot_u32 inputs must have equal length")]
1339 fn u8_dot_u32_rejects_longer_second_slice() {
1340 let a = [1u8; 1];
1341 let b = [2u8; 16];
1342
1343 let _ = u8_dot_u32(&a, &b);
1344 }
1345
1346 #[test]
1347 #[should_panic(expected = "approx_dot input codes must match codec dims")]
1348 fn approx_dot_rejects_mismatched_codes() {
1349 let vectors = rand_vecs(2, 16, 1);
1350 let codec = Sq8Codec::train(&vectors);
1351 let a = codec.encode(&vectors[0]);
1352 let mut b = codec.encode(&vectors[1]);
1353 b.codes.pop();
1354
1355 let _ = codec.approx_dot(&a, &b);
1356 }
1357
1358 #[test]
1359 #[should_panic(expected = "approx_dot input codes must match codec dims")]
1360 fn approx_dot_rejects_longer_first_codes() {
1361 let vectors = rand_vecs(2, 16, 1);
1362 let codec = Sq8Codec::train(&vectors);
1363 let mut a = codec.encode(&vectors[0]);
1364 let b = codec.encode(&vectors[1]);
1365 a.codes.push(0);
1366
1367 let _ = codec.approx_dot(&a, &b);
1368 }
1369
1370 #[test]
1371 #[should_panic(expected = "approx_l2_sq input codes must match codec dims")]
1372 fn approx_l2_sq_rejects_mismatched_codes() {
1373 let vectors = rand_vecs(2, 16, 1);
1374 let codec = Sq8Codec::train(&vectors);
1375 let a = codec.encode(&vectors[0]);
1376 let mut b = codec.encode(&vectors[1]);
1377 b.codes.pop();
1378
1379 let _ = codec.approx_l2_sq(&a, &b);
1380 }
1381
1382 #[test]
1383 #[should_panic(expected = "approx_l2_sq input codes must match codec dims")]
1384 fn approx_l2_sq_rejects_longer_first_codes() {
1385 let vectors = rand_vecs(2, 16, 1);
1386 let codec = Sq8Codec::train(&vectors);
1387 let mut a = codec.encode(&vectors[0]);
1388 let b = codec.encode(&vectors[1]);
1389 a.codes.push(0);
1390
1391 let _ = codec.approx_l2_sq(&a, &b);
1392 }
1393
1394 #[test]
1395 fn u8_helpers_tail_path_max_diff() {
1396 for len in [1usize, 7, 15, 17, 100, 383] {
1397 let a = vec![255u8; len];
1398 let b = vec![0u8; len];
1399 assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
1400 assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
1401 }
1402 }
1403
1404 #[test]
1405 fn u8_l2sq_u32_matches_scalar() {
1406 let a: Vec<u8> = (0u8..=255).take(384).collect();
1407 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1408 let scalar: u32 = a
1409 .iter()
1410 .zip(b.iter())
1411 .map(|(&x, &y)| {
1412 let d = (x as i32) - (y as i32);
1413 (d * d) as u32
1414 })
1415 .sum();
1416 assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
1417 }
1418
1419 #[test]
1420 #[should_panic(expected = "u8_l2sq_u32 inputs must have equal length")]
1421 fn u8_l2sq_u32_rejects_shorter_second_slice() {
1422 let a = [1u8; 16];
1423 let b = [2u8; 1];
1424
1425 let _ = u8_l2sq_u32(&a, &b);
1426 }
1427
1428 #[test]
1436 fn gs_l2_sq_anisotropic_ordering_preserved() {
1437 let corpus = vec![
1438 vec![0.0f32, 0.0f32], vec![1.0f32, 1.0f32], vec![1.0f32, 4001.0f32], ];
1442 let codec = GsSq8Codec::train(&corpus);
1443
1444 let enc_origin = codec.encode(&corpus[0]);
1445 let enc_near = codec.encode(&corpus[1]);
1446 let enc_far = codec.encode(&corpus[2]);
1447
1448 let d_near = codec.l2_sq(&enc_origin, &enc_near);
1449 let d_far = codec.l2_sq(&enc_origin, &enc_far);
1450
1451 assert!(
1452 d_near < d_far,
1453 "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
1454 (anisotropy_ratio={:.1})",
1455 codec.anisotropy_ratio
1456 );
1457 }
1458
1459 #[test]
1460 fn gs_l2_sq_isotropic_small_error() {
1461 let vecs = rand_vecs(200, 64, 202);
1462 let codec = GsSq8Codec::train(&vecs);
1463 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1464
1465 let mut max_rel = 0.0f32;
1466 for i in 0..vecs.len() {
1467 for j in (i + 1)..vecs.len().min(i + 10) {
1468 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1469 let approx = codec.l2_sq(&encoded[i], &encoded[j]);
1470 let denom = true_l2.max(1e-6);
1471 let rel = (approx - true_l2).abs() / denom;
1472 if rel > max_rel {
1473 max_rel = rel;
1474 }
1475 }
1476 }
1477 assert!(
1478 max_rel < 0.15,
1479 "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
1480 );
1481 }
1482
1483 #[test]
1484 fn gs_train_flat_matches_train_rows() {
1485 let vecs = rand_vecs(50, 16, 321);
1486 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1487
1488 let codec_rows = GsSq8Codec::train(&vecs);
1489 let codec_flat = GsSq8Codec::train_flat(&flat, 16);
1490
1491 assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
1492 for d in 0..16 {
1493 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1494 }
1495 }
1496
1497 #[test]
1498 fn gs_l2_sq_order_preservation_triplets() {
1499 let vecs = rand_vecs(300, 64, 303);
1500 let codec = GsSq8Codec::train(&vecs);
1501 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1502
1503 let n = vecs.len();
1504 let mut agree = 0usize;
1505 let mut total = 0usize;
1506
1507 for anchor in 0..50 {
1508 let a = &vecs[anchor];
1509 let ea = &encoded[anchor];
1510 for b_idx in 0..n {
1511 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1512 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1513 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1514
1515 let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
1516 let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
1517
1518 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1519 continue;
1520 }
1521
1522 let true_closer_b = dist_ab_true < dist_ac_true;
1523 let approx_closer_b = dist_ab_approx < dist_ac_approx;
1524 if true_closer_b == approx_closer_b {
1525 agree += 1;
1526 }
1527 total += 1;
1528 }
1529 }
1530 }
1531
1532 let rate = agree as f64 / total.max(1) as f64;
1533 assert!(
1534 rate >= 0.95,
1535 "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1536 );
1537 }
1538}