velesdb_core/index/hnsw/native/
quantization.rs1use std::sync::Arc;
27
28#[inline]
42fn distance_l2_quantized_simd(a: &[u8], b: &[u8]) -> u32 {
43 debug_assert_eq!(a.len(), b.len());
44
45 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 for i in 0..chunks {
56 let base = i * 8;
57
58 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 #[allow(clippy::cast_sign_loss)] {
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 let base = chunks * 8;
81 for i in 0..remainder {
82 let diff = i32::from(a[base + i]) - i32::from(b[base + i]);
83 #[allow(clippy::cast_sign_loss)]
85 {
86 sum0 += (diff * diff) as u32;
87 }
88 }
89
90 sum0 + sum1 + sum2 + sum3
91}
92
93#[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#[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#[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#[derive(Debug, Clone)]
184pub struct ScalarQuantizer {
185 pub min_vals: Vec<f32>,
187 pub scales: Vec<f32>,
189 pub inv_scales: Vec<f32>,
191 pub dimension: usize,
193}
194
195#[derive(Debug, Clone)]
197pub struct QuantizedVector {
198 pub data: Vec<u8>,
200}
201
202#[derive(Debug, Clone)]
204pub struct QuantizedVectorStore {
205 quantizer: Arc<ScalarQuantizer>,
207 data: Vec<u8>,
209 count: usize,
211}
212
213impl ScalarQuantizer {
214 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 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 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 } else {
257 255.0 / range
258 }
259 })
260 .collect();
261
262 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 #[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 #[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 f32::from(q) * inv_scale + min
304 })
305 .collect()
306 }
307
308 #[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 #[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 #[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 #[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 #[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 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 #[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 #[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 #[must_use]
397 pub fn len(&self) -> usize {
398 self.count
399 }
400
401 #[must_use]
403 pub fn is_empty(&self) -> bool {
404 self.count == 0
405 }
406
407 #[must_use]
409 pub fn quantizer(&self) -> &ScalarQuantizer {
410 &self.quantizer
411 }
412}