1use alloc::vec;
14use alloc::vec::Vec;
15use core::fmt;
16
17#[derive(Debug, Clone, PartialEq)]
25pub struct Sq8Vector {
26 pub min: f32,
27 pub max: f32,
28 pub bytes: Vec<u8>,
29}
30
31impl Sq8Vector {
32 #[must_use]
34 pub fn dim(&self) -> usize {
35 self.bytes.len()
36 }
37}
38
39const RANGE_FLOOR: f32 = 1e-12;
43
44#[must_use]
51pub fn quantize(v: &[f32]) -> Sq8Vector {
52 if v.is_empty() {
53 return Sq8Vector {
54 min: 0.0,
55 max: 0.0,
56 bytes: Vec::new(),
57 };
58 }
59 let mut min = v[0];
60 let mut max = v[0];
61 for &x in &v[1..] {
62 if x < min {
63 min = x;
64 }
65 if x > max {
66 max = x;
67 }
68 }
69 let range = max - min;
70 let bytes: Vec<u8> = if range <= RANGE_FLOOR {
71 vec![0u8; v.len()]
72 } else {
73 let scale = 255.0 / range;
74 v.iter()
75 .map(|&x| {
76 let mapped = ((x - min) * scale) + 0.5;
77 clamp_to_u8(mapped)
78 })
79 .collect()
80 };
81 Sq8Vector { min, max, bytes }
82}
83
84#[must_use]
88pub fn dequantize(q: &Sq8Vector) -> Vec<f32> {
89 if q.bytes.is_empty() {
90 return Vec::new();
91 }
92 let range = q.max - q.min;
93 if range <= RANGE_FLOOR {
94 return vec![q.min; q.bytes.len()];
95 }
96 let inv = range / 255.0;
97 q.bytes
98 .iter()
99 .map(|&b| q.min + f32::from(b) * inv)
100 .collect()
101}
102
103#[inline]
107#[allow(
108 clippy::cast_possible_truncation,
109 clippy::cast_sign_loss,
110 reason = "guarded by NaN check + (0.0, 255.0) range bracket above"
111)]
112fn clamp_to_u8(x: f32) -> u8 {
113 if x.is_nan() {
114 return 0;
115 }
116 if x <= 0.0 {
117 0
118 } else if x >= 255.0 {
119 255
120 } else {
121 x as u8
122 }
123}
124
125#[must_use]
143pub fn sq8_l2_distance_sq(a: &Sq8Vector, b: &Sq8Vector) -> f32 {
144 if a.bytes.len() != b.bytes.len() {
145 return f32::INFINITY;
146 }
147 let inv_a = sq8_step(a);
148 let inv_b = sq8_step(b);
149 let mut acc: f32 = 0.0;
150 for (&ba, &bb) in a.bytes.iter().zip(b.bytes.iter()) {
151 let xa = a.min + f32::from(ba) * inv_a;
152 let xb = b.min + f32::from(bb) * inv_b;
153 let d = xa - xb;
154 acc += d * d;
155 }
156 acc
157}
158
159#[must_use]
167pub fn sq8_l2_distance_sq_asymmetric(a: &Sq8Vector, q: &[f32]) -> f32 {
168 if a.bytes.len() != q.len() {
169 return f32::INFINITY;
170 }
171 #[cfg(target_arch = "aarch64")]
172 {
173 let n = a.bytes.len();
174 if n >= 16 && n.is_multiple_of(16) {
175 return unsafe { sq8_l2_distance_sq_asymmetric_neon(a, q) };
178 }
179 }
180 sq8_l2_distance_sq_asymmetric_scalar(a, q)
181}
182
183fn sq8_l2_distance_sq_asymmetric_scalar(a: &Sq8Vector, q: &[f32]) -> f32 {
184 let inv_a = sq8_step(a);
185 let mut acc: f32 = 0.0;
186 for (&ba, &qx) in a.bytes.iter().zip(q.iter()) {
187 let xa = a.min + f32::from(ba) * inv_a;
188 let d = xa - qx;
189 acc += d * d;
190 }
191 acc
192}
193
194#[cfg(target_arch = "aarch64")]
195#[target_feature(enable = "neon")]
196#[allow(clippy::many_single_char_names)] unsafe fn sq8_l2_distance_sq_asymmetric_neon(a: &Sq8Vector, q: &[f32]) -> f32 {
198 use core::arch::aarch64::{
199 float32x4_t, vaddq_f32, vaddvq_f32, vcvtq_f32_u32, vdupq_n_f32, vfmaq_f32, vget_high_u16,
200 vget_low_u16, vld1_u8, vld1q_f32, vmovl_u8, vmovl_u16, vsubq_f32,
201 };
202 unsafe {
203 let step = vdupq_n_f32(sq8_step(a));
204 let bias = vdupq_n_f32(a.min);
205 let zero: float32x4_t = vdupq_n_f32(0.0);
206 let mut acc0 = zero;
207 let mut acc1 = zero;
208 let n = a.bytes.len();
209 let mut i = 0usize;
210 while i + 16 <= n {
211 let lo8 = vld1_u8(a.bytes.as_ptr().add(i));
215 let hi8 = vld1_u8(a.bytes.as_ptr().add(i + 8));
216 let lo16 = vmovl_u8(lo8); let hi16 = vmovl_u8(hi8);
218 let xa0 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))));
219 let xa1 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))));
220 let xa2 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))));
221 let xa3 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))));
222 let q0 = vld1q_f32(q.as_ptr().add(i));
223 let q1 = vld1q_f32(q.as_ptr().add(i + 4));
224 let q2 = vld1q_f32(q.as_ptr().add(i + 8));
225 let q3 = vld1q_f32(q.as_ptr().add(i + 12));
226 let d0 = vsubq_f32(xa0, q0);
227 let d1 = vsubq_f32(xa1, q1);
228 let d2 = vsubq_f32(xa2, q2);
229 let d3 = vsubq_f32(xa3, q3);
230 acc0 = vfmaq_f32(acc0, d0, d0);
231 acc1 = vfmaq_f32(acc1, d1, d1);
232 acc0 = vfmaq_f32(acc0, d2, d2);
233 acc1 = vfmaq_f32(acc1, d3, d3);
234 i += 16;
235 }
236 vaddvq_f32(vaddq_f32(acc0, acc1))
237 }
238}
239
240#[must_use]
243pub fn sq8_inner_product(a: &Sq8Vector, b: &Sq8Vector) -> f32 {
244 if a.bytes.len() != b.bytes.len() {
245 return f32::INFINITY;
246 }
247 let inv_a = sq8_step(a);
248 let inv_b = sq8_step(b);
249 let mut dot: f32 = 0.0;
250 for (&ba, &bb) in a.bytes.iter().zip(b.bytes.iter()) {
251 let xa = a.min + f32::from(ba) * inv_a;
252 let xb = b.min + f32::from(bb) * inv_b;
253 dot += xa * xb;
254 }
255 -dot
256}
257
258#[must_use]
262pub fn sq8_inner_product_asymmetric(a: &Sq8Vector, q: &[f32]) -> f32 {
263 if a.bytes.len() != q.len() {
264 return f32::INFINITY;
265 }
266 #[cfg(target_arch = "aarch64")]
267 {
268 let n = a.bytes.len();
269 if n >= 16 && n.is_multiple_of(16) {
270 return -unsafe { sq8_dot_asymmetric_neon(a, q) };
272 }
273 }
274 -sq8_dot_asymmetric_scalar(a, q)
275}
276
277fn sq8_dot_asymmetric_scalar(a: &Sq8Vector, q: &[f32]) -> f32 {
278 let inv_a = sq8_step(a);
279 let mut dot: f32 = 0.0;
280 for (&ba, &qx) in a.bytes.iter().zip(q.iter()) {
281 let xa = a.min + f32::from(ba) * inv_a;
282 dot += xa * qx;
283 }
284 dot
285}
286
287#[cfg(target_arch = "aarch64")]
288#[target_feature(enable = "neon")]
289#[allow(clippy::many_single_char_names)]
290unsafe fn sq8_dot_asymmetric_neon(a: &Sq8Vector, q: &[f32]) -> f32 {
291 use core::arch::aarch64::{
292 float32x4_t, vaddq_f32, vaddvq_f32, vcvtq_f32_u32, vdupq_n_f32, vfmaq_f32, vget_high_u16,
293 vget_low_u16, vld1_u8, vld1q_f32, vmovl_u8, vmovl_u16,
294 };
295 unsafe {
296 let step = vdupq_n_f32(sq8_step(a));
297 let bias = vdupq_n_f32(a.min);
298 let zero: float32x4_t = vdupq_n_f32(0.0);
299 let mut acc0 = zero;
300 let mut acc1 = zero;
301 let n = a.bytes.len();
302 let mut i = 0usize;
303 while i + 16 <= n {
304 let lo8 = vld1_u8(a.bytes.as_ptr().add(i));
305 let hi8 = vld1_u8(a.bytes.as_ptr().add(i + 8));
306 let lo16 = vmovl_u8(lo8);
307 let hi16 = vmovl_u8(hi8);
308 let xa0 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))));
309 let xa1 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))));
310 let xa2 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))));
311 let xa3 = vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))));
312 acc0 = vfmaq_f32(acc0, xa0, vld1q_f32(q.as_ptr().add(i)));
313 acc1 = vfmaq_f32(acc1, xa1, vld1q_f32(q.as_ptr().add(i + 4)));
314 acc0 = vfmaq_f32(acc0, xa2, vld1q_f32(q.as_ptr().add(i + 8)));
315 acc1 = vfmaq_f32(acc1, xa3, vld1q_f32(q.as_ptr().add(i + 12)));
316 i += 16;
317 }
318 vaddvq_f32(vaddq_f32(acc0, acc1))
319 }
320}
321
322#[must_use]
326pub fn sq8_cosine_distance(a: &Sq8Vector, b: &Sq8Vector) -> f32 {
327 if a.bytes.len() != b.bytes.len() {
328 return f32::INFINITY;
329 }
330 let inv_a = sq8_step(a);
331 let inv_b = sq8_step(b);
332 let (mut dot, mut na, mut nb) = (0.0_f32, 0.0_f32, 0.0_f32);
333 for (&ba, &bb) in a.bytes.iter().zip(b.bytes.iter()) {
334 let xa = a.min + f32::from(ba) * inv_a;
335 let xb = b.min + f32::from(bb) * inv_b;
336 dot += xa * xb;
337 na += xa * xa;
338 nb += xb * xb;
339 }
340 if na == 0.0 || nb == 0.0 {
341 return f32::INFINITY;
342 }
343 1.0 - dot / (sqrt_finite(na) * sqrt_finite(nb))
344}
345
346#[must_use]
350pub fn sq8_cosine_distance_asymmetric(a: &Sq8Vector, q: &[f32]) -> f32 {
351 if a.bytes.len() != q.len() {
352 return f32::INFINITY;
353 }
354 let (dot, na, nq);
355 #[cfg(target_arch = "aarch64")]
356 {
357 let n = a.bytes.len();
358 if n >= 16 && n.is_multiple_of(16) {
359 let (d, a2, q2) = unsafe { sq8_cosine_accumulators_asymmetric_neon(a, q) };
361 dot = d;
362 na = a2;
363 nq = q2;
364 } else {
365 let (d, a2, q2) = sq8_cosine_accumulators_asymmetric_scalar(a, q);
366 dot = d;
367 na = a2;
368 nq = q2;
369 }
370 }
371 #[cfg(not(target_arch = "aarch64"))]
372 {
373 let (d, a2, q2) = sq8_cosine_accumulators_asymmetric_scalar(a, q);
374 dot = d;
375 na = a2;
376 nq = q2;
377 }
378 if na == 0.0 || nq == 0.0 {
379 return f32::INFINITY;
380 }
381 1.0 - dot / (sqrt_finite(na) * sqrt_finite(nq))
382}
383
384fn sq8_cosine_accumulators_asymmetric_scalar(a: &Sq8Vector, q: &[f32]) -> (f32, f32, f32) {
385 let inv_a = sq8_step(a);
386 let (mut dot, mut na, mut nq) = (0.0_f32, 0.0_f32, 0.0_f32);
387 for (&ba, &qx) in a.bytes.iter().zip(q.iter()) {
388 let xa = a.min + f32::from(ba) * inv_a;
389 dot += xa * qx;
390 na += xa * xa;
391 nq += qx * qx;
392 }
393 (dot, na, nq)
394}
395
396#[cfg(target_arch = "aarch64")]
397#[target_feature(enable = "neon")]
398#[allow(clippy::many_single_char_names, clippy::similar_names)]
399unsafe fn sq8_cosine_accumulators_asymmetric_neon(a: &Sq8Vector, q: &[f32]) -> (f32, f32, f32) {
400 use core::arch::aarch64::{
401 float32x4_t, vaddvq_f32, vcvtq_f32_u32, vdupq_n_f32, vfmaq_f32, vget_high_u16,
402 vget_low_u16, vld1_u8, vld1q_f32, vmovl_u8, vmovl_u16,
403 };
404 unsafe {
405 let step = vdupq_n_f32(sq8_step(a));
406 let bias = vdupq_n_f32(a.min);
407 let zero: float32x4_t = vdupq_n_f32(0.0);
408 let mut acc_dot = zero;
409 let mut acc_na = zero;
410 let mut acc_nq = zero;
411 let n = a.bytes.len();
412 let mut i = 0usize;
413 while i + 16 <= n {
414 let lo8 = vld1_u8(a.bytes.as_ptr().add(i));
415 let hi8 = vld1_u8(a.bytes.as_ptr().add(i + 8));
416 let lo16 = vmovl_u8(lo8);
417 let hi16 = vmovl_u8(hi8);
418 let xs = [
419 vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16)))),
420 vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16)))),
421 vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16)))),
422 vfmaq_f32(bias, step, vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16)))),
423 ];
424 let qs = [
425 vld1q_f32(q.as_ptr().add(i)),
426 vld1q_f32(q.as_ptr().add(i + 4)),
427 vld1q_f32(q.as_ptr().add(i + 8)),
428 vld1q_f32(q.as_ptr().add(i + 12)),
429 ];
430 for k in 0..4 {
431 acc_dot = vfmaq_f32(acc_dot, xs[k], qs[k]);
432 acc_na = vfmaq_f32(acc_na, xs[k], xs[k]);
433 acc_nq = vfmaq_f32(acc_nq, qs[k], qs[k]);
434 }
435 i += 16;
436 }
437 (vaddvq_f32(acc_dot), vaddvq_f32(acc_na), vaddvq_f32(acc_nq))
438 }
439}
440
441#[inline]
445fn sq8_step(q: &Sq8Vector) -> f32 {
446 let range = q.max - q.min;
447 if range <= RANGE_FLOOR {
448 0.0
449 } else {
450 range / 255.0
451 }
452}
453
454#[inline]
458fn sqrt_finite(x: f32) -> f32 {
459 if x <= 0.0 {
460 return 0.0;
461 }
462 let mut y = if x >= 1.0 { x * 0.5 } else { (x + 1.0) * 0.5 };
463 for _ in 0..6 {
464 y = 0.5 * (y + x / y);
465 }
466 y
467}
468
469#[cfg(test)]
470#[allow(
471 clippy::cast_lossless,
472 clippy::cast_possible_truncation,
473 clippy::cast_precision_loss,
474 clippy::cast_sign_loss,
475 clippy::doc_markdown,
476 clippy::useless_conversion,
477 clippy::similar_names,
478 clippy::unreadable_literal,
479 clippy::items_after_statements,
480 clippy::too_many_lines,
481 clippy::float_cmp,
482 clippy::suboptimal_flops,
483 clippy::cast_possible_wrap
484)]
485mod tests {
486 use super::*;
487
488 struct SplitMix64 {
492 state: u64,
493 }
494
495 impl SplitMix64 {
496 const fn new(seed: u64) -> Self {
497 Self { state: seed }
498 }
499
500 fn next_u64(&mut self) -> u64 {
501 self.state = self.state.wrapping_add(0x9E37_79B9_7F4A_7C15);
502 let mut z = self.state;
503 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
504 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
505 z ^ (z >> 31)
506 }
507
508 fn next_unit_f32(&mut self) -> f32 {
510 let bits = (self.next_u64() >> 40) as u32;
512 (bits as f32) / ((1u32 << 24) as f32)
513 }
514
515 fn next_gaussian_f32(&mut self) -> f32 {
517 let u = 1.0 - self.next_unit_f32();
519 let v = self.next_unit_f32();
520 let r = sqrt_f32(-2.0 * ln_f32(u));
521 let theta = 2.0 * core::f32::consts::PI * v;
522 r * cos_f32(theta)
523 }
524 }
525
526 fn sqrt_f32(x: f32) -> f32 {
530 if x <= 0.0 {
531 return 0.0;
532 }
533 let mut y = if x >= 1.0 { x * 0.5 } else { (x + 1.0) * 0.5 };
534 for _ in 0..6 {
535 y = 0.5 * (y + x / y);
536 }
537 y
538 }
539
540 fn ln_f32(x: f32) -> f32 {
544 if x <= 0.0 {
545 return f32::NEG_INFINITY;
546 }
547 let mut k: i32 = 0;
549 let mut m = x;
550 while m >= 1.0 {
551 m *= 0.5;
552 k += 1;
553 }
554 while m < 0.5 {
555 m *= 2.0;
556 k -= 1;
557 }
558 let u = (m - 1.0) / (m + 1.0);
560 let u2 = u * u;
561 let mut term = u;
562 let mut sum = 0.0;
563 for i in 0..16 {
564 sum += term / ((2 * i + 1) as f32);
565 term *= u2;
566 }
567 2.0 * sum + (k as f32) * core::f32::consts::LN_2
568 }
569
570 fn cos_f32(theta: f32) -> f32 {
574 let two_pi = 2.0 * core::f32::consts::PI;
575 let mut t = theta % two_pi;
576 if t > core::f32::consts::PI {
577 t -= two_pi;
578 } else if t < -core::f32::consts::PI {
579 t += two_pi;
580 }
581 let t2 = t * t;
582 1.0 - t2 / 2.0 + t2 * t2 / 24.0 - t2 * t2 * t2 / 720.0 + t2 * t2 * t2 * t2 / 40_320.0
584 - t2 * t2 * t2 * t2 * t2 / 3_628_800.0
585 }
586
587 fn random_gaussian_vec(rng: &mut SplitMix64, dim: usize) -> Vec<f32> {
588 (0..dim).map(|_| rng.next_gaussian_f32()).collect()
589 }
590
591 fn random_unit_vec(rng: &mut SplitMix64, dim: usize) -> Vec<f32> {
592 (0..dim).map(|_| rng.next_unit_f32() * 2.0 - 1.0).collect()
593 }
594
595 fn linf_error(a: &[f32], b: &[f32]) -> f32 {
596 let mut e: f32 = 0.0;
597 for (x, y) in a.iter().zip(b.iter()) {
598 let d = (x - y).abs();
599 if d > e {
600 e = d;
601 }
602 }
603 e
604 }
605
606 #[test]
607 fn quantize_empty_vector_is_zero_dim() {
608 let q = quantize(&[]);
609 assert_eq!(q.dim(), 0);
610 assert_eq!(q.min, 0.0);
611 assert_eq!(q.max, 0.0);
612 assert!(dequantize(&q).is_empty());
613 }
614
615 #[test]
616 fn quantize_single_element_roundtrips_exactly() {
617 let q = quantize(&[3.25]);
618 assert_eq!(q.dim(), 1);
619 assert_eq!(q.min, 3.25);
620 assert_eq!(q.max, 3.25);
621 let d = dequantize(&q);
622 assert_eq!(d.len(), 1);
623 assert!((d[0] - 3.25).abs() < 1e-6);
625 }
626
627 #[test]
628 fn quantize_constant_vector_roundtrips_exactly() {
629 let v = vec![7.5_f32; 64];
630 let q = quantize(&v);
631 assert_eq!(q.min, 7.5);
632 assert_eq!(q.max, 7.5);
633 let d = dequantize(&q);
634 for x in &d {
635 assert!((x - 7.5).abs() < 1e-6);
636 }
637 }
638
639 #[test]
640 fn quantize_min_and_max_endpoints_reconstruct_exactly() {
641 let v = vec![-2.0_f32, 0.0, 5.0, 3.0, -2.0, 5.0];
642 let q = quantize(&v);
643 assert_eq!(q.min, -2.0);
644 assert_eq!(q.max, 5.0);
645 let d = dequantize(&q);
646 assert!((d[0] - (-2.0)).abs() < 1e-5);
649 assert!((d[2] - 5.0).abs() < 1e-5);
650 assert!((d[4] - (-2.0)).abs() < 1e-5);
651 assert!((d[5] - 5.0).abs() < 1e-5);
652 }
653
654 #[test]
655 fn quantize_dequantize_roundtrip_bounded_error_gaussian() {
656 let mut rng = SplitMix64::new(0xDEAD_BEEF_CAFE_F00D);
657 for dim in [32_usize, 128, 512, 1024] {
658 for _trial in 0..250 {
659 let v = random_gaussian_vec(&mut rng, dim);
660 let q = quantize(&v);
661 let r = dequantize(&q);
662 let step = (q.max - q.min) / 510.0;
666 let bound = step + 1e-6_f32.max(step * 1e-3);
667 let err = linf_error(&v, &r);
668 assert!(
669 err <= bound,
670 "dim={dim} err={err} bound={bound} range={}",
671 q.max - q.min
672 );
673 }
674 }
675 }
676
677 fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
680 a.iter().zip(b.iter()).map(|(x, y)| (x - y).powi(2)).sum()
681 }
682
683 fn inner_product_f32(a: &[f32], b: &[f32]) -> f32 {
684 -a.iter().zip(b.iter()).map(|(x, y)| x * y).sum::<f32>()
685 }
686
687 fn cosine_distance_f32(a: &[f32], b: &[f32]) -> f32 {
688 let (mut dot, mut na, mut nb) = (0.0_f32, 0.0_f32, 0.0_f32);
689 for (x, y) in a.iter().zip(b.iter()) {
690 dot += x * y;
691 na += x * x;
692 nb += y * y;
693 }
694 if na == 0.0 || nb == 0.0 {
695 return f32::INFINITY;
696 }
697 1.0 - dot / (sqrt_f32(na) * sqrt_f32(nb))
698 }
699
700 fn float_tolerance_for_dim(dim: usize) -> f32 {
709 1e-4 * dim as f32
712 }
713
714 #[test]
715 fn sq8_l2_distance_matches_dequantize_then_f32() {
716 let mut rng = SplitMix64::new(0xABCD_0001_2345_6789);
717 for dim in [32_usize, 128, 512, 1024] {
718 let tol = float_tolerance_for_dim(dim);
719 for _ in 0..2500 {
720 let a = random_gaussian_vec(&mut rng, dim);
721 let b = random_gaussian_vec(&mut rng, dim);
722 let qa = quantize(&a);
723 let qb = quantize(&b);
724 let dqa = dequantize(&qa);
725 let dqb = dequantize(&qb);
726 let want_sym = l2_sq_f32(&dqa, &dqb);
727 let want_asym = l2_sq_f32(&dqa, &b);
728 let got_sym = sq8_l2_distance_sq(&qa, &qb);
729 let got_asym = sq8_l2_distance_sq_asymmetric(&qa, &b);
730 let err_sym = (got_sym - want_sym).abs();
731 let err_asym = (got_asym - want_asym).abs();
732 let scale = want_sym.abs().max(want_asym.abs()).max(1.0);
733 assert!(
734 err_sym <= tol * scale,
735 "dim={dim} sym got={got_sym} want={want_sym} err={err_sym} tol={}",
736 tol * scale
737 );
738 assert!(
739 err_asym <= tol * scale,
740 "dim={dim} asym got={got_asym} want={want_asym} err={err_asym} tol={}",
741 tol * scale
742 );
743 }
744 }
745 }
746
747 #[test]
748 fn sq8_inner_product_matches_dequantize_then_f32() {
749 let mut rng = SplitMix64::new(0xABCD_0002_2345_6789);
750 for dim in [32_usize, 128, 512, 1024] {
751 let tol = float_tolerance_for_dim(dim);
752 for _ in 0..2500 {
753 let a = random_gaussian_vec(&mut rng, dim);
754 let b = random_gaussian_vec(&mut rng, dim);
755 let qa = quantize(&a);
756 let qb = quantize(&b);
757 let dqa = dequantize(&qa);
758 let dqb = dequantize(&qb);
759 let want_sym = inner_product_f32(&dqa, &dqb);
760 let want_asym = inner_product_f32(&dqa, &b);
761 let got_sym = sq8_inner_product(&qa, &qb);
762 let got_asym = sq8_inner_product_asymmetric(&qa, &b);
763 let scale = want_sym.abs().max(want_asym.abs()).max(1.0);
764 let err_sym = (got_sym - want_sym).abs();
765 let err_asym = (got_asym - want_asym).abs();
766 assert!(
767 err_sym <= tol * scale,
768 "dim={dim} sym got={got_sym} want={want_sym} err={err_sym}"
769 );
770 assert!(
771 err_asym <= tol * scale,
772 "dim={dim} asym got={got_asym} want={want_asym} err={err_asym}"
773 );
774 }
775 }
776 }
777
778 #[test]
779 fn sq8_cosine_distance_matches_dequantize_then_f32() {
780 let mut rng = SplitMix64::new(0xABCD_0003_2345_6789);
781 for dim in [32_usize, 128, 512, 1024] {
782 let tol = float_tolerance_for_dim(dim);
783 for _ in 0..2500 {
784 let a = random_gaussian_vec(&mut rng, dim);
785 let b = random_gaussian_vec(&mut rng, dim);
786 let qa = quantize(&a);
787 let qb = quantize(&b);
788 let dqa = dequantize(&qa);
789 let dqb = dequantize(&qb);
790 let want_sym = cosine_distance_f32(&dqa, &dqb);
791 let want_asym = cosine_distance_f32(&dqa, &b);
792 let got_sym = sq8_cosine_distance(&qa, &qb);
793 let got_asym = sq8_cosine_distance_asymmetric(&qa, &b);
794 let bound = tol;
797 assert!(
798 (got_sym - want_sym).abs() <= bound,
799 "dim={dim} sym got={got_sym} want={want_sym}"
800 );
801 assert!(
802 (got_asym - want_asym).abs() <= bound,
803 "dim={dim} asym got={got_asym} want={want_asym}"
804 );
805 }
806 }
807 }
808
809 #[test]
810 fn sq8_distance_handles_dim_mismatch_with_infinity() {
811 let a = quantize(&[1.0, 2.0, 3.0]);
812 let b = quantize(&[1.0, 2.0]);
813 assert_eq!(sq8_l2_distance_sq(&a, &b), f32::INFINITY);
814 assert_eq!(sq8_inner_product(&a, &b), f32::INFINITY);
815 assert_eq!(sq8_cosine_distance(&a, &b), f32::INFINITY);
816 assert_eq!(sq8_l2_distance_sq_asymmetric(&a, &[1.0]), f32::INFINITY);
817 }
818
819 #[test]
820 fn sq8_cosine_handles_zero_norm_with_infinity() {
821 let zero = quantize(&[0.0_f32; 8]);
822 let nonzero = quantize(&[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]);
823 assert_eq!(sq8_cosine_distance(&zero, &nonzero), f32::INFINITY);
824 assert_eq!(
825 sq8_cosine_distance_asymmetric(&zero, &[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]),
826 f32::INFINITY
827 );
828 }
829
830 #[test]
831 fn quantize_dequantize_roundtrip_bounded_error_uniform() {
832 let mut rng = SplitMix64::new(0xF0F0_F0F0_F0F0_F0F0);
833 for dim in [32_usize, 128, 512, 1024] {
834 for _trial in 0..250 {
835 let v = random_unit_vec(&mut rng, dim);
836 let q = quantize(&v);
837 let r = dequantize(&q);
838 let step = (q.max - q.min) / 510.0;
839 let bound = step + 1e-6_f32.max(step * 1e-3);
840 let err = linf_error(&v, &r);
841 assert!(
842 err <= bound,
843 "dim={dim} err={err} bound={bound} range={}",
844 q.max - q.min
845 );
846 }
847 }
848 }
849
850 fn topk_indices_l2(corpus: &[Vec<f32>], query: &[f32], k: usize) -> Vec<usize> {
856 let mut scored: Vec<(f32, usize)> = corpus
857 .iter()
858 .enumerate()
859 .map(|(i, v)| (l2_sq_f32(v, query), i))
860 .collect();
861 scored.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(core::cmp::Ordering::Equal));
862 scored.into_iter().take(k).map(|(_, i)| i).collect()
863 }
864
865 fn topk_indices_l2_sq8_asym(corpus: &[Sq8Vector], query: &[f32], k: usize) -> Vec<usize> {
866 let mut scored: Vec<(f32, usize)> = corpus
867 .iter()
868 .enumerate()
869 .map(|(i, qv)| (sq8_l2_distance_sq_asymmetric(qv, query), i))
870 .collect();
871 scored.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(core::cmp::Ordering::Equal));
872 scored.into_iter().take(k).map(|(_, i)| i).collect()
873 }
874
875 fn overlap_fraction(a: &[usize], b: &[usize]) -> f32 {
876 let mut hits = 0;
877 for &x in a {
878 if b.contains(&x) {
879 hits += 1;
880 }
881 }
882 hits as f32 / a.len() as f32
883 }
884
885 #[test]
886 fn sq8_recall_at_10_above_0_95_gaussian() {
887 const N: usize = 10_000;
888 const Q: usize = 100;
889 const K: usize = 10;
890 const DIM: usize = 128;
891
892 let mut rng = SplitMix64::new(0x5EED_5EED_5EED_5EED);
893 let corpus_f32: Vec<Vec<f32>> =
894 (0..N).map(|_| random_gaussian_vec(&mut rng, DIM)).collect();
895 let corpus_sq8: Vec<Sq8Vector> = corpus_f32.iter().map(|v| quantize(v)).collect();
896
897 let mut total_recall: f32 = 0.0;
898 for _ in 0..Q {
899 let query = random_gaussian_vec(&mut rng, DIM);
900 let truth = topk_indices_l2(&corpus_f32, &query, K);
901 let sq8_top = topk_indices_l2_sq8_asym(&corpus_sq8, &query, K);
902 total_recall += overlap_fraction(&truth, &sq8_top);
903 }
904 let avg = total_recall / Q as f32;
905 assert!(
906 avg >= 0.95,
907 "Gaussian recall@10 average = {avg} (need ≥ 0.95)"
908 );
909 }
910
911 #[test]
912 fn sq8_recall_at_10_above_0_93_uniform_unit_sphere() {
913 const N: usize = 10_000;
914 const Q: usize = 100;
915 const K: usize = 10;
916 const DIM: usize = 128;
917
918 let mut rng = SplitMix64::new(0xC0DE_C0DE_C0DE_C0DE);
919 let normalise = |mut v: Vec<f32>| -> Vec<f32> {
921 let n = sqrt_f32(v.iter().map(|x| x * x).sum::<f32>()).max(1e-12);
922 for x in &mut v {
923 *x /= n;
924 }
925 v
926 };
927 let corpus_f32: Vec<Vec<f32>> = (0..N)
928 .map(|_| normalise(random_gaussian_vec(&mut rng, DIM)))
929 .collect();
930 let corpus_sq8: Vec<Sq8Vector> = corpus_f32.iter().map(|v| quantize(v)).collect();
931
932 let mut total_recall: f32 = 0.0;
933 for _ in 0..Q {
934 let query = normalise(random_gaussian_vec(&mut rng, DIM));
935 let truth = topk_indices_l2(&corpus_f32, &query, K);
936 let sq8_top = topk_indices_l2_sq8_asym(&corpus_sq8, &query, K);
937 total_recall += overlap_fraction(&truth, &sq8_top);
938 }
939 let avg = total_recall / Q as f32;
940 assert!(
941 avg >= 0.93,
942 "Unit-sphere recall@10 average = {avg} (need ≥ 0.93)"
943 );
944 }
945
946 #[test]
949 fn sq8_serde_roundtrip_preserves_all_fields() {
950 let mut rng = SplitMix64::new(0xBEEF_F00D_DEAD_0123);
951 for dim in [0_usize, 1, 7, 32, 128, 1024] {
952 for _ in 0..200 {
953 let v = random_gaussian_vec(&mut rng, dim);
954 let q = quantize(&v);
955 let bytes = q.to_bytes();
956 assert_eq!(bytes.len(), Sq8Vector::encoded_size_for(dim));
957 let back = Sq8Vector::from_bytes(&bytes).expect("from_bytes");
958 assert_eq!(back, q, "dim={dim} roundtrip mismatch");
959 }
960 }
961 }
962
963 #[test]
964 fn sq8_from_bytes_rejects_truncated_header() {
965 for short in [0_usize, 1, 4, 8, 11] {
966 let buf = vec![0u8; short];
967 assert_eq!(Sq8Vector::from_bytes(&buf), Err(QuantizeError::Truncated));
968 }
969 }
970
971 #[cfg(target_arch = "aarch64")]
972 #[test]
973 fn sq8_adc_ip_asymmetric_neon_matches_scalar() {
974 let dims = [16usize, 32, 64, 128, 256, 512, 1024];
978 for &d in &dims {
979 let mut rng = SplitMix64::new(0xBEEF_DEAD_1234_A5A5u64 ^ d as u64);
980 for _ in 0..16 {
981 let v = random_gaussian_vec(&mut rng, d);
982 let q = random_gaussian_vec(&mut rng, d);
983 let sq = quantize(&v);
984 let scalar = -sq8_dot_asymmetric_scalar(&sq, &q);
985 let neon = -unsafe { sq8_dot_asymmetric_neon(&sq, &q) };
986 let tol = (scalar.abs().max(1e-6)) * 1e-4 + (d as f32) * 1e-5;
987 assert!(
988 (scalar - neon).abs() <= tol,
989 "IP asym dim={d}: scalar={scalar} neon={neon} diff={}",
990 (scalar - neon).abs()
991 );
992 }
993 }
994 }
995
996 #[cfg(target_arch = "aarch64")]
997 #[test]
998 fn sq8_adc_cosine_asymmetric_neon_matches_scalar() {
999 let dims = [16usize, 32, 64, 128, 256, 512, 1024];
1003 for &d in &dims {
1004 let mut rng = SplitMix64::new(0xC0DE_F00D_1234_5678u64 ^ d as u64);
1005 for _ in 0..16 {
1006 let v = random_gaussian_vec(&mut rng, d);
1007 let q = random_gaussian_vec(&mut rng, d);
1008 let sq = quantize(&v);
1009 let (dot_s, na_s, nq_s) = sq8_cosine_accumulators_asymmetric_scalar(&sq, &q);
1010 let (dot_n, na_n, nq_n) =
1011 unsafe { sq8_cosine_accumulators_asymmetric_neon(&sq, &q) };
1012 let tol = |x: f32| (x.abs().max(1e-6)) * 1e-4 + (d as f32) * 1e-5;
1013 assert!(
1014 (dot_s - dot_n).abs() <= tol(dot_s),
1015 "cos dot dim={d}: scalar={dot_s} neon={dot_n}"
1016 );
1017 assert!(
1018 (na_s - na_n).abs() <= tol(na_s),
1019 "cos na dim={d}: scalar={na_s} neon={na_n}"
1020 );
1021 assert!(
1022 (nq_s - nq_n).abs() <= tol(nq_s),
1023 "cos nq dim={d}: scalar={nq_s} neon={nq_n}"
1024 );
1025 }
1026 }
1027 }
1028
1029 #[cfg(target_arch = "aarch64")]
1030 #[test]
1031 fn sq8_adc_l2_asymmetric_neon_matches_scalar() {
1032 let dims = [16usize, 32, 48, 64, 128, 256, 512, 1024];
1039 for &d in &dims {
1040 let mut rng = SplitMix64::new(0xA5A5_1234_DEAD_BEEFu64 ^ d as u64);
1041 for _ in 0..16 {
1042 let v = random_gaussian_vec(&mut rng, d);
1043 let q = random_gaussian_vec(&mut rng, d);
1044 let sq = quantize(&v);
1045 let scalar = sq8_l2_distance_sq_asymmetric_scalar(&sq, &q);
1046 let neon = unsafe { sq8_l2_distance_sq_asymmetric_neon(&sq, &q) };
1047 let tol = (scalar.abs().max(1e-6)) * 1e-4 + (d as f32) * 1e-5;
1048 assert!(
1049 (scalar - neon).abs() <= tol,
1050 "L2 asym dim={d}: scalar={scalar} neon={neon} diff={}",
1051 (scalar - neon).abs()
1052 );
1053 }
1054 }
1055 }
1056
1057 #[test]
1058 fn sq8_from_bytes_rejects_dim_mismatch() {
1059 let mut buf: Vec<u8> = Vec::new();
1061 buf.extend_from_slice(&4u32.to_le_bytes());
1062 buf.extend_from_slice(&0.0f32.to_le_bytes());
1063 buf.extend_from_slice(&1.0f32.to_le_bytes());
1064 buf.extend_from_slice(&[10u8, 200u8]);
1065 assert_eq!(
1066 Sq8Vector::from_bytes(&buf),
1067 Err(QuantizeError::DimMismatch {
1068 expected: 4,
1069 got: 2
1070 })
1071 );
1072 }
1073}
1074
1075#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1077pub enum QuantizeError {
1078 Truncated,
1080 DimMismatch { expected: u32, got: u32 },
1082}
1083
1084impl fmt::Display for QuantizeError {
1085 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1086 match self {
1087 Self::Truncated => write!(f, "sq8 input truncated"),
1088 Self::DimMismatch { expected, got } => write!(
1089 f,
1090 "sq8 dim mismatch: expected {expected}, payload carries {got}"
1091 ),
1092 }
1093 }
1094}
1095
1096impl Sq8Vector {
1108 #[must_use]
1114 pub fn to_bytes(&self) -> Vec<u8> {
1115 let dim = u32::try_from(self.bytes.len())
1116 .expect("Sq8Vector dim fits in u32 by DataType::Vector contract");
1117 let mut out = Vec::with_capacity(12 + self.bytes.len());
1118 out.extend_from_slice(&dim.to_le_bytes());
1119 out.extend_from_slice(&self.min.to_le_bytes());
1120 out.extend_from_slice(&self.max.to_le_bytes());
1121 out.extend_from_slice(&self.bytes);
1122 out
1123 }
1124
1125 pub fn from_bytes(input: &[u8]) -> Result<Self, QuantizeError> {
1128 if input.len() < 12 {
1129 return Err(QuantizeError::Truncated);
1130 }
1131 let dim = u32::from_le_bytes([input[0], input[1], input[2], input[3]]);
1132 let min = f32::from_le_bytes([input[4], input[5], input[6], input[7]]);
1133 let max = f32::from_le_bytes([input[8], input[9], input[10], input[11]]);
1134 let body = &input[12..];
1135 if body.len() != dim as usize {
1136 let got = u32::try_from(body.len()).unwrap_or(u32::MAX);
1137 return Err(QuantizeError::DimMismatch { expected: dim, got });
1138 }
1139 Ok(Self {
1140 min,
1141 max,
1142 bytes: body.to_vec(),
1143 })
1144 }
1145
1146 #[must_use]
1149 pub const fn encoded_size_for(dim: usize) -> usize {
1150 12 + dim
1151 }
1152}