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//!
29//! See `docs/api/codecs.md` for the full function-by-function reference and
30//! `docs/design.md` for the anisotropy-gating rationale behind `GsSq8Codec`.
31
32#[cfg(feature = "parallel")]
33use rayon::prelude::*;
34
35// ─── NEON helpers ─────────────────────────────────────────────────────────────
36
37/// `Σ a_i * b_i` over equal-length `u8` slices as a `u32` accumulator (NEON on
38/// aarch64, chunked portable fallback elsewhere). See `docs/api/codecs.md`.
39///
40/// Panics if the slices have different lengths.
41#[inline(always)]
42fn u8_dot_u32(a: &[u8], b: &[u8]) -> u32 {
43    assert_eq!(a.len(), b.len(), "u8_dot_u32 inputs must have equal length");
44
45    #[cfg(target_arch = "aarch64")]
46    {
47        use std::arch::aarch64::*;
48        let n = a.len();
49        let chunks = n / 16;
50        let rem = n % 16;
51
52        let mut acc0: uint32x4_t;
53        let mut acc1: uint32x4_t;
54        let mut acc2: uint32x4_t;
55        let mut acc3: uint32x4_t;
56
57        unsafe {
58            acc0 = vdupq_n_u32(0);
59            acc1 = vdupq_n_u32(0);
60            acc2 = vdupq_n_u32(0);
61            acc3 = vdupq_n_u32(0);
62
63            for i in 0..chunks {
64                let ap = a.as_ptr().add(i * 16);
65                let bp = b.as_ptr().add(i * 16);
66
67                let va = vld1q_u8(ap);
68                let vb = vld1q_u8(bp);
69
70                let lo_u16 = vmull_u8(vget_low_u8(va), vget_low_u8(vb));
71                let hi_u16 = vmull_high_u8(va, vb);
72
73                acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
74                acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
75                acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
76                acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
77            }
78
79            let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
80            let mut total = vaddvq_u32(sum4);
81
82            for i in (n - rem)..n {
83                total += a[i] as u32 * b[i] as u32;
84            }
85            total
86        }
87    }
88
89    #[cfg(not(target_arch = "aarch64"))]
90    {
91        a.chunks(8)
92            .zip(b.chunks(8))
93            .map(|(ac, bc)| {
94                ac.iter()
95                    .zip(bc.iter())
96                    .map(|(&x, &y)| (x as u32) * (y as u32))
97                    .sum::<u32>()
98            })
99            .sum()
100    }
101}
102
103/// `Σ (a_i - b_i)²` over equal-length `u8` slices as a `u32` accumulator
104/// (NEON on aarch64, chunked portable fallback elsewhere). See
105/// `docs/api/codecs.md` for the kernel breakdown.
106///
107/// Safety: both slices must have the same length (panics otherwise).
108#[inline(always)]
109pub fn u8_l2sq_u32(a: &[u8], b: &[u8]) -> u32 {
110    assert_eq!(
111        a.len(),
112        b.len(),
113        "u8_l2sq_u32 inputs must have equal length"
114    );
115
116    #[cfg(target_arch = "aarch64")]
117    {
118        use std::arch::aarch64::*;
119        let n = a.len();
120        let chunks = n / 16;
121        let rem = n % 16;
122
123        let mut acc0: uint32x4_t;
124        let mut acc1: uint32x4_t;
125        let mut acc2: uint32x4_t;
126        let mut acc3: uint32x4_t;
127
128        unsafe {
129            acc0 = vdupq_n_u32(0);
130            acc1 = vdupq_n_u32(0);
131            acc2 = vdupq_n_u32(0);
132            acc3 = vdupq_n_u32(0);
133
134            for i in 0..chunks {
135                let ap = a.as_ptr().add(i * 16);
136                let bp = b.as_ptr().add(i * 16);
137
138                let va = vld1q_u8(ap);
139                let vb = vld1q_u8(bp);
140
141                let diff = vabdq_u8(va, vb);
142
143                let lo_u16 = vmull_u8(vget_low_u8(diff), vget_low_u8(diff));
144                let hi_u16 = vmull_high_u8(diff, diff);
145
146                acc0 = vaddq_u32(acc0, vmovl_u16(vget_low_u16(lo_u16)));
147                acc1 = vaddq_u32(acc1, vmovl_high_u16(lo_u16));
148                acc2 = vaddq_u32(acc2, vmovl_u16(vget_low_u16(hi_u16)));
149                acc3 = vaddq_u32(acc3, vmovl_high_u16(hi_u16));
150            }
151
152            let sum4 = vaddq_u32(vaddq_u32(acc0, acc1), vaddq_u32(acc2, acc3));
153            let mut total = vaddvq_u32(sum4);
154
155            for i in (n - rem)..n {
156                let d = (a[i] as i32) - (b[i] as i32);
157                total += (d * d) as u32;
158            }
159            total
160        }
161    }
162
163    #[cfg(not(target_arch = "aarch64"))]
164    {
165        a.chunks(8)
166            .zip(b.chunks(8))
167            .map(|(ac, bc)| {
168                ac.iter()
169                    .zip(bc.iter())
170                    .map(|(&x, &y)| {
171                        let d = (x as i32) - (y as i32);
172                        (d * d) as u32
173                    })
174                    .sum::<u32>()
175            })
176            .sum()
177    }
178}
179
180// ─── Validation errors (QUANT-AUD-002) ─────────────────────────────────────────
181
182/// Errors returned by the `try_*` codec constructors and encoders.
183///
184/// Public train/encode input is caller-controlled (corpus data, external
185/// vectors); shape mismatches here must be a typed error rather than a panic
186/// or a silently truncated/malformed encoded vector.
187#[derive(Debug, Clone, PartialEq)]
188pub enum QuantError {
189    /// The training corpus contained zero rows.
190    EmptyCorpus,
191    /// `dims` was zero (flat API) or the first row was empty (row API).
192    ZeroDims,
193    /// A flat vector's length was not a multiple of `dims`.
194    FlatLengthNotDivisible { len: usize, dims: usize },
195    /// A training row's length did not match the dims established by row 0.
196    RaggedRow {
197        row: usize,
198        expected: usize,
199        got: usize,
200    },
201    /// A vector passed to `encode`/`encode_flat_par` did not match the
202    /// codec's trained dims.
203    EncodeLengthMismatch { expected: usize, got: usize },
204}
205
206impl std::fmt::Display for QuantError {
207    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
208        match self {
209            Self::EmptyCorpus => write!(f, "cannot train on empty corpus"),
210            Self::ZeroDims => write!(f, "dims must be > 0"),
211            Self::FlatLengthNotDivisible { len, dims } => write!(
212                f,
213                "flat vector length {len} is not a multiple of dims {dims}"
214            ),
215            Self::RaggedRow { row, expected, got } => write!(
216                f,
217                "row {row} has length {got}, expected {expected} (dims fixed by row 0)"
218            ),
219            Self::EncodeLengthMismatch { expected, got } => write!(
220                f,
221                "vector length {got} does not match codec dims {expected}"
222            ),
223        }
224    }
225}
226
227impl std::error::Error for QuantError {}
228
229/// Per-dimension min/max over row-major flat vectors; validates `dims > 0`,
230/// non-empty corpus, `vectors.len()` a multiple of `dims`. Non-finite values
231/// are skipped. See [`finalize_min_max`] for the empty-dimension default.
232fn flat_min_max(vectors: &[f32], dims: usize) -> Result<(Vec<f32>, Vec<f32>), QuantError> {
233    if dims == 0 {
234        return Err(QuantError::ZeroDims);
235    }
236    if vectors.is_empty() {
237        return Err(QuantError::EmptyCorpus);
238    }
239    if !vectors.len().is_multiple_of(dims) {
240        return Err(QuantError::FlatLengthNotDivisible {
241            len: vectors.len(),
242            dims,
243        });
244    }
245
246    let n = vectors.len() / dims;
247    let mut min = vec![f32::INFINITY; dims];
248    let mut max = vec![f32::NEG_INFINITY; dims];
249
250    for row in 0..n {
251        let v = &vectors[row * dims..(row + 1) * dims];
252        for (d, &x) in v.iter().enumerate() {
253            if x.is_finite() {
254                if x < min[d] {
255                    min[d] = x;
256                }
257                if x > max[d] {
258                    max[d] = x;
259                }
260            }
261        }
262    }
263    finalize_min_max(&mut min, &mut max);
264    Ok((min, max))
265}
266
267/// Per-dimension min/max over row vectors; validates non-empty corpus,
268/// `dims > 0` (row 0's length), and every row matching row 0's length.
269fn row_min_max(vectors: &[Vec<f32>]) -> Result<(usize, Vec<f32>, Vec<f32>), QuantError> {
270    if vectors.is_empty() {
271        return Err(QuantError::EmptyCorpus);
272    }
273    let dims = vectors[0].len();
274    if dims == 0 {
275        return Err(QuantError::ZeroDims);
276    }
277
278    let mut min = vec![f32::INFINITY; dims];
279    let mut max = vec![f32::NEG_INFINITY; dims];
280
281    for (row, v) in vectors.iter().enumerate() {
282        if v.len() != dims {
283            return Err(QuantError::RaggedRow {
284                row,
285                expected: dims,
286                got: v.len(),
287            });
288        }
289        for (d, &x) in v.iter().enumerate() {
290            if x.is_finite() {
291                if x < min[d] {
292                    min[d] = x;
293                }
294                if x > max[d] {
295                    max[d] = x;
296                }
297            }
298        }
299    }
300    finalize_min_max(&mut min, &mut max);
301    Ok((dims, min, max))
302}
303
304/// Default a dimension with no finite observation to `min=0, max=1`; widen a
305/// degenerate `max <= min` to a unit range so `scale`/`gs` never divide by zero.
306fn finalize_min_max(min: &mut [f32], max: &mut [f32]) {
307    for d in 0..min.len() {
308        if !min[d].is_finite() {
309            min[d] = 0.0;
310        }
311        if !max[d].is_finite() || max[d] <= min[d] {
312            max[d] = min[d] + 1.0;
313        }
314    }
315}
316
317// ─── Sq8Codec (per-dim scale — dot product / cosine) ──────────────────────────
318
319/// Per-dimension affine SQ8 codec for dot product and cosine distance.
320///
321/// Encodes `f32` dimensions to `u8` via: `code = round((x - min) / scale)`,
322/// where `scale_i = (max_i - min_i) / 255`.
323///
324/// For L2 distance on Vamana builds, use [`GsSq8Codec`] instead — its global
325/// shared scale makes L2 algebraically exact in code space without a residual pass.
326#[derive(Debug, Clone)]
327pub struct Sq8Codec {
328    /// Per-dimension minimum values.
329    pub min: Vec<f32>,
330    /// Per-dimension scale: `(max - min) / 255`.
331    pub scale: Vec<f32>,
332    /// Per-dimension `scale²` precomputed for fast L2 and dot product.
333    pub scale_sq: Vec<f32>,
334    /// Mean of `scale_sq` across all dimensions — used as the integer-pass multiplier.
335    pub mean_scale_sq: f32,
336    /// Residual: `scale_sq_i - mean_scale_sq` (zero-mean, small magnitude).
337    pub scale_sq_residual: Vec<f32>,
338    /// `Σ_i min_i²` precomputed for dot-product correction.
339    pub offset_sq_sum: f32,
340}
341
342/// A corpus vector encoded by [`Sq8Codec`].
343#[derive(Debug, Clone)]
344pub struct EncodedVector {
345    /// SQ8 u8 codes, one per dimension.
346    pub codes: Vec<u8>,
347    /// L2 norm of the original f32 vector (for cosine distance).
348    pub norm: f32,
349    /// `Σ_i scale_i * min_i * code_i` — per-vector correction term for dot product.
350    pub soc_sum: f32,
351    /// `Σ_i scale_sq_residual_i * code_i` precomputed at encode time.
352    pub residual_dot_bias: f32,
353}
354
355impl Sq8Codec {
356    fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
357        let dims = min.len();
358        let scale: Vec<f32> = (0..dims).map(|d| (max[d] - min[d]) / 255.0).collect();
359        let scale_sq: Vec<f32> = scale.iter().map(|s| s * s).collect();
360        let mean_scale_sq = scale_sq.iter().sum::<f32>() / dims as f32;
361        let scale_sq_residual: Vec<f32> = scale_sq.iter().map(|&ss| ss - mean_scale_sq).collect();
362        let offset_sq_sum: f32 = min.iter().map(|o| o * o).sum();
363
364        Self {
365            min,
366            scale,
367            scale_sq,
368            mean_scale_sq,
369            scale_sq_residual,
370            offset_sq_sum,
371        }
372    }
373
374    /// Train a codec from row-major flat vectors.
375    ///
376    /// Panics on invalid input. See [`Self::try_train_flat`] for a fallible
377    /// variant that returns [`QuantError`] instead.
378    pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
379        Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
380    }
381
382    /// Fallible variant of [`Self::train_flat`]. Validates `dims > 0`, a
383    /// non-empty corpus, and that `vectors.len()` is a multiple of `dims`.
384    pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
385        let (min, max) = flat_min_max(vectors, dims)?;
386        Ok(Self::build_from_min_max(min, max))
387    }
388
389    /// Train from a slice of row vectors (each a `Vec<f32>`).
390    ///
391    /// Panics on invalid input. See [`Self::try_train`] for a fallible
392    /// variant that returns [`QuantError`] instead.
393    pub fn train(vectors: &[Vec<f32>]) -> Self {
394        Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
395    }
396
397    /// Fallible variant of [`Self::train`]. Validates a non-empty corpus,
398    /// `dims > 0` (row 0's length), and that every row is the same length
399    /// (rectangular corpus); a ragged row returns [`QuantError::RaggedRow`]
400    /// instead of panicking on out-of-bounds indexing.
401    pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
402        let (_dims, min, max) = row_min_max(vectors)?;
403        Ok(Self::build_from_min_max(min, max))
404    }
405
406    /// Encode a single vector into SQ8 codes + correction metadata.
407    ///
408    /// Panics if `v.len()` does not match the codec's trained dims. See
409    /// [`Self::try_encode`] for a fallible variant that returns
410    /// [`QuantError`] instead.
411    pub fn encode(&self, v: &[f32]) -> EncodedVector {
412        self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
413    }
414
415    /// Fallible variant of [`Self::encode`]. Validates `v.len()` against the
416    /// codec's trained dims before encoding. See `docs/design.md` (QUANT-AUD-002)
417    /// for why this check must be a typed error, not a debug-only assertion.
418    pub fn try_encode(&self, v: &[f32]) -> Result<EncodedVector, QuantError> {
419        let dims = self.min.len();
420        if v.len() != dims {
421            return Err(QuantError::EncodeLengthMismatch {
422                expected: dims,
423                got: v.len(),
424            });
425        }
426        Ok(self.encode_unchecked(v))
427    }
428
429    /// Encode a vector already validated to have `v.len() == self.min.len()`.
430    fn encode_unchecked(&self, v: &[f32]) -> EncodedVector {
431        let dims = self.min.len();
432        let mut codes = Vec::with_capacity(dims);
433        let mut soc_sum = 0.0f32;
434        let mut residual_dot_bias = 0.0f32;
435        let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
436
437        for (d, &x) in v.iter().enumerate() {
438            let s = self.scale[d];
439            let inv_s = if s > 1e-12 { 1.0 / s } else { 0.0 };
440            let raw = (x - self.min[d]) * inv_s;
441            let code = raw.round().clamp(0.0, 255.0) as u8;
442            codes.push(code);
443            soc_sum += s * self.min[d] * code as f32;
444            residual_dot_bias += self.scale_sq_residual[d] * code as f32;
445        }
446
447        EncodedVector {
448            codes,
449            norm,
450            soc_sum,
451            residual_dot_bias,
452        }
453    }
454
455    /// Encode a batch of flat-row vectors, using Rayon when the `parallel` feature is enabled.
456    ///
457    /// Panics on invalid input. See [`Self::try_encode_flat_par`] for a
458    /// fallible variant that returns [`QuantError`] instead.
459    pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<EncodedVector> {
460        self.try_encode_flat_par(vectors, dims)
461            .unwrap_or_else(|e| panic!("{e}"))
462    }
463
464    /// Fallible variant of [`Self::encode_flat_par`]. Validates `dims > 0`,
465    /// divisibility, and that `dims` matches the codec's trained dims before
466    /// dividing `vectors.len() / dims`. See `docs/design.md` (QUANT-AUD-002).
467    pub fn try_encode_flat_par(
468        &self,
469        vectors: &[f32],
470        dims: usize,
471    ) -> Result<Vec<EncodedVector>, QuantError> {
472        if dims == 0 {
473            return Err(QuantError::ZeroDims);
474        }
475        if !vectors.len().is_multiple_of(dims) {
476            return Err(QuantError::FlatLengthNotDivisible {
477                len: vectors.len(),
478                dims,
479            });
480        }
481        if dims != self.min.len() {
482            return Err(QuantError::EncodeLengthMismatch {
483                expected: self.min.len(),
484                got: dims,
485            });
486        }
487        let n = vectors.len() / dims;
488        #[cfg(feature = "parallel")]
489        let encoded = (0..n)
490            .into_par_iter()
491            .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
492            .collect();
493        #[cfg(not(feature = "parallel"))]
494        let encoded = (0..n)
495            .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
496            .collect();
497        Ok(encoded)
498    }
499
500    /// Encode a batch of row vectors, using Rayon when the `parallel` feature is enabled.
501    ///
502    /// Panics if any row's length does not match the codec's trained dims.
503    /// See [`Self::try_encode_par`] for a fallible variant that returns
504    /// [`QuantError`] instead.
505    pub fn encode_par(&self, vectors: &[Vec<f32>]) -> Vec<EncodedVector> {
506        self.try_encode_par(vectors)
507            .unwrap_or_else(|e| panic!("{e}"))
508    }
509
510    /// Fallible variant of [`Self::encode_par`]. Validates every row's length
511    /// against the codec's trained dims before dispatching to the thread pool.
512    pub fn try_encode_par(&self, vectors: &[Vec<f32>]) -> Result<Vec<EncodedVector>, QuantError> {
513        let dims = self.min.len();
514        for v in vectors {
515            if v.len() != dims {
516                return Err(QuantError::EncodeLengthMismatch {
517                    expected: dims,
518                    got: v.len(),
519                });
520            }
521        }
522        #[cfg(feature = "parallel")]
523        let encoded = vectors
524            .par_iter()
525            .map(|v| self.encode_unchecked(v))
526            .collect();
527        #[cfg(not(feature = "parallel"))]
528        let encoded = vectors.iter().map(|v| self.encode_unchecked(v)).collect();
529        Ok(encoded)
530    }
531
532    /// Approximate dot product between two encoded vectors (same codec).
533    ///
534    /// Full-precision correction identity (same min/scale for both):
535    /// `dot(a, b) = Σ s²·a·b + soc_a + soc_b + offset_sq_sum`
536    ///
537    /// The integer pass (`u8_dot_u32`) computes `raw = Σ a_i*b_i` as `u32` using
538    /// NEON (16-wide on aarch64). The scale correction then applies `mean_scale_sq`
539    /// plus a compact per-dim residual f32 pass for accuracy.
540    #[inline]
541    pub fn approx_dot(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
542        let raw = u8_dot_u32(&a.codes, &b.codes) as f32;
543        let residual_hot: f32 = self
544            .scale_sq_residual
545            .iter()
546            .zip(a.codes.iter())
547            .zip(b.codes.iter())
548            .map(|((r, &ac), &bc)| r * (ac as f32) * (bc as f32))
549            .sum();
550        self.mean_scale_sq * raw + residual_hot + a.soc_sum + b.soc_sum + self.offset_sq_sum
551    }
552
553    /// Approximate cosine distance between two encoded vectors (same codec).
554    ///
555    /// Returns `1 - dot / (norm_a * norm_b)`. Falls back to 1.0 for zero norms.
556    #[inline]
557    pub fn approx_cosine_dist(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
558        let denom = a.norm * b.norm;
559        if !denom.is_finite() || denom <= 0.0 {
560            return 1.0;
561        }
562        let dot = self.approx_dot(a, b);
563        let cosine = (dot / denom).clamp(-1.0, 1.0);
564        1.0 - cosine
565    }
566
567    /// Approximate squared L2 distance — per-dim residual corrected.
568    ///
569    /// Full-precision identity: `||a-b||² = Σ scale_sq_i * (a_i-b_i)²`.
570    /// Offsets cancel because both vectors share the same codec.
571    ///
572    /// The integer pass (`u8_l2sq_u32`) computes `raw = Σ (a_i-b_i)²` using NEON
573    /// `vabdq_u8` + `vmull_u8`. The residual correction keeps ordinal accuracy
574    /// across anisotropic corpora.
575    ///
576    /// For Vamana L2 acquisition use [`GsSq8Codec::l2_sq`] — algebraically exact
577    /// in code space and ~2× faster (no residual pass).
578    #[inline]
579    pub fn approx_l2_sq(&self, a: &EncodedVector, b: &EncodedVector) -> f32 {
580        let raw = u8_l2sq_u32(&a.codes, &b.codes) as f32;
581        let residual_hot: f32 = self
582            .scale_sq_residual
583            .iter()
584            .zip(a.codes.iter())
585            .zip(b.codes.iter())
586            .map(|((r, &ac), &bc)| {
587                let d = (ac as i32) - (bc as i32);
588                r * (d as f32) * (d as f32)
589            })
590            .sum();
591        self.mean_scale_sq * raw + residual_hot
592    }
593
594    /// Number of dimensions.
595    pub fn dims(&self) -> usize {
596        self.min.len()
597    }
598}
599
600// ─── GsSq8Codec (global-scale — L2 / Vamana acquisition) ─────────────────────
601
602/// Global-scale SQ8 codec for L2 distance — the Vamana acquisition path.
603///
604/// A single shared scale `gs = max_range_across_dims / 255` is used for all
605/// dims; per-dim offsets are still subtracted before quantizing. Encoding is
606/// **lossy** (rounded + clamped to u8); L2² in code space is exact after that
607/// lossy encode, but round-trip error vs. true f32 L2² can reach ~15% for
608/// anisotropic/OOD data — no residual pass, no gate, no silent fallback.
609/// Callers needing correctness on OOD queries must check
610/// [`Self::is_in_distribution`] and fall back to exact f32 themselves.
611/// See `docs/api/codecs.md` for the full accuracy discussion and
612/// `docs/design.md` for why this replaced the earlier per-dim anisotropy-gated design.
613#[derive(Debug, Clone)]
614pub struct GsSq8Codec {
615    /// Per-dimension minimum values.
616    pub min: Vec<f32>,
617    /// Global scale: `max_range / 255` where `max_range = max_i(max_i - min_i)`.
618    pub gs: f32,
619    /// `gs²` precomputed for L2.
620    pub gs_sq: f32,
621    /// Anisotropy ratio measured at train time: `max(range_i) / min(nonzero range_i)`.
622    /// Informational only — never used for dispatch decisions.
623    pub anisotropy_ratio: f32,
624}
625
626/// A corpus vector encoded by [`GsSq8Codec`].
627#[derive(Debug, Clone)]
628pub struct GsEncodedVector {
629    /// SQ8 u8 codes, one per dimension.
630    pub codes: Vec<u8>,
631}
632
633impl GsSq8Codec {
634    fn build_from_min_max(min: Vec<f32>, max: Vec<f32>) -> Self {
635        let dims = min.len();
636        let ranges: Vec<f32> = (0..dims).map(|d| max[d] - min[d]).collect();
637        let max_range = ranges.iter().cloned().fold(0.0f32, f32::max);
638        let gs = if max_range > 1e-12 {
639            max_range / 255.0
640        } else {
641            1.0 / 255.0
642        };
643
644        let min_range_nonzero = ranges
645            .iter()
646            .cloned()
647            .filter(|&r| r > 1e-12)
648            .fold(f32::INFINITY, f32::min);
649        let anisotropy_ratio = if min_range_nonzero.is_finite() && min_range_nonzero > 0.0 {
650            max_range / min_range_nonzero
651        } else {
652            1.0
653        };
654
655        Self {
656            min,
657            gs,
658            gs_sq: gs * gs,
659            anisotropy_ratio,
660        }
661    }
662
663    /// Train from row-major flat vectors.
664    ///
665    /// Panics on invalid input. See [`Self::try_train_flat`] for a fallible
666    /// variant that returns [`QuantError`] instead.
667    pub fn train_flat(vectors: &[f32], dims: usize) -> Self {
668        Self::try_train_flat(vectors, dims).unwrap_or_else(|e| panic!("{e}"))
669    }
670
671    /// Fallible variant of [`Self::train_flat`]. Validates `dims > 0`, a
672    /// non-empty corpus, and that `vectors.len()` is a multiple of `dims`.
673    pub fn try_train_flat(vectors: &[f32], dims: usize) -> Result<Self, QuantError> {
674        let (min, max) = flat_min_max(vectors, dims)?;
675        Ok(Self::build_from_min_max(min, max))
676    }
677
678    /// Train from a slice of row vectors.
679    ///
680    /// Panics on invalid input. See [`Self::try_train`] for a fallible
681    /// variant that returns [`QuantError`] instead.
682    pub fn train(vectors: &[Vec<f32>]) -> Self {
683        Self::try_train(vectors).unwrap_or_else(|e| panic!("{e}"))
684    }
685
686    /// Fallible variant of [`Self::train`]. Validates a non-empty corpus,
687    /// `dims > 0` (row 0's length), and that every row is the same length
688    /// (rectangular corpus); a ragged row returns [`QuantError::RaggedRow`]
689    /// instead of panicking on out-of-bounds indexing.
690    pub fn try_train(vectors: &[Vec<f32>]) -> Result<Self, QuantError> {
691        let (_dims, min, max) = row_min_max(vectors)?;
692        Ok(Self::build_from_min_max(min, max))
693    }
694
695    /// Encode a single vector.
696    ///
697    /// Panics if `v.len()` does not match the codec's trained dims. See
698    /// [`Self::try_encode`] for a fallible variant that returns
699    /// [`QuantError`] instead.
700    #[inline]
701    pub fn encode(&self, v: &[f32]) -> GsEncodedVector {
702        self.try_encode(v).unwrap_or_else(|e| panic!("{e}"))
703    }
704
705    /// Fallible variant of [`Self::encode`]. Validates `v.len()` against the
706    /// codec's trained dims before encoding — an unchecked shape mismatch
707    /// (e.g. `v = &[]`) could otherwise score as a false exact match; see
708    /// `docs/design.md` (QUANT-AUD-002).
709    pub fn try_encode(&self, v: &[f32]) -> Result<GsEncodedVector, QuantError> {
710        let dims = self.min.len();
711        if v.len() != dims {
712            return Err(QuantError::EncodeLengthMismatch {
713                expected: dims,
714                got: v.len(),
715            });
716        }
717        Ok(self.encode_unchecked(v))
718    }
719
720    /// Encode a vector already validated to have `v.len() == self.min.len()`.
721    #[inline]
722    fn encode_unchecked(&self, v: &[f32]) -> GsEncodedVector {
723        let inv_gs = if self.gs > 1e-12 { 1.0 / self.gs } else { 0.0 };
724        let codes = v
725            .iter()
726            .enumerate()
727            .map(|(d, &x)| ((x - self.min[d]) * inv_gs).round().clamp(0.0, 255.0) as u8)
728            .collect();
729        GsEncodedVector { codes }
730    }
731
732    /// Encode a batch of flat-row vectors, using Rayon when the `parallel` feature is enabled.
733    ///
734    /// Panics on invalid input. See [`Self::try_encode_flat_par`] for a
735    /// fallible variant that returns [`QuantError`] instead.
736    pub fn encode_flat_par(&self, vectors: &[f32], dims: usize) -> Vec<GsEncodedVector> {
737        self.try_encode_flat_par(vectors, dims)
738            .unwrap_or_else(|e| panic!("{e}"))
739    }
740
741    /// Fallible variant of [`Self::encode_flat_par`]. Validates `dims > 0`,
742    /// divisibility, and that `dims` matches the codec's trained dims before
743    /// dividing `vectors.len() / dims`. See `docs/design.md` (QUANT-AUD-002).
744    pub fn try_encode_flat_par(
745        &self,
746        vectors: &[f32],
747        dims: usize,
748    ) -> Result<Vec<GsEncodedVector>, QuantError> {
749        if dims == 0 {
750            return Err(QuantError::ZeroDims);
751        }
752        if !vectors.len().is_multiple_of(dims) {
753            return Err(QuantError::FlatLengthNotDivisible {
754                len: vectors.len(),
755                dims,
756            });
757        }
758        if dims != self.min.len() {
759            return Err(QuantError::EncodeLengthMismatch {
760                expected: self.min.len(),
761                got: dims,
762            });
763        }
764        let n = vectors.len() / dims;
765        #[cfg(feature = "parallel")]
766        let encoded = (0..n)
767            .into_par_iter()
768            .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
769            .collect();
770        #[cfg(not(feature = "parallel"))]
771        let encoded = (0..n)
772            .map(|i| self.encode_unchecked(&vectors[i * dims..(i + 1) * dims]))
773            .collect();
774        Ok(encoded)
775    }
776
777    /// Approximate squared L2 distance.
778    ///
779    /// `||a-b||² ≈ gs² × Σ (a_i - b_i)²`
780    ///
781    /// Exact in code space (offset terms cancel, `gs²` factorizes) after the
782    /// lossy f32→u8 encode. Per-round-trip L2 error can reach ~15%; recall
783    /// safety is established by probe, not by this formula.
784    /// The NEON path runs ~13 ns at 384-d.
785    #[inline]
786    pub fn l2_sq(&self, a: &GsEncodedVector, b: &GsEncodedVector) -> f32 {
787        self.gs_sq * u8_l2sq_u32(&a.codes, &b.codes) as f32
788    }
789
790    /// [`Self::l2_sq`] over raw code slices, for callers that store codes in a
791    /// flat (possibly memory-mapped) buffer rather than per-vector allocations.
792    #[inline]
793    pub fn l2_sq_codes(&self, a: &[u8], b: &[u8]) -> f32 {
794        self.gs_sq * u8_l2sq_u32(a, b) as f32
795    }
796
797    /// Number of dimensions.
798    pub fn dims(&self) -> usize {
799        self.min.len()
800    }
801
802    /// Returns `true` if every component of `v` falls within the trained range
803    /// `[min_d, min_d + 255 * gs]` (i.e., encoding would produce no clamping).
804    ///
805    /// When this returns `false` at least one dimension is out-of-distribution;
806    /// callers that need correctness guarantees should fall back to exact f32.
807    #[inline]
808    pub fn is_in_distribution(&self, v: &[f32]) -> bool {
809        let max_code = 255.0 * self.gs;
810        v.iter()
811            .zip(self.min.iter())
812            .all(|(&x, &mn)| x >= mn && x <= mn + max_code)
813    }
814}
815
816#[cfg(test)]
817mod tests {
818    use super::*;
819
820    fn rand_vecs(n: usize, dims: usize, seed: u64) -> Vec<Vec<f32>> {
821        let mut h = seed;
822        (0..n)
823            .map(|_| {
824                (0..dims)
825                    .map(|_| {
826                        h = h
827                            .wrapping_mul(0x6c62_272e_07bb_0142)
828                            .wrapping_add(0x62b8_2175_62d9_6b1a);
829                        let bits = (h >> 33) as u32;
830                        (bits as f32) / (u32::MAX as f32) * 2.0 - 1.0
831                    })
832                    .collect()
833            })
834            .collect()
835    }
836
837    fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
838        a.iter().zip(b).map(|(x, y)| x * y).sum()
839    }
840
841    fn l2_sq_f32(a: &[f32], b: &[f32]) -> f32 {
842        a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
843    }
844
845    // ── QUANT-AUD-002: validation regression tests ──────────────────────────
846
847    #[test]
848    fn sq8_try_train_ragged_rows_returns_error_not_panic() {
849        let vecs = vec![vec![0.0], vec![1.0, 2.0]];
850        let err = Sq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
851        assert_eq!(
852            err,
853            QuantError::RaggedRow {
854                row: 1,
855                expected: 1,
856                got: 2
857            }
858        );
859    }
860
861    #[test]
862    fn gs_try_train_ragged_rows_returns_error_not_panic() {
863        let vecs = vec![vec![0.0], vec![1.0, 2.0]];
864        let err = GsSq8Codec::try_train(&vecs).expect_err("ragged rows must be rejected");
865        assert_eq!(
866            err,
867            QuantError::RaggedRow {
868                row: 1,
869                expected: 1,
870                got: 2
871            }
872        );
873    }
874
875    #[test]
876    fn sq8_try_train_empty_corpus_returns_error() {
877        let vecs: Vec<Vec<f32>> = vec![];
878        assert_eq!(
879            Sq8Codec::try_train(&vecs).unwrap_err(),
880            QuantError::EmptyCorpus
881        );
882    }
883
884    #[test]
885    fn sq8_try_train_flat_zero_dims_returns_error() {
886        assert_eq!(
887            Sq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
888            QuantError::ZeroDims
889        );
890    }
891
892    #[test]
893    fn gs_try_train_flat_zero_dims_returns_error() {
894        assert_eq!(
895            GsSq8Codec::try_train_flat(&[1.0, 2.0], 0).unwrap_err(),
896            QuantError::ZeroDims
897        );
898    }
899
900    #[test]
901    fn sq8_try_train_flat_remainder_returns_error() {
902        // 5 elements is not a multiple of dims=2.
903        let err = Sq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
904        assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
905    }
906
907    #[test]
908    fn gs_try_train_flat_remainder_returns_error() {
909        let err = GsSq8Codec::try_train_flat(&[1.0, 2.0, 3.0, 4.0, 5.0], 2).unwrap_err();
910        assert_eq!(err, QuantError::FlatLengthNotDivisible { len: 5, dims: 2 });
911    }
912
913    #[test]
914    fn sq8_try_encode_shorter_input_returns_error_not_malformed_vector() {
915        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
916        let err = codec.try_encode(&[0.0, 1.0]).unwrap_err();
917        assert_eq!(
918            err,
919            QuantError::EncodeLengthMismatch {
920                expected: 4,
921                got: 2
922            }
923        );
924    }
925
926    #[test]
927    fn sq8_try_encode_longer_input_returns_error() {
928        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
929        let err = codec.try_encode(&[0.0, 1.0, 2.0, 3.0, 4.0]).unwrap_err();
930        assert_eq!(
931            err,
932            QuantError::EncodeLengthMismatch {
933                expected: 4,
934                got: 5
935            }
936        );
937    }
938
939    #[test]
940    fn gs_try_encode_empty_input_returns_error_not_malformed_vector() {
941        // QUANT-AUD-002: previously encode(&[]) on a trained (dims=1) codec
942        // silently returned an EMPTY code vector in release builds (the
943        // length check was debug_assert-only), which is_in_distribution
944        // vacuously accepted and l2_sq scored as 0.0. Must now be a typed
945        // error, never a malformed zero-length code vector.
946        let codec = GsSq8Codec::train_flat(&[1.0, 2.0, 3.0, 4.0], 1);
947        let err = codec.try_encode(&[]).unwrap_err();
948        assert_eq!(
949            err,
950            QuantError::EncodeLengthMismatch {
951                expected: 1,
952                got: 0
953            }
954        );
955    }
956
957    #[test]
958    fn sq8_try_encode_flat_par_zero_dims_returns_error() {
959        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
960        assert_eq!(
961            codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
962            QuantError::ZeroDims
963        );
964    }
965
966    #[test]
967    fn gs_try_encode_flat_par_zero_dims_returns_error() {
968        let codec = GsSq8Codec::train(&rand_vecs(10, 4, 1));
969        assert_eq!(
970            codec.try_encode_flat_par(&[1.0, 2.0], 0).unwrap_err(),
971            QuantError::ZeroDims
972        );
973    }
974
975    #[test]
976    fn sq8_try_encode_flat_par_remainder_returns_error() {
977        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
978        let flat: Vec<f32> = (0..9).map(|i| i as f32).collect(); // 9 not a multiple of 4
979        assert_eq!(
980            codec.try_encode_flat_par(&flat, 4).unwrap_err(),
981            QuantError::FlatLengthNotDivisible { len: 9, dims: 4 }
982        );
983    }
984
985    #[test]
986    fn sq8_try_encode_flat_par_dims_mismatch_returns_error() {
987        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
988        let flat: Vec<f32> = (0..6).map(|i| i as f32).collect();
989        let err = codec.try_encode_flat_par(&flat, 3).unwrap_err();
990        assert_eq!(
991            err,
992            QuantError::EncodeLengthMismatch {
993                expected: 4,
994                got: 3
995            }
996        );
997    }
998
999    #[test]
1000    #[should_panic(expected = "row 1 has length 2, expected 1")]
1001    fn sq8_train_still_panics_with_typed_message_on_ragged_rows() {
1002        // The panicking convenience wrapper is preserved for existing callers,
1003        // but must now surface the typed QuantError message rather than
1004        // panicking from a raw out-of-bounds index.
1005        let _ = Sq8Codec::train(&[vec![0.0], vec![1.0, 2.0]]);
1006    }
1007
1008    #[test]
1009    fn sq8_try_train_and_encode_roundtrip_matches_panicking_api() {
1010        let vecs = rand_vecs(20, 8, 7);
1011        let a = Sq8Codec::train(&vecs);
1012        let b = Sq8Codec::try_train(&vecs).expect("valid corpus must train");
1013        assert_eq!(a.min, b.min);
1014        assert_eq!(a.scale, b.scale);
1015        let ea = a.encode(&vecs[0]);
1016        let eb = b.try_encode(&vecs[0]).expect("valid vector must encode");
1017        assert_eq!(ea.codes, eb.codes);
1018    }
1019
1020    // ── Sq8Codec tests ──────────────────────────────────────────────────────
1021
1022    #[test]
1023    fn encode_decode_roundtrip_is_bounded() {
1024        let vecs = rand_vecs(100, 32, 42);
1025        let codec = Sq8Codec::train(&vecs);
1026        for v in &vecs {
1027            let ev = codec.encode(v);
1028            assert_eq!(ev.codes.len(), v.len());
1029            for (d, &code) in ev.codes.iter().enumerate() {
1030                let decoded = code as f32 * codec.scale[d] + codec.min[d];
1031                let err = (decoded - v[d]).abs();
1032                assert!(
1033                    err <= codec.scale[d] + 1e-5,
1034                    "dim {d}: err={err} scale={}",
1035                    codec.scale[d]
1036                );
1037            }
1038        }
1039    }
1040
1041    #[test]
1042    fn approx_dot_relative_error_bounded() {
1043        let vecs = rand_vecs(200, 64, 77);
1044        let codec = Sq8Codec::train(&vecs);
1045        let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1046
1047        let mut max_rel_err = 0.0f32;
1048        for i in 0..vecs.len() {
1049            for j in (i + 1)..vecs.len().min(i + 10) {
1050                let true_dot = dot_f32(&vecs[i], &vecs[j]);
1051                let approx = codec.approx_dot(&encoded[i], &encoded[j]);
1052                let denom = true_dot.abs().max(1e-3);
1053                let rel = (approx - true_dot).abs() / denom;
1054                if rel > max_rel_err {
1055                    max_rel_err = rel;
1056                }
1057            }
1058        }
1059        assert!(
1060            max_rel_err < 0.15,
1061            "max relative dot error {max_rel_err:.4} >= 0.15"
1062        );
1063    }
1064
1065    #[test]
1066    fn approx_l2_sq_relative_error_bounded() {
1067        let vecs = rand_vecs(200, 64, 88);
1068        let codec = Sq8Codec::train(&vecs);
1069        let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1070
1071        let mut max_rel_err = 0.0f32;
1072        for i in 0..vecs.len() {
1073            for j in (i + 1)..vecs.len().min(i + 10) {
1074                let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1075                let approx = codec.approx_l2_sq(&encoded[i], &encoded[j]);
1076                let denom = true_l2.max(1e-6);
1077                let rel = (approx - true_l2).abs() / denom;
1078                if rel > max_rel_err {
1079                    max_rel_err = rel;
1080                }
1081            }
1082        }
1083        assert!(
1084            max_rel_err < 0.15,
1085            "max relative L2² error {max_rel_err:.4} >= 0.15"
1086        );
1087    }
1088
1089    #[test]
1090    fn order_preservation_triplets_cosine() {
1091        let vecs = rand_vecs(300, 64, 99);
1092        let codec = Sq8Codec::train(&vecs);
1093        let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1094
1095        let n = vecs.len();
1096        let mut agree = 0usize;
1097        let mut total = 0usize;
1098
1099        for anchor in 0..50 {
1100            let a = &vecs[anchor];
1101            let ea = &encoded[anchor];
1102            for b_idx in 0..n {
1103                for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1104                    let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1105                    let norm_b: f32 = vecs[b_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1106                    let norm_c: f32 = vecs[c_idx].iter().map(|x| x * x).sum::<f32>().sqrt();
1107
1108                    let cos_ab = dot_f32(a, &vecs[b_idx]) / (norm_a * norm_b).max(1e-9);
1109                    let cos_ac = dot_f32(a, &vecs[c_idx]) / (norm_a * norm_c).max(1e-9);
1110                    let dist_ab_true = 1.0 - cos_ab;
1111                    let dist_ac_true = 1.0 - cos_ac;
1112
1113                    let dist_ab_approx = codec.approx_cosine_dist(ea, &encoded[b_idx]);
1114                    let dist_ac_approx = codec.approx_cosine_dist(ea, &encoded[c_idx]);
1115
1116                    if (dist_ab_true - dist_ac_true).abs() < 0.01 {
1117                        continue;
1118                    }
1119
1120                    let true_closer_b = dist_ab_true < dist_ac_true;
1121                    let approx_closer_b = dist_ab_approx < dist_ac_approx;
1122                    if true_closer_b == approx_closer_b {
1123                        agree += 1;
1124                    }
1125                    total += 1;
1126                }
1127            }
1128        }
1129
1130        let rate = agree as f64 / total.max(1) as f64;
1131        assert!(
1132            rate >= 0.95,
1133            "order preservation {rate:.3} < 0.95 ({agree}/{total})"
1134        );
1135    }
1136
1137    #[test]
1138    fn order_preservation_triplets_l2() {
1139        let vecs = rand_vecs(300, 64, 101);
1140        let codec = Sq8Codec::train(&vecs);
1141        let encoded: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1142
1143        let n = vecs.len();
1144        let mut agree = 0usize;
1145        let mut total = 0usize;
1146
1147        for anchor in 0..50 {
1148            let a = &vecs[anchor];
1149            let ea = &encoded[anchor];
1150            for b_idx in 0..n {
1151                for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1152                    let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1153                    let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1154
1155                    let dist_ab_approx = codec.approx_l2_sq(ea, &encoded[b_idx]);
1156                    let dist_ac_approx = codec.approx_l2_sq(ea, &encoded[c_idx]);
1157
1158                    if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1159                        continue;
1160                    }
1161
1162                    let true_closer_b = dist_ab_true < dist_ac_true;
1163                    let approx_closer_b = dist_ab_approx < dist_ac_approx;
1164                    if true_closer_b == approx_closer_b {
1165                        agree += 1;
1166                    }
1167                    total += 1;
1168                }
1169            }
1170        }
1171
1172        let rate = agree as f64 / total.max(1) as f64;
1173        assert!(
1174            rate >= 0.95,
1175            "L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1176        );
1177    }
1178
1179    #[test]
1180    fn train_flat_matches_train_rows() {
1181        let vecs = rand_vecs(50, 16, 123);
1182        let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1183
1184        let codec_rows = Sq8Codec::train(&vecs);
1185        let codec_flat = Sq8Codec::train_flat(&flat, 16);
1186
1187        for d in 0..16 {
1188            assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1189            assert!((codec_rows.scale[d] - codec_flat.scale[d]).abs() < 1e-6);
1190        }
1191    }
1192
1193    #[test]
1194    fn encode_par_matches_sequential() {
1195        let vecs = rand_vecs(50, 32, 555);
1196        let codec = Sq8Codec::train(&vecs);
1197
1198        let seq: Vec<EncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1199        let par = codec.encode_par(&vecs);
1200
1201        assert_eq!(seq.len(), par.len());
1202        for (s, p) in seq.iter().zip(par.iter()) {
1203            assert_eq!(s.codes, p.codes);
1204            assert!((s.soc_sum - p.soc_sum).abs() < 1e-5);
1205        }
1206    }
1207
1208    #[test]
1209    fn sq8_try_encode_par_short_row_returns_error_not_panic() {
1210        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1211        let mut rows = rand_vecs(5, 4, 2);
1212        rows[3] = vec![0.0, 1.0];
1213        let err = codec.try_encode_par(&rows).unwrap_err();
1214        assert_eq!(
1215            err,
1216            QuantError::EncodeLengthMismatch {
1217                expected: 4,
1218                got: 2
1219            }
1220        );
1221    }
1222
1223    #[test]
1224    fn sq8_try_encode_par_long_row_returns_error_not_panic() {
1225        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1226        let mut rows = rand_vecs(5, 4, 2);
1227        rows[3] = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
1228        let err = codec.try_encode_par(&rows).unwrap_err();
1229        assert_eq!(
1230            err,
1231            QuantError::EncodeLengthMismatch {
1232                expected: 4,
1233                got: 6
1234            }
1235        );
1236    }
1237
1238    #[test]
1239    #[should_panic(expected = "vector length 2 does not match codec dims 4")]
1240    fn sq8_encode_par_still_panics_with_typed_message_on_short_row() {
1241        let codec = Sq8Codec::train(&rand_vecs(10, 4, 1));
1242        let mut rows = rand_vecs(5, 4, 2);
1243        rows[3] = vec![0.0, 1.0];
1244        let _ = codec.encode_par(&rows);
1245    }
1246
1247    #[test]
1248    fn u8_dot_u32_matches_scalar() {
1249        let a: Vec<u8> = (0u8..=255).take(384).collect();
1250        let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1251        let scalar: u32 = a
1252            .iter()
1253            .zip(b.iter())
1254            .map(|(&x, &y)| x as u32 * y as u32)
1255            .sum();
1256        assert_eq!(u8_dot_u32(&a, &b), scalar, "u8_dot_u32 mismatch");
1257    }
1258
1259    #[test]
1260    #[should_panic(expected = "u8_dot_u32 inputs must have equal length")]
1261    fn u8_dot_u32_rejects_shorter_second_slice() {
1262        let a = [1u8; 16];
1263        let b = [2u8; 1];
1264
1265        let _ = u8_dot_u32(&a, &b);
1266    }
1267
1268    #[test]
1269    #[should_panic(expected = "u8_dot_u32 inputs must have equal length")]
1270    fn u8_dot_u32_rejects_longer_second_slice() {
1271        let a = [1u8; 1];
1272        let b = [2u8; 16];
1273
1274        let _ = u8_dot_u32(&a, &b);
1275    }
1276
1277    #[test]
1278    #[should_panic(expected = "u8_dot_u32 inputs must have equal length")]
1279    fn approx_dot_rejects_mismatched_codes() {
1280        let vectors = rand_vecs(2, 16, 1);
1281        let codec = Sq8Codec::train(&vectors);
1282        let a = codec.encode(&vectors[0]);
1283        let mut b = codec.encode(&vectors[1]);
1284        b.codes.pop();
1285
1286        let _ = codec.approx_dot(&a, &b);
1287    }
1288
1289    #[test]
1290    fn u8_helpers_tail_path_max_diff() {
1291        for len in [1usize, 7, 15, 17, 100, 383] {
1292            let a = vec![255u8; len];
1293            let b = vec![0u8; len];
1294            assert_eq!(u8_l2sq_u32(&a, &b), len as u32 * 255 * 255, "l2 len={len}");
1295            assert_eq!(u8_dot_u32(&a, &a), len as u32 * 255 * 255, "dot len={len}");
1296        }
1297    }
1298
1299    #[test]
1300    fn u8_l2sq_u32_matches_scalar() {
1301        let a: Vec<u8> = (0u8..=255).take(384).collect();
1302        let b: Vec<u8> = (0u8..=255).rev().take(384).collect();
1303        let scalar: u32 = a
1304            .iter()
1305            .zip(b.iter())
1306            .map(|(&x, &y)| {
1307                let d = (x as i32) - (y as i32);
1308                (d * d) as u32
1309            })
1310            .sum();
1311        assert_eq!(u8_l2sq_u32(&a, &b), scalar, "u8_l2sq_u32 mismatch");
1312    }
1313
1314    #[test]
1315    #[should_panic(expected = "u8_l2sq_u32 inputs must have equal length")]
1316    fn u8_l2sq_u32_rejects_shorter_second_slice() {
1317        let a = [1u8; 16];
1318        let b = [2u8; 1];
1319
1320        let _ = u8_l2sq_u32(&a, &b);
1321    }
1322
1323    // ── GsSq8Codec tests ────────────────────────────────────────────────────
1324
1325    /// Regression counterexample (2026-06-12): ranges [0,1] and [0,1e6].
1326    ///
1327    /// Without global-scale, the per-dim fast path reversed near/far ordering by >6 OOM.
1328    /// With GsSq8Codec the global scale is dominated by the wide dim; the narrow dim
1329    /// loses code resolution but contributes proportionally little to L2 — ordering is preserved.
1330    #[test]
1331    fn gs_l2_sq_anisotropic_ordering_preserved() {
1332        let corpus = vec![
1333            vec![0.0f32, 0.0f32],    // origin
1334            vec![1.0f32, 1.0f32],    // near: exact L2² = 2.0
1335            vec![1.0f32, 4001.0f32], // far: exact L2² ~ 16_000_002
1336        ];
1337        let codec = GsSq8Codec::train(&corpus);
1338
1339        let enc_origin = codec.encode(&corpus[0]);
1340        let enc_near = codec.encode(&corpus[1]);
1341        let enc_far = codec.encode(&corpus[2]);
1342
1343        let d_near = codec.l2_sq(&enc_origin, &enc_near);
1344        let d_far = codec.l2_sq(&enc_origin, &enc_far);
1345
1346        assert!(
1347            d_near < d_far,
1348            "GsSq8Codec reversed near/far on anisotropic corpus: near={d_near} far={d_far} \
1349             (anisotropy_ratio={:.1})",
1350            codec.anisotropy_ratio
1351        );
1352    }
1353
1354    #[test]
1355    fn gs_l2_sq_isotropic_small_error() {
1356        let vecs = rand_vecs(200, 64, 202);
1357        let codec = GsSq8Codec::train(&vecs);
1358        let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1359
1360        let mut max_rel = 0.0f32;
1361        for i in 0..vecs.len() {
1362            for j in (i + 1)..vecs.len().min(i + 10) {
1363                let true_l2 = l2_sq_f32(&vecs[i], &vecs[j]);
1364                let approx = codec.l2_sq(&encoded[i], &encoded[j]);
1365                let denom = true_l2.max(1e-6);
1366                let rel = (approx - true_l2).abs() / denom;
1367                if rel > max_rel {
1368                    max_rel = rel;
1369                }
1370            }
1371        }
1372        assert!(
1373            max_rel < 0.15,
1374            "GsSq8Codec max relative L2² error {max_rel:.4} >= 0.15"
1375        );
1376    }
1377
1378    #[test]
1379    fn gs_train_flat_matches_train_rows() {
1380        let vecs = rand_vecs(50, 16, 321);
1381        let flat: Vec<f32> = vecs.iter().flatten().copied().collect();
1382
1383        let codec_rows = GsSq8Codec::train(&vecs);
1384        let codec_flat = GsSq8Codec::train_flat(&flat, 16);
1385
1386        assert!((codec_rows.gs - codec_flat.gs).abs() < 1e-7);
1387        for d in 0..16 {
1388            assert!((codec_rows.min[d] - codec_flat.min[d]).abs() < 1e-6);
1389        }
1390    }
1391
1392    #[test]
1393    fn gs_l2_sq_order_preservation_triplets() {
1394        let vecs = rand_vecs(300, 64, 303);
1395        let codec = GsSq8Codec::train(&vecs);
1396        let encoded: Vec<GsEncodedVector> = vecs.iter().map(|v| codec.encode(v)).collect();
1397
1398        let n = vecs.len();
1399        let mut agree = 0usize;
1400        let mut total = 0usize;
1401
1402        for anchor in 0..50 {
1403            let a = &vecs[anchor];
1404            let ea = &encoded[anchor];
1405            for b_idx in 0..n {
1406                for c_idx in (b_idx + 1)..n.min(b_idx + 5) {
1407                    let dist_ab_true = l2_sq_f32(a, &vecs[b_idx]);
1408                    let dist_ac_true = l2_sq_f32(a, &vecs[c_idx]);
1409
1410                    let dist_ab_approx = codec.l2_sq(ea, &encoded[b_idx]);
1411                    let dist_ac_approx = codec.l2_sq(ea, &encoded[c_idx]);
1412
1413                    if (dist_ab_true - dist_ac_true).abs() < 0.001 {
1414                        continue;
1415                    }
1416
1417                    let true_closer_b = dist_ab_true < dist_ac_true;
1418                    let approx_closer_b = dist_ab_approx < dist_ac_approx;
1419                    if true_closer_b == approx_closer_b {
1420                        agree += 1;
1421                    }
1422                    total += 1;
1423                }
1424            }
1425        }
1426
1427        let rate = agree as f64 / total.max(1) as f64;
1428        assert!(
1429            rate >= 0.95,
1430            "GsSq8Codec L2 order preservation {rate:.3} < 0.95 ({agree}/{total})"
1431        );
1432    }
1433}