Skip to main content

nodedb_vector/
quantize.rs

1//! Scalar Quantization (SQ8): FP32 → INT8 per-dimension.
2//!
3//! Each dimension is independently quantized to `[0, 255]` using per-dimension
4//! min/max calibration. Provides 4x RAM reduction with <1% recall loss.
5//!
6//! Distance computation uses asymmetric mode: query stays in FP32,
7//! candidates are in INT8. This avoids quantizing the query and
8//! preserves accuracy.
9//!
10//! Storage: D bytes per vector (vs 4D bytes for FP32).
11
12use serde::{Deserialize, Serialize};
13
14/// SQ8 calibration parameters: per-dimension min/max.
15#[derive(Clone, Serialize, Deserialize)]
16pub struct Sq8Codec {
17    dim: usize,
18    /// Per-dimension minimum observed during calibration.
19    mins: Vec<f32>,
20    /// Per-dimension maximum observed during calibration.
21    maxs: Vec<f32>,
22    /// Pre-computed per-dimension scale: `(max - min) / 255.0`.
23    scales: Vec<f32>,
24    /// Pre-computed per-dimension inverse scale: `255.0 / (max - min)`.
25    inv_scales: Vec<f32>,
26}
27
28impl Sq8Codec {
29    /// Calibrate min/max from a set of training vectors.
30    ///
31    /// Scans all vectors to find per-dimension min/max bounds.
32    /// At least 1000 vectors recommended for stable calibration.
33    pub fn calibrate(vectors: &[&[f32]], dim: usize) -> Self {
34        assert!(!vectors.is_empty(), "cannot calibrate on empty set");
35        assert!(dim > 0);
36
37        let mut mins = vec![f32::MAX; dim];
38        let mut maxs = vec![f32::MIN; dim];
39
40        for v in vectors {
41            debug_assert_eq!(v.len(), dim);
42            for d in 0..dim {
43                if v[d] < mins[d] {
44                    mins[d] = v[d];
45                }
46                if v[d] > maxs[d] {
47                    maxs[d] = v[d];
48                }
49            }
50        }
51
52        let mut scales = vec![0.0f32; dim];
53        let mut inv_scales = vec![0.0f32; dim];
54        for d in 0..dim {
55            let range = maxs[d] - mins[d];
56            if range > f32::EPSILON {
57                scales[d] = range / 255.0;
58                inv_scales[d] = 255.0 / range;
59            }
60        }
61
62        Self {
63            dim,
64            mins,
65            maxs,
66            scales,
67            inv_scales,
68        }
69    }
70
71    /// Quantize a single FP32 vector to INT8.
72    pub fn quantize(&self, vector: &[f32]) -> Vec<u8> {
73        debug_assert_eq!(vector.len(), self.dim);
74        let mut out = Vec::with_capacity(self.dim);
75        for ((&v, &min), (&max, &inv_scale)) in vector
76            .iter()
77            .zip(self.mins.iter())
78            .zip(self.maxs.iter().zip(self.inv_scales.iter()))
79        {
80            let clamped = v.clamp(min, max);
81            let q = ((clamped - min) * inv_scale).round() as u8;
82            out.push(q);
83        }
84        out
85    }
86
87    /// Batch quantize: quantize all vectors into a contiguous byte array.
88    ///
89    /// Returns `dim * N` bytes laid out as `[v0_d0, v0_d1, ..., v1_d0, ...]`.
90    pub fn quantize_batch(&self, vectors: &[&[f32]]) -> Vec<u8> {
91        let mut out = Vec::with_capacity(self.dim * vectors.len());
92        for v in vectors {
93            out.extend(self.quantize(v));
94        }
95        out
96    }
97
98    /// Dequantize INT8 back to FP32 (lossy reconstruction).
99    pub fn dequantize(&self, quantized: &[u8]) -> Vec<f32> {
100        debug_assert_eq!(quantized.len(), self.dim);
101        let mut out = Vec::with_capacity(self.dim);
102        for ((&q, &min), &scale) in quantized
103            .iter()
104            .zip(self.mins.iter())
105            .zip(self.scales.iter())
106        {
107            out.push(min + q as f32 * scale);
108        }
109        out
110    }
111
112    /// Asymmetric L2 squared distance: query (FP32) vs candidate (INT8).
113    ///
114    /// This is the hot-path function used during HNSW traversal.
115    #[inline]
116    pub fn asymmetric_l2(&self, query: &[f32], candidate: &[u8]) -> f32 {
117        debug_assert_eq!(query.len(), self.dim);
118        debug_assert_eq!(candidate.len(), self.dim);
119        let mut sum = 0.0f32;
120        for d in 0..self.dim {
121            let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
122            let diff = query[d] - dequant;
123            sum += diff * diff;
124        }
125        sum
126    }
127
128    /// Asymmetric cosine distance: query (FP32) vs candidate (INT8).
129    #[inline]
130    pub fn asymmetric_cosine(&self, query: &[f32], candidate: &[u8]) -> f32 {
131        debug_assert_eq!(query.len(), self.dim);
132        debug_assert_eq!(candidate.len(), self.dim);
133        let mut dot = 0.0f32;
134        let mut norm_q = 0.0f32;
135        let mut norm_c = 0.0f32;
136        for d in 0..self.dim {
137            let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
138            dot += query[d] * dequant;
139            norm_q += query[d] * query[d];
140            norm_c += dequant * dequant;
141        }
142        let denom = (norm_q * norm_c).sqrt();
143        if denom < f32::EPSILON {
144            return 1.0;
145        }
146        (1.0 - dot / denom).max(0.0)
147    }
148
149    /// Asymmetric negative inner product: query (FP32) vs candidate (INT8).
150    #[inline]
151    pub fn asymmetric_ip(&self, query: &[f32], candidate: &[u8]) -> f32 {
152        debug_assert_eq!(query.len(), self.dim);
153        debug_assert_eq!(candidate.len(), self.dim);
154        let mut dot = 0.0f32;
155        for d in 0..self.dim {
156            let dequant = self.mins[d] + candidate[d] as f32 * self.scales[d];
157            dot += query[d] * dequant;
158        }
159        -dot
160    }
161
162    /// Vector dimension count.
163    pub fn dimensions(&self) -> usize {
164        self.dim
165    }
166}
167
168#[cfg(test)]
169mod tests {
170    use super::*;
171
172    fn make_vectors() -> Vec<Vec<f32>> {
173        (0..100)
174            .map(|i| vec![i as f32 * 0.1, (i as f32).sin(), (i as f32).cos()])
175            .collect()
176    }
177
178    #[test]
179    fn quantize_dequantize_roundtrip() {
180        let vecs = make_vectors();
181        let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
182        let codec = Sq8Codec::calibrate(&refs, 3);
183
184        for v in &vecs {
185            let q = codec.quantize(v);
186            let dq = codec.dequantize(&q);
187            for d in 0..3 {
188                let error = (v[d] - dq[d]).abs();
189                let range = codec.maxs[d] - codec.mins[d];
190                assert!(
191                    error <= range / 255.0 + 1e-6,
192                    "d={d}: error={error}, max_step={}",
193                    range / 255.0
194                );
195            }
196        }
197    }
198
199    #[test]
200    fn asymmetric_l2_close_to_exact() {
201        let vecs = make_vectors();
202        let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
203        let codec = Sq8Codec::calibrate(&refs, 3);
204
205        let query = &[5.0, 0.5, -0.5];
206        for v in &vecs {
207            let q = codec.quantize(v);
208            let exact = crate::distance::l2_squared(query, v);
209            let approx = codec.asymmetric_l2(query, &q);
210            let rel_error = if exact > 0.01 {
211                (exact - approx).abs() / exact
212            } else {
213                (exact - approx).abs()
214            };
215            assert!(
216                rel_error < 0.05 || (exact - approx).abs() < 0.1,
217                "exact={exact}, approx={approx}, rel_error={rel_error}"
218            );
219        }
220    }
221
222    #[test]
223    fn batch_quantize() {
224        let vecs = make_vectors();
225        let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
226        let codec = Sq8Codec::calibrate(&refs, 3);
227
228        let batch = codec.quantize_batch(&refs);
229        assert_eq!(batch.len(), 3 * 100);
230
231        let single = codec.quantize(&vecs[0]);
232        assert_eq!(&batch[0..3], &single[..]);
233    }
234
235    #[test]
236    fn constant_dimension_handled() {
237        let vecs: Vec<Vec<f32>> = (0..10).map(|i| vec![5.0, i as f32]).collect();
238        let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
239        let codec = Sq8Codec::calibrate(&refs, 2);
240
241        let q = codec.quantize(&[5.0, 3.0]);
242        assert_eq!(q[0], 0); // constant dim
243    }
244}