Skip to main content

velesdb_core/index/hnsw/native/
quantization.rs

1//! Scalar Quantization (SQ8) for fast HNSW traversal.
2//!
3//! Based on VSAG paper (arXiv:2503.17911): dual-precision architecture
4//! using int8 for graph traversal and float32 for final re-ranking.
5//!
6//! # Performance Benefits
7//!
8//! - **4x memory bandwidth reduction** during traversal
9//! - **SIMD-friendly**: 32 int8 values fit in 256-bit register (vs 8 float32)
10//! - **Cache efficiency**: More vectors fit in L1/L2 cache
11//!
12//! # Algorithm
13//!
14//! For each dimension:
15//! - Compute min/max from training data
16//! - Scale to [0, 255] range: `q = round((x - min) / (max - min) * 255)`
17//! - Store scale and offset for reconstruction
18//!
19//! # Safety (EPIC-032/US-007)
20//!
21//! All `as u32` casts in distance computation are proven safe:
22//! - Input: u8 values in [0, 255]
23//! - Difference: i32 in [-255, 255]
24//! - Squared: i32 in [0, 65025] (always non-negative, fits in u32)
25
26use std::sync::Arc;
27
28// =============================================================================
29// SIMD-optimized distance computation for int8 quantized vectors
30// =============================================================================
31
32/// Computes L2 squared distance between two quantized vectors using SIMD.
33///
34/// Uses 8-wide unrolling for better instruction-level parallelism.
35/// On x86_64 with AVX2, processes 32 bytes per iteration.
36///
37/// # Performance
38///
39/// - **4x memory bandwidth reduction** vs float32
40/// - **Better SIMD utilization**: 32 int8 fit in 256-bit register vs 8 float32
41#[inline]
42fn distance_l2_quantized_simd(a: &[u8], b: &[u8]) -> u32 {
43    debug_assert_eq!(a.len(), b.len());
44
45    // Process in chunks of 8 for better ILP (Instruction Level Parallelism)
46    let chunks = a.len() / 8;
47    let remainder = a.len() % 8;
48
49    let mut sum0: u32 = 0;
50    let mut sum1: u32 = 0;
51    let mut sum2: u32 = 0;
52    let mut sum3: u32 = 0;
53
54    // Main loop: 8-wide unrolling
55    for i in 0..chunks {
56        let base = i * 8;
57
58        // Unroll 8 iterations with 4 accumulators
59        let d0 = i32::from(a[base]) - i32::from(b[base]);
60        let d1 = i32::from(a[base + 1]) - i32::from(b[base + 1]);
61        let d2 = i32::from(a[base + 2]) - i32::from(b[base + 2]);
62        let d3 = i32::from(a[base + 3]) - i32::from(b[base + 3]);
63        let d4 = i32::from(a[base + 4]) - i32::from(b[base + 4]);
64        let d5 = i32::from(a[base + 5]) - i32::from(b[base + 5]);
65        let d6 = i32::from(a[base + 6]) - i32::from(b[base + 6]);
66        let d7 = i32::from(a[base + 7]) - i32::from(b[base + 7]);
67
68        // SAFETY (EPIC-032/US-007): d_i in [-255, 255], so d_i*d_i in [0, 65025]
69        // This is always non-negative and fits in u32 (max 4,294,967,295)
70        #[allow(clippy::cast_sign_loss)] // Proven non-negative: square of integer
71        {
72            sum0 += (d0 * d0) as u32 + (d4 * d4) as u32;
73            sum1 += (d1 * d1) as u32 + (d5 * d5) as u32;
74            sum2 += (d2 * d2) as u32 + (d6 * d6) as u32;
75            sum3 += (d3 * d3) as u32 + (d7 * d7) as u32;
76        }
77    }
78
79    // Handle remainder
80    let base = chunks * 8;
81    for i in 0..remainder {
82        let diff = i32::from(a[base + i]) - i32::from(b[base + i]);
83        // SAFETY (EPIC-032/US-007): diff in [-255, 255], diff*diff in [0, 65025]
84        #[allow(clippy::cast_sign_loss)]
85        {
86            sum0 += (diff * diff) as u32;
87        }
88    }
89
90    sum0 + sum1 + sum2 + sum3
91}
92
93/// Computes asymmetric L2 distance: float32 query vs quantized candidate.
94///
95/// Uses precomputed lookup tables for efficient SIMD execution.
96/// Based on VSAG paper's ADT (Asymmetric Distance Table) approach.
97#[inline]
98fn distance_l2_asymmetric_simd(
99    query: &[f32],
100    quantized: &[u8],
101    min_vals: &[f32],
102    inv_scales: &[f32],
103) -> f32 {
104    debug_assert_eq!(query.len(), quantized.len());
105    debug_assert_eq!(query.len(), min_vals.len());
106    debug_assert_eq!(query.len(), inv_scales.len());
107
108    let chunks = query.len() / 4;
109    let remainder = query.len() % 4;
110
111    let (sum0, sum1, sum2, sum3) =
112        asymmetric_chunked_sum(query, quantized, min_vals, inv_scales, chunks);
113
114    let remainder_sum = asymmetric_remainder_sum(
115        query,
116        quantized,
117        min_vals,
118        inv_scales,
119        chunks * 4,
120        remainder,
121    );
122
123    (sum0 + sum1 + sum2 + sum3 + remainder_sum).sqrt()
124}
125
126/// Computes the main chunked (4-wide) sum for asymmetric L2 distance.
127#[inline]
128fn asymmetric_chunked_sum(
129    query: &[f32],
130    quantized: &[u8],
131    min_vals: &[f32],
132    inv_scales: &[f32],
133    chunks: usize,
134) -> (f32, f32, f32, f32) {
135    let mut sum0: f32 = 0.0;
136    let mut sum1: f32 = 0.0;
137    let mut sum2: f32 = 0.0;
138    let mut sum3: f32 = 0.0;
139
140    for i in 0..chunks {
141        let base = i * 4;
142
143        let dq0 = f32::from(quantized[base]) * inv_scales[base] + min_vals[base];
144        let dq1 = f32::from(quantized[base + 1]) * inv_scales[base + 1] + min_vals[base + 1];
145        let dq2 = f32::from(quantized[base + 2]) * inv_scales[base + 2] + min_vals[base + 2];
146        let dq3 = f32::from(quantized[base + 3]) * inv_scales[base + 3] + min_vals[base + 3];
147
148        let d0 = query[base] - dq0;
149        let d1 = query[base + 1] - dq1;
150        let d2 = query[base + 2] - dq2;
151        let d3 = query[base + 3] - dq3;
152
153        sum0 += d0 * d0;
154        sum1 += d1 * d1;
155        sum2 += d2 * d2;
156        sum3 += d3 * d3;
157    }
158
159    (sum0, sum1, sum2, sum3)
160}
161
162/// Computes the remainder sum for asymmetric L2 distance (elements not covered by 4-wide chunks).
163#[inline]
164fn asymmetric_remainder_sum(
165    query: &[f32],
166    quantized: &[u8],
167    min_vals: &[f32],
168    inv_scales: &[f32],
169    base: usize,
170    remainder: usize,
171) -> f32 {
172    let mut sum = 0.0_f32;
173    for i in 0..remainder {
174        let idx = base + i;
175        let dq = f32::from(quantized[idx]) * inv_scales[idx] + min_vals[idx];
176        let diff = query[idx] - dq;
177        sum += diff * diff;
178    }
179    sum
180}
181
182/// Quantization parameters learned from training data.
183#[derive(Debug, Clone)]
184pub struct ScalarQuantizer {
185    /// Minimum value per dimension
186    pub min_vals: Vec<f32>,
187    /// Scale factor per dimension: 255 / (max - min)
188    pub scales: Vec<f32>,
189    /// Inverse scale factor: 1 / scale (precomputed for fast dequantization)
190    pub inv_scales: Vec<f32>,
191    /// Vector dimension
192    pub dimension: usize,
193}
194
195/// Quantized vector storage (int8 per dimension).
196#[derive(Debug, Clone)]
197pub struct QuantizedVector {
198    /// Quantized values [0, 255]
199    pub data: Vec<u8>,
200}
201
202/// Quantized vector storage with shared quantizer reference.
203#[derive(Debug, Clone)]
204pub struct QuantizedVectorStore {
205    /// Shared quantizer parameters
206    quantizer: Arc<ScalarQuantizer>,
207    /// Quantized vectors (flattened: node_id * dimension + dim_idx)
208    data: Vec<u8>,
209    /// Number of vectors stored
210    count: usize,
211}
212
213impl ScalarQuantizer {
214    /// Creates a new quantizer from training vectors.
215    ///
216    /// # Arguments
217    ///
218    /// * `vectors` - Training vectors to compute min/max per dimension
219    ///
220    /// # Errors
221    ///
222    /// Returns `Error::InvalidQuantizerConfig` if `vectors` is empty or
223    /// vectors have inconsistent dimensions.
224    pub fn train(vectors: &[&[f32]]) -> crate::error::Result<Self> {
225        if vectors.is_empty() {
226            return Err(crate::error::Error::InvalidQuantizerConfig(
227                "cannot train on empty vectors".to_string(),
228            ));
229        }
230        let dimension = vectors[0].len();
231        if !vectors.iter().all(|v| v.len() == dimension) {
232            return Err(crate::error::Error::InvalidQuantizerConfig(
233                "all vectors must have same dimension".to_string(),
234            ));
235        }
236
237        let mut min_vals = vec![f32::MAX; dimension];
238        let mut max_vals = vec![f32::MIN; dimension];
239
240        // Find min/max per dimension
241        for vec in vectors {
242            for (i, &val) in vec.iter().enumerate() {
243                min_vals[i] = min_vals[i].min(val);
244                max_vals[i] = max_vals[i].max(val);
245            }
246        }
247
248        // Compute scales (avoid division by zero)
249        let scales: Vec<f32> = min_vals
250            .iter()
251            .zip(max_vals.iter())
252            .map(|(&min, &max)| {
253                let range = max - min;
254                if range.abs() < 1e-10 {
255                    1.0 // Constant dimension, scale doesn't matter
256                } else {
257                    255.0 / range
258                }
259            })
260            .collect();
261
262        // Precompute inverse scales for fast dequantization
263        let inv_scales: Vec<f32> = scales.iter().map(|&s| 1.0 / s).collect();
264
265        Ok(Self {
266            min_vals,
267            scales,
268            inv_scales,
269            dimension,
270        })
271    }
272
273    /// Quantizes a float32 vector to int8.
274    #[must_use]
275    pub fn quantize(&self, vector: &[f32]) -> QuantizedVector {
276        debug_assert_eq!(vector.len(), self.dimension);
277
278        let data: Vec<u8> = vector
279            .iter()
280            .zip(self.min_vals.iter())
281            .zip(self.scales.iter())
282            .map(|((&val, &min), &scale)| {
283                let q = ((val - min) * scale).round();
284                q.clamp(0.0, 255.0) as u8
285            })
286            .collect();
287
288        QuantizedVector { data }
289    }
290
291    /// Dequantizes an int8 vector back to float32.
292    #[must_use]
293    pub fn dequantize(&self, quantized: &QuantizedVector) -> Vec<f32> {
294        debug_assert_eq!(quantized.data.len(), self.dimension);
295
296        quantized
297            .data
298            .iter()
299            .zip(self.min_vals.iter())
300            .zip(self.inv_scales.iter())
301            .map(|((&q, &min), &inv_scale)| {
302                // x = q * inv_scale + min (multiplication is faster than division)
303                f32::from(q) * inv_scale + min
304            })
305            .collect()
306    }
307
308    /// Computes approximate L2 distance between quantized vectors.
309    ///
310    /// This is ~4x faster than float32 due to SIMD efficiency.
311    #[inline]
312    #[must_use]
313    pub fn distance_l2_quantized(&self, a: &QuantizedVector, b: &QuantizedVector) -> u32 {
314        debug_assert_eq!(a.data.len(), b.data.len());
315        distance_l2_quantized_simd(&a.data, &b.data)
316    }
317
318    /// Computes approximate L2 distance using raw slices (zero-copy).
319    ///
320    /// Useful for QuantizedVectorStore.get_slice() access pattern.
321    #[inline]
322    #[must_use]
323    pub fn distance_l2_quantized_slice(&self, a: &[u8], b: &[u8]) -> u32 {
324        debug_assert_eq!(a.len(), b.len());
325        distance_l2_quantized_simd(a, b)
326    }
327
328    /// Computes approximate L2 distance: quantized vs float32 query.
329    ///
330    /// Asymmetric distance: query stays in float32, candidates in int8.
331    /// This is the VSAG "ADT" (Asymmetric Distance Table) approach.
332    #[inline]
333    #[must_use]
334    pub fn distance_l2_asymmetric(&self, query: &[f32], quantized: &QuantizedVector) -> f32 {
335        debug_assert_eq!(query.len(), self.dimension);
336        debug_assert_eq!(quantized.data.len(), self.dimension);
337
338        distance_l2_asymmetric_simd(query, &quantized.data, &self.min_vals, &self.inv_scales)
339    }
340
341    /// Computes asymmetric L2 distance using raw slice (zero-copy).
342    #[inline]
343    #[must_use]
344    pub fn distance_l2_asymmetric_slice(&self, query: &[f32], quantized: &[u8]) -> f32 {
345        debug_assert_eq!(query.len(), self.dimension);
346        debug_assert_eq!(quantized.len(), self.dimension);
347
348        distance_l2_asymmetric_simd(query, quantized, &self.min_vals, &self.inv_scales)
349    }
350}
351
352impl QuantizedVectorStore {
353    /// Creates a new quantized vector store.
354    #[must_use]
355    pub fn new(quantizer: Arc<ScalarQuantizer>, capacity: usize) -> Self {
356        let dimension = quantizer.dimension;
357        Self {
358            quantizer,
359            data: Vec::with_capacity(capacity * dimension),
360            count: 0,
361        }
362    }
363
364    /// Adds a vector to the store (quantizes it first).
365    pub fn push(&mut self, vector: &[f32]) {
366        let quantized = self.quantizer.quantize(vector);
367        self.data.extend(quantized.data);
368        self.count += 1;
369    }
370
371    /// Gets a quantized vector by index.
372    #[must_use]
373    pub fn get(&self, index: usize) -> Option<QuantizedVector> {
374        if index >= self.count {
375            return None;
376        }
377        let start = index * self.quantizer.dimension;
378        let end = start + self.quantizer.dimension;
379        Some(QuantizedVector {
380            data: self.data[start..end].to_vec(),
381        })
382    }
383
384    /// Gets raw slice for a quantized vector (zero-copy).
385    #[must_use]
386    pub fn get_slice(&self, index: usize) -> Option<&[u8]> {
387        if index >= self.count {
388            return None;
389        }
390        let start = index * self.quantizer.dimension;
391        let end = start + self.quantizer.dimension;
392        Some(&self.data[start..end])
393    }
394
395    /// Returns the number of vectors.
396    #[must_use]
397    pub fn len(&self) -> usize {
398        self.count
399    }
400
401    /// Returns true if empty.
402    #[must_use]
403    pub fn is_empty(&self) -> bool {
404        self.count == 0
405    }
406
407    /// Returns reference to quantizer.
408    #[must_use]
409    pub fn quantizer(&self) -> &ScalarQuantizer {
410        &self.quantizer
411    }
412}