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 #[cfg(target_arch = "aarch64")]
106 {
107 use std::arch::aarch64::*;
108 let n = a.len();
109 let chunks = n / 16;
110 let rem = n % 16;
111
112 let mut acc0: uint32x4_t;
113 let mut acc1: uint32x4_t;
114 let mut acc2: uint32x4_t;
115 let mut acc3: uint32x4_t;
116
117 unsafe {
118 acc0 = vdupq_n_u32(0);
119 acc1 = vdupq_n_u32(0);
120 acc2 = vdupq_n_u32(0);
121 acc3 = vdupq_n_u32(0);
122
123 for i in 0..chunks {
124 let ap = a.as_ptr().add(i * 16);
125 let bp = b.as_ptr().add(i * 16);
126
127 let va = vld1q_u8(ap);
128 let vb = vld1q_u8(bp);
129
130 let diff = vabdq_u8(va, vb);
131
132 let lo_u16 = vmull_u8(vget_low_u8(diff), vget_low_u8(diff));
133 let hi_u16 = vmull_high_u8(diff, diff);
134
135 acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
136 acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
137 acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
138 acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
139 }
140
141 let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
142 let mut total = vaddvq_u32(sum4);
143
144 for i in (n - rem)..n {
145 let d = (a[i] as i32) - (b[i] as i32);
146 total += (d * d) as u32;
147 }
148 total
149 }
150 }
151
152 #[cfg(not(target_arch = "aarch64"))]
153 {
154 a.chunks(8)
155 .zip(b.chunks(8))
156 .map(|(ac, bc)| {
157 ac.iter()
158 .zip(bc.iter())
159 .map(|(&x, &y)| {
160 let d = (x as i32) - (y as i32);
161 (d * d) as u32
162 })
163 .sum::<u32>()
164 })
165 .sum()
166 }
167}
168
169#[derive(Debug, Clone)]
179pub struct Sq8Codec {
180 pub min: Vec<f32>,
182 pub scale: Vec<f32>,
184 pub scale_sq: Vec<f32>,
186 pub mean_scale_sq: f32,
188 pub scale_sq_residual: Vec<f32>,
190 pub offset_sq_sum: f32,
192}
193
194#[derive(Debug, Clone)]
196pub struct EncodedVector {
197 pub codes: Vec<u8>,
199 pub norm: f32,
201 pub soc_sum: f32,
203 pub residual_dot_bias: f32,
205}
206
207impl Sq8Codec {
208 fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
209 let dims = min.len();
210 let scale: Vec<f32> = (0..dims).map(|d| (max[d] - min[d]) / 255.0).collect();
211 let scale_sq: Vec<f32> = scale.iter().map(|s| s * s).collect();
212 let mean_scale_sq = scale_sq.iter().sum::<f32>() / dims as f32;
213 let scale_sq_residual: Vec<f32> = scale_sq.iter().map(|&ss| ss - mean_scale_sq).collect();
214 let offset_sq_sum: f32 = min.iter().map(|o| o * o).sum();
215
216 Self {
217 min,
218 scale,
219 scale_sq,
220 mean_scale_sq,
221 scale_sq_residual,
222 offset_sq_sum,
223 }
224 }
225
226 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
228 assert!(dims > 0, "dims must be > 0");
229 assert!(!vectors.is_empty(), "cannot train on empty corpus");
230 assert_eq!(
231 vectors.len() % dims,
232 0,
233 "vectors length must be a multiple of dims"
234 );
235
236 let n = vectors.len() / dims;
237 let mut min = vec![f32::INFINITY; dims];
238 let mut max = vec![f32::NEG_INFINITY; dims];
239
240 for row in 0..n {
241 let v = &vectors[row * dims..(row + 1) * dims];
242 for (d, &x) in v.iter().enumerate() {
243 if x.is_finite() {
244 if x < min[d] {
245 min[d] = x;
246 }
247 if x > max[d] {
248 max[d] = x;
249 }
250 }
251 }
252 }
253
254 for d in 0..dims {
255 if !min[d].is_finite() {
256 min[d] = 0.0;
257 }
258 if !max[d].is_finite() || max[d] <= min[d] {
259 max[d] = min[d] + 1.0;
260 }
261 }
262
263 Self::build_from_min_max(min, max)
264 }
265
266 pub fn train(vectors: &[Vec<f32>]) -> Self {
268 assert!(!vectors.is_empty(), "cannot train on empty corpus");
269 let dims = vectors[0].len();
270 assert!(dims > 0, "dims must be > 0");
271
272 let mut min = vec![f32::INFINITY; dims];
273 let mut max = vec![f32::NEG_INFINITY; dims];
274
275 for v in vectors {
276 for (d, &x) in v.iter().enumerate() {
277 if x.is_finite() {
278 if x < min[d] {
279 min[d] = x;
280 }
281 if x > max[d] {
282 max[d] = x;
283 }
284 }
285 }
286 }
287
288 for d in 0..dims {
289 if !min[d].is_finite() {
290 min[d] = 0.0;
291 }
292 if !max[d].is_finite() || max[d] <= min[d] {
293 max[d] = min[d] + 1.0;
294 }
295 }
296
297 Self::build_from_min_max(min, max)
298 }
299
300 pub fn encode(&self, v: &[f32]) -> EncodedVector {
302 let dims = self.min.len();
303 debug_assert_eq!(v.len(), dims, "vector length must match codec dims");
304
305 let mut codes = Vec::with_capacity(dims);
306 let mut soc_sum = 0.0f32;
307 let mut residual_dot_bias = 0.0f32;
308 let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
309
310 for (d, &x) in v.iter().enumerate() {
311 let s = self.scale[d];
312 let inv_s = if s > 1e-12 { 1.0 / s } else { 0.0 };
313 let raw = (x - self.min[d]) * inv_s;
314 let code = raw.round().clamp(0.0, 255.0) as u8;
315 codes.push(code);
316 soc_sum += s * self.min[d] * code as f32;
317 residual_dot_bias += self.scale_sq_residual[d] * code as f32;
318 }
319
320 EncodedVector {
321 codes,
322 norm,
323 soc_sum,
324 residual_dot_bias,
325 }
326 }
327
328 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<EncodedVector> {
330 let n = vectors.len() / dims;
331 (0..n)
332 .into_par_iter()
333 .map(|i| self.encode(&vectors[i * dims..(i + 1) * dims]))
334 .collect()
335 }
336
337 pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
339 vectors.par_iter().map(|v| self.encode(v)).collect()
340 }
341
342 #[inline]
351 pub fn approx_dot(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
352 let raw = u8_dot_u32(&a.codes, &b.codes) as f32;
353 let residual_hot: f32 = self
354 .scale_sq_residual
355 .iter()
356 .zip(a.codes.iter())
357 .zip(b.codes.iter())
358 .map(|((r, &ac), &bc)| r * (ac as f32) * (bc as f32))
359 .sum();
360 self.mean_scale_sq * raw + residual_hot + a.soc_sum + b.soc_sum + self.offset_sq_sum
361 }
362
363 #[inline]
367 pub fn approx_cosine_dist(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
368 let denom = a.norm * b.norm;
369 if !denom.is_finite() || denom <= 0.0 {
370 return 1.0;
371 }
372 let dot = self.approx_dot(a, b);
373 let cosine = (dot / denom).clamp(-1.0, 1.0);
374 1.0 - cosine
375 }
376
377 #[inline]
389 pub fn approx_l2_sq(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
390 let raw = u8_l2sq_u32(&a.codes, &b.codes) as f32;
391 let residual_hot: f32 = self
392 .scale_sq_residual
393 .iter()
394 .zip(a.codes.iter())
395 .zip(b.codes.iter())
396 .map(|((r, &ac), &bc)| {
397 let d = (ac as i32) - (bc as i32);
398 r * (d as f32) * (d as f32)
399 })
400 .sum();
401 self.mean_scale_sq * raw + residual_hot
402 }
403
404 pub fn dims(&self) -> usize {
406 self.min.len()
407 }
408}
409
410#[derive(Debug, Clone)]
432pub struct GsSq8Codec {
433 pub min: Vec<f32>,
435 pub gs: f32,
437 pub gs_sq: f32,
439 pub anisotropy_ratio: f32,
442}
443
444#[derive(Debug, Clone)]
446pub struct GsEncodedVector {
447 pub codes: Vec<u8>,
449}
450
451impl GsSq8Codec {
452 pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
454 assert!(dims > 0, "dims must be > 0");
455 assert!(!vectors.is_empty(), "cannot train on empty corpus");
456 assert_eq!(
457 vectors.len() % dims,
458 0,
459 "vectors length must be a multiple of dims"
460 );
461
462 let n = vectors.len() / dims;
463 let mut min = vec![f32::INFINITY; dims];
464 let mut max = vec![f32::NEG_INFINITY; dims];
465
466 for row in 0..n {
467 let v = &vectors[row * dims..(row + 1) * dims];
468 for (d, &x) in v.iter().enumerate() {
469 if x.is_finite() {
470 if x < min[d] {
471 min[d] = x;
472 }
473 if x > max[d] {
474 max[d] = x;
475 }
476 }
477 }
478 }
479
480 for d in 0..dims {
481 if !min[d].is_finite() {
482 min[d] = 0.0;
483 }
484 if !max[d].is_finite() || max[d] <= min[d] {
485 max[d] = min[d] + 1.0;
486 }
487 }
488
489 let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
490 let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
491 let gs = if max_range > 1e-12 {
492 max_range / 255.0
493 } else {
494 1.0 / 255.0
495 };
496
497 let min_range_nonzero = ranges
498 .iter()
499 .cloned()
500 .filter(|&r| r > 1e-12)
501 .fold(f32::INFINITY, f32::min);
502 let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
503 max_range / min_range_nonzero
504 } else {
505 1.0
506 };
507
508 Self {
509 min,
510 gs,
511 gs_sq: gs * gs,
512 anisotropy_ratio,
513 }
514 }
515
516 pub fn train(vectors: &[Vec<f32>]) -> Self {
518 assert!(!vectors.is_empty(), "cannot train on empty corpus");
519 let dims = vectors[0].len();
520 assert!(dims > 0, "dims must be > 0");
521
522 let mut min = vec![f32::INFINITY; dims];
523 let mut max = vec![f32::NEG_INFINITY; dims];
524
525 for v in vectors {
526 for (d, &x) in v.iter().enumerate() {
527 if x.is_finite() {
528 if x < min[d] {
529 min[d] = x;
530 }
531 if x > max[d] {
532 max[d] = x;
533 }
534 }
535 }
536 }
537
538 for d in 0..dims {
539 if !min[d].is_finite() {
540 min[d] = 0.0;
541 }
542 if !max[d].is_finite() || max[d] <= min[d] {
543 max[d] = min[d] + 1.0;
544 }
545 }
546
547 let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
548 let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
549 let gs = if max_range > 1e-12 {
550 max_range / 255.0
551 } else {
552 1.0 / 255.0
553 };
554
555 let min_range_nonzero = ranges
556 .iter()
557 .cloned()
558 .filter(|&r| r > 1e-12)
559 .fold(f32::INFINITY, f32::min);
560 let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
561 max_range / min_range_nonzero
562 } else {
563 1.0
564 };
565
566 Self {
567 min,
568 gs,
569 gs_sq: gs * gs,
570 anisotropy_ratio,
571 }
572 }
573
574 #[inline]
576 pub fn encode(&self, v: &[f32]) -> GsEncodedVector {
577 debug_assert_eq!(
578 v.len(),
579 self.min.len(),
580 "vector length must match codec dims"
581 );
582 let inv_gs = if self.gs > 1e-12 { 1.0 / self.gs } else { 0.0 };
583 let codes = v
584 .iter()
585 .enumerate()
586 .map(|(d, &x)| ((x - self.min[d]) * inv_gs).round().clamp(0.0, 255.0) as u8)
587 .collect();
588 GsEncodedVector { codes }
589 }
590
591 pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<GsEncodedVector> {
593 let n = vectors.len() / dims;
594 (0..n)
595 .into_par_iter()
596 .map(|i| self.encode(&vectors[i * dims..(i + 1) * dims]))
597 .collect()
598 }
599
600 #[inline]
609 pub fn l2_sq(&self, a: &GsEncodedVector, b: &GsEncodedVector) -> f32 {
610 self.gs_sq * u8_l2sq_u32(&a.codes, &b.codes) as f32
611 }
612
613 pub fn dims(&self) -> usize {
615 self.min.len()
616 }
617
618 #[inline]
624 pub fn is_in_distribution(&self, v: &[f32]) -> bool {
625 let max_code = 255.0 * self.gs;
626 v.iter()
627 .zip(self.min.iter())
628 .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
629 }
630}
631
632#[cfg(test)]
633mod tests {
634 use super::*;
635
636 fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
637 let mut h = seed;
638 (0..n)
639 .map(|_| {
640 (0..dims)
641 .map(|_| {
642 h = h
643 .wrapping_mul(0x6c62_272e_07bb_0142)
644 .wrapping_add(0x62b8_2175_62d9_6b1a);
645 let bits = (h >> 33) as u32;
646 (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
647 })
648 .collect()
649 })
650 .collect()
651 }
652
653 fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
654 a.iter().zip(b).map(|(x, y)| x * y).sum()
655 }
656
657 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
658 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
659 }
660
661 #[test]
664 fn encode_decode_roundtrip_is_bounded() {
665 let vecs = rand_vecs(100, 32, 42);
666 let codec = Sq8Codec::train(&vecs);
667 for v in &vecs {
668 let ev = codec.encode(v);
669 assert_eq!(ev.codes.len(), v.len());
670 for (d, &code) in ev.codes.iter().enumerate() {
671 let decoded = code as f32 * codec.scale[d] + codec.min[d];
672 let err = (decoded - v[d]).abs();
673 assert!(
674 err <= codec.scale[d] + 1e-5,
675 "dim {d}: err={err} scale={}",
676 codec.scale[d]
677 );
678 }
679 }
680 }
681
682 #[test]
683 fn approx_dot_relative_error_bounded() {
684 let vecs = rand_vecs(200, 64, 77);
685 let codec = Sq8Codec::train(&vecs);
686 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
687
688 let mut max_rel_err = 0.0f32;
689 for i in 0..vecs.len() {
690 for j in (i + 1)..vecs.len().min(i + 10) {
691 let true_dot = dot_f32(&vecs[i], &vecs[j]);
692 let approx = codec.approx_dot(&encoded[i], &encoded[j]);
693 let denom = true_dot.abs().max(1e-3);
694 let rel = (approx - true_dot).abs() / denom;
695 if rel > max_rel_err {
696 max_rel_err = rel;
697 }
698 }
699 }
700 assert!(
701 max_rel_err < 0.15,
702 "max relative dot error {max_rel_err:.4} >= 0.15"
703 );
704 }
705
706 #[test]
707 fn approx_l2_sq_relative_error_bounded() {
708 let vecs = rand_vecs(200, 64, 88);
709 let codec = Sq8Codec::train(&vecs);
710 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
711
712 let mut max_rel_err = 0.0f32;
713 for i in 0..vecs.len() {
714 for j in (i + 1)..vecs.len().min(i + 10) {
715 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
716 let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
717 let denom = true_l2.max(1e-6);
718 let rel = (approx - true_l2).abs() / denom;
719 if rel > max_rel_err {
720 max_rel_err = rel;
721 }
722 }
723 }
724 assert!(
725 max_rel_err < 0.15,
726 "max relative L2² error {max_rel_err:.4} >= 0.15"
727 );
728 }
729
730 #[test]
731 fn order_preservation_triplets_cosine() {
732 let vecs = rand_vecs(300, 64, 99);
733 let codec = Sq8Codec::train(&vecs);
734 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
735
736 let n = vecs.len();
737 let mut agree = 0usize;
738 let mut total = 0usize;
739
740 for anchor in 0..50 {
741 let a = &vecs[anchor];
742 let ea = &encoded[anchor];
743 for b_idx in 0..n {
744 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
745 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
746 let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
747 let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
748
749 let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
750 let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
751 let dist_ab_true = 1.0 - cos_ab;
752 let dist_ac_true = 1.0 - cos_ac;
753
754 let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
755 let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
756
757 if (dist_ab_true - dist_ac_true).abs() < 0.01 {
758 continue;
759 }
760
761 let true_closer_b = dist_ab_true < dist_ac_true;
762 let approx_closer_b = dist_ab_approx < dist_ac_approx;
763 if true_closer_b == approx_closer_b {
764 agree += 1;
765 }
766 total += 1;
767 }
768 }
769 }
770
771 let rate = agree as f64 / total.max(1) as f64;
772 assert!(
773 rate >= 0.95,
774 "order preservation {rate:.3} < 0.95 ({agree}/{total})"
775 );
776 }
777
778 #[test]
779 fn order_preservation_triplets_l2() {
780 let vecs = rand_vecs(300, 64, 101);
781 let codec = Sq8Codec::train(&vecs);
782 let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
783
784 let n = vecs.len();
785 let mut agree = 0usize;
786 let mut total = 0usize;
787
788 for anchor in 0..50 {
789 let a = &vecs[anchor];
790 let ea = &encoded[anchor];
791 for b_idx in 0..n {
792 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
793 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
794 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
795
796 let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
797 let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
798
799 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
800 continue;
801 }
802
803 let true_closer_b = dist_ab_true < dist_ac_true;
804 let approx_closer_b = dist_ab_approx < dist_ac_approx;
805 if true_closer_b == approx_closer_b {
806 agree += 1;
807 }
808 total += 1;
809 }
810 }
811 }
812
813 let rate = agree as f64 / total.max(1) as f64;
814 assert!(
815 rate >= 0.95,
816 "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
817 );
818 }
819
820 #[test]
821 fn train_flat_matches_train_rows() {
822 let vecs = rand_vecs(50, 16, 123);
823 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
824
825 let codec_rows = Sq8Codec::train(&vecs);
826 let codec_flat = Sq8Codec::train_flat(&flat, 16);
827
828 for d in 0..16 {
829 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
830 assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
831 }
832 }
833
834 #[test]
835 fn encode_par_matches_sequential() {
836 let vecs = rand_vecs(50, 32, 555);
837 let codec = Sq8Codec::train(&vecs);
838
839 let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
840 let par = codec.encode_par(&vecs);
841
842 assert_eq!(seq.len(), par.len());
843 for (s, p) in seq.iter().zip(par.iter()) {
844 assert_eq!(s.codes, p.codes);
845 assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
846 }
847 }
848
849 #[test]
850 fn u8_dot_u32_matches_scalar() {
851 let a: Vec<u8> = (0u8..=255).take(384).collect();
852 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
853 let scalar: u32 = a
854 .iter()
855 .zip(b.iter())
856 .map(|(&x, &y)| x as u32 * y as u32)
857 .sum();
858 assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
859 }
860
861 #[test]
862 fn u8_helpers_tail_path_max_diff() {
863 for len in [1usize, 7, 15, 17, 100, 383] {
864 let a = vec![255u8; len];
865 let b = vec![0u8; len];
866 assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
867 assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
868 }
869 }
870
871 #[test]
872 fn u8_l2sq_u32_matches_scalar() {
873 let a: Vec<u8> = (0u8..=255).take(384).collect();
874 let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
875 let scalar: u32 = a
876 .iter()
877 .zip(b.iter())
878 .map(|(&x, &y)| {
879 let d = (x as i32) - (y as i32);
880 (d * d) as u32
881 })
882 .sum();
883 assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
884 }
885
886 #[test]
894 fn gs_l2_sq_anisotropic_ordering_preserved() {
895 let corpus = vec![
896 vec![0.0f32, 0.0f32], vec![1.0f32, 1.0f32], vec![1.0f32, 4001.0f32], ];
900 let codec = GsSq8Codec::train(&corpus);
901
902 let enc_origin = codec.encode(&corpus[0]);
903 let enc_near = codec.encode(&corpus[1]);
904 let enc_far = codec.encode(&corpus[2]);
905
906 let d_near = codec.l2_sq(&enc_origin, &enc_near);
907 let d_far = codec.l2_sq(&enc_origin, &enc_far);
908
909 assert!(
910 d_near < d_far,
911 "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
912 (anisotropy_ratio={:.1})",
913 codec.anisotropy_ratio
914 );
915 }
916
917 #[test]
918 fn gs_l2_sq_isotropic_small_error() {
919 let vecs = rand_vecs(200, 64, 202);
920 let codec = GsSq8Codec::train(&vecs);
921 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
922
923 let mut max_rel = 0.0f32;
924 for i in 0..vecs.len() {
925 for j in (i + 1)..vecs.len().min(i + 10) {
926 let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
927 let approx = codec.l2_sq(&encoded[i], &encoded[j]);
928 let denom = true_l2.max(1e-6);
929 let rel = (approx - true_l2).abs() / denom;
930 if rel > max_rel {
931 max_rel = rel;
932 }
933 }
934 }
935 assert!(
936 max_rel < 0.15,
937 "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
938 );
939 }
940
941 #[test]
942 fn gs_train_flat_matches_train_rows() {
943 let vecs = rand_vecs(50, 16, 321);
944 let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
945
946 let codec_rows = GsSq8Codec::train(&vecs);
947 let codec_flat = GsSq8Codec::train_flat(&flat, 16);
948
949 assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
950 for d in 0..16 {
951 assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
952 }
953 }
954
955 #[test]
956 fn gs_l2_sq_order_preservation_triplets() {
957 let vecs = rand_vecs(300, 64, 303);
958 let codec = GsSq8Codec::train(&vecs);
959 let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
960
961 let n = vecs.len();
962 let mut agree = 0usize;
963 let mut total = 0usize;
964
965 for anchor in 0..50 {
966 let a = &vecs[anchor];
967 let ea = &encoded[anchor];
968 for b_idx in 0..n {
969 for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
970 let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
971 let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
972
973 let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
974 let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
975
976 if (dist_ab_true - dist_ac_true).abs() < 0.001 {
977 continue;
978 }
979
980 let true_closer_b = dist_ab_true < dist_ac_true;
981 let approx_closer_b = dist_ab_approx < dist_ac_approx;
982 if true_closer_b == approx_closer_b {
983 agree += 1;
984 }
985 total += 1;
986 }
987 }
988 }
989
990 let rate = agree as f64 / total.max(1) as f64;
991 assert!(
992 rate >= 0.95,
993 "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
994 );
995 }
996}