Skip to main content

khive_quant/
lib.rs

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