Skip to main content

khive_quant/
lib.rs

1//! SQ8 scalar quantization codecs for approximate distance computation in ANN indexes.
2//!
3//! Two codecs with different encoding strategies:
4//!
5//! ## `Sq8Codec` — per-dimension affine, for dot product / cosine
6//!
7//! Each dimension is mapped to [0, 255] using its own observed min/max.
8//! Dot product and cosine require per-dim scale accuracy; the residual-corrected
9//! path (`approx_dot`, `approx_cosine_dist`) preserves ordinal ranking.
10//!
11//! ## `GsSq8Codec` — global-scale affine, for L2 (Vamana acquisition)
12//!
13//! A single shared scale `gs = max_range_across_dims / 255` is used for all dims;
14//! per-dim `min_i` offsets are still subtracted before quantizing.
15//!
16//! L2² in code space: `gs² × Σ (a_i - b_i)²` — exact after the lossy f32→u8 encode
17//! (offsets cancel, one scalar factorizes out). No anisotropy gate or residual pass for
18//! in-distribution vectors; OOD queries (components outside the trained range) fall back
19//! to exact f32 in the caller (see `VamanaIndex::search`). Small-range dims contribute
20//! proportionally fewer codes and proportionally less L2 signal — an honest trade-off
21//! documented in ADR-052.
22//!
23//! # Hot-loop NEON helpers (`u8_dot_u32`, `u8_l2sq_u32`)
24//!
25//! Both codecs share these inner functions:
26//! - `u8_dot_u32`: NEON `vmull_u8` (16-wide u8→u16→u32) or chunked portable fallback.
27//! - `u8_l2sq_u32`: NEON `vabdq_u8` + `vmull_u8` squaring or chunked portable fallback.
28
29use rayon::prelude::*;
30
31// ─── NEON helpers ─────────────────────────────────────────────────────────────
32
33/// Compute `Σ a_i * b_i` over `u8` slices as a `u32` accumulator using NEON
34/// `vmull_u8` (8-wide u8→u16 widening multiply) on aarch64, or a chunked
35/// portable widening fallback elsewhere.
36///
37/// Safety: both slices must have the same length.
38#[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/// Compute `Σ (a_i - b_i)²` over `u8` slices as a `u32` accumulator using NEON
99/// `vabdq_u8` (absolute difference) + `vmull_u8` squaring on aarch64, or a
100/// chunked portable fallback elsewhere.
101///
102/// Safety: both slices must have the same length.
103#[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// ─── Sq8Codec (per-dim scale — dot product / cosine) ──────────────────────────
170
171/// Per-dimension affine SQ8 codec for dot product and cosine distance.
172///
173/// Encodes `f32` dimensions to `u8` via: `code = round((x - min) / scale)`,
174/// where `scale_i = (max_i - min_i) / 255`.
175///
176/// For L2 distance on Vamana builds, use [`GsSq8Codec`] instead — its global
177/// shared scale makes L2 algebraically exact in code space without a residual pass.
178#[derive(Debug, Clone)]
179pub struct Sq8Codec {
180    /// Per-dimension minimum values.
181    pub min: Vec<f32>,
182    /// Per-dimension scale: `(max - min) / 255`.
183    pub scale: Vec<f32>,
184    /// Per-dimension `scale²` precomputed for fast L2 and dot product.
185    pub scale_sq: Vec<f32>,
186    /// Mean of `scale_sq` across all dimensions — used as the integer-pass multiplier.
187    pub mean_scale_sq: f32,
188    /// Residual: `scale_sq_i - mean_scale_sq` (zero-mean, small magnitude).
189    pub scale_sq_residual: Vec<f32>,
190    /// `Σ_i min_i²` precomputed for dot-product correction.
191    pub offset_sq_sum: f32,
192}
193
194/// A corpus vector encoded by [`Sq8Codec`].
195#[derive(Debug, Clone)]
196pub struct EncodedVector {
197    /// SQ8 u8 codes, one per dimension.
198    pub codes: Vec<u8>,
199    /// L2 norm of the original f32 vector (for cosine distance).
200    pub norm: f32,
201    /// `Σ_i scale_i * min_i * code_i` — per-vector correction term for dot product.
202    pub soc_sum: f32,
203    /// `Σ_i scale_sq_residual_i * code_i` precomputed at encode time.
204    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    /// Train a codec from row-major flat vectors.
227    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    /// Train from a slice of row vectors (each a `Vec<f32>`).
267    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    /// Encode a single vector into SQ8 codes + correction metadata.
301    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    /// Encode a batch of flat-row vectors in parallel.
329    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    /// Encode a batch of row vectors in parallel.
338    pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
339        vectors.par_iter().map(|v| self.encode(v)).collect()
340    }
341
342    /// Approximate dot product between two encoded vectors (same codec).
343    ///
344    /// Full-precision correction identity (same min/scale for both):
345    /// `dot(a, b) = Σ s²·a·b + soc_a + soc_b + offset_sq_sum`
346    ///
347    /// The integer pass (`u8_dot_u32`) computes `raw = Σ a_i*b_i` as `u32` using
348    /// NEON (16-wide on aarch64). The scale correction then applies `mean_scale_sq`
349    /// plus a compact per-dim residual f32 pass for accuracy.
350    #[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    /// Approximate cosine distance between two encoded vectors (same codec).
364    ///
365    /// Returns `1 - dot / (norm_a * norm_b)`. Falls back to 1.0 for zero norms.
366    #[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    /// Approximate squared L2 distance — per-dim residual corrected.
378    ///
379    /// Full-precision identity: `||a-b||² = Σ scale_sq_i * (a_i-b_i)²`.
380    /// Offsets cancel because both vectors share the same codec.
381    ///
382    /// The integer pass (`u8_l2sq_u32`) computes `raw = Σ (a_i-b_i)²` using NEON
383    /// `vabdq_u8` + `vmull_u8`. The residual correction keeps ordinal accuracy
384    /// across anisotropic corpora.
385    ///
386    /// For Vamana L2 acquisition use [`GsSq8Codec::l2_sq`] — algebraically exact
387    /// in code space and ~2× faster (no residual pass).
388    #[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    /// Number of dimensions.
405    pub fn dims(&self) -> usize {
406        self.min.len()
407    }
408}
409
410// ─── GsSq8Codec (global-scale — L2 / Vamana acquisition) ─────────────────────
411
412/// Global-scale SQ8 codec for L2 distance — the Vamana acquisition path.
413///
414/// A single shared scale `gs = max_range_across_dims / 255` is used for all dims.
415/// Per-dim offsets (`min_i`) are still subtracted before quantizing so codes span
416/// [0, 255] for the widest dim and fewer levels for narrower dims (honest trade-off).
417///
418/// Encoding is **lossy**: f32 components are rounded and clamped to u8 before storage.
419/// L2² in code space (`gs² × Σ (a_i - b_i)²`) is exact *after* that lossy encode —
420/// offset terms cancel and `gs²` factorizes — but the round-trip error relative to
421/// the original f32 L2² can reach ~15% for anisotropic or out-of-distribution data.
422/// Recall safety must be established by probe (see `sq8_recall_parity_vs_f32_oracle`
423/// and `sq8_ood_fallback_deterministic_ranking_flip`), not by an exactness argument.
424/// No residual pass, no gate, no silent fallback for anisotropic data.
425///
426/// Historical note: the predecessor per-dim codec required `approx_l2_sq_fast` + an
427/// anisotropy gate (ratio ≤ 4.0) to achieve the integer-only hot path. The gate was
428/// calibrated on an LCG corpus that gave ratio ≈ 4.0; real transformer embeddings
429/// have rogue dimensions (ratio 10–32) that silently fell back to the full residual
430/// path, defeating the purpose. Global-scale eliminates the gate entirely — see ADR-052.
431#[derive(Debug, Clone)]
432pub struct GsSq8Codec {
433    /// Per-dimension minimum values.
434    pub min: Vec<f32>,
435    /// Global scale: `max_range / 255` where `max_range = max_i(max_i - min_i)`.
436    pub gs: f32,
437    /// `gs²` precomputed for L2.
438    pub gs_sq: f32,
439    /// Anisotropy ratio measured at train time: `max(range_i) / min(nonzero range_i)`.
440    /// Informational only — never used for dispatch decisions.
441    pub anisotropy_ratio: f32,
442}
443
444/// A corpus vector encoded by [`GsSq8Codec`].
445#[derive(Debug, Clone)]
446pub struct GsEncodedVector {
447    /// SQ8 u8 codes, one per dimension.
448    pub codes: Vec<u8>,
449}
450
451impl GsSq8Codec {
452    /// Train from row-major flat vectors.
453    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    /// Train from a slice of row vectors.
517    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    /// Encode a single vector.
575    #[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    /// Encode a batch of flat-row vectors in parallel.
592    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    /// Approximate squared L2 distance.
601    ///
602    /// `||a-b||² ≈ gs² × Σ (a_i - b_i)²`
603    ///
604    /// Exact in code space (offset terms cancel, `gs²` factorizes) after the
605    /// lossy f32→u8 encode. Per-round-trip L2 error can reach ~15%; recall
606    /// safety is established by probe, not by this formula.
607    /// The NEON path runs ~13 ns at 384-d.
608    #[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    /// Number of dimensions.
614    pub fn dims(&self) -> usize {
615        self.min.len()
616    }
617
618    /// Returns `true` if every component of `v` falls within the trained range
619    /// `[min_d, min_d + 255 * gs]` (i.e., encoding would produce no clamping).
620    ///
621    /// When this returns `false` at least one dimension is out-of-distribution;
622    /// callers that need correctness guarantees should fall back to exact f32.
623    #[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    // ── Sq8Codec tests ──────────────────────────────────────────────────────
662
663    #[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    // ── GsSq8Codec tests ────────────────────────────────────────────────────
887
888    /// Codex counterexample (2026-06-12): ranges [0,1] and [0,1e6].
889    ///
890    /// Without global-scale, the per-dim fast path reversed near/far ordering by >6 OOM.
891    /// With GsSq8Codec the global scale is dominated by the wide dim; the narrow dim
892    /// loses code resolution but contributes proportionally little to L2 — ordering is preserved.
893    #[test]
894    fn gs_l2_sq_anisotropic_ordering_preserved() {
895        let corpus = vec![
896            vec![0.0f32, 0.0f32],    // origin
897            vec![1.0f32, 1.0f32],    // near: exact L2² = 2.0
898            vec![1.0f32, 4001.0f32], // far: exact L2² ~ 16_000_002
899        ];
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}