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