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