Skip to main content

edgehdf5_memory/
pq.rs

1//! Product Quantization (PQ) for approximate nearest neighbor search.
2//!
3//! Compresses high-dimensional vectors into compact codes by splitting each
4//! vector into subvectors and quantizing each subvector to its nearest
5//! centroid from a learned codebook.
6//!
7//! Default: 384-dim → 48 subvectors × 256 centroids = 48 bytes per vector (8x compression).
8
9use crate::cosine_similarity_prenorm;
10
11/// Product Quantizer with learned codebooks.
12pub struct ProductQuantizer {
13    /// Number of sub-vector segments.
14    pub num_subvectors: usize,
15    /// Number of centroids per sub-vector (max 256 for u8 codes).
16    pub num_centroids: usize,
17    /// Original vector dimension.
18    pub dim: usize,
19    /// Dimension of each sub-vector.
20    pub sub_dim: usize,
21    /// Codebook: `[num_subvectors][num_centroids][sub_dim]` stored flat.
22    /// Layout: codebook[sv * num_centroids * sub_dim + c * sub_dim + d]
23    pub codebook: Vec<f32>,
24}
25
26impl ProductQuantizer {
27    /// Train a product quantizer from a set of vectors using k-means.
28    ///
29    /// `num_subvectors` must evenly divide the vector dimension.
30    /// `num_centroids` must be <= 256 (for u8 encoding).
31    pub fn train(
32        vectors: &[Vec<f32>],
33        dim: usize,
34        num_subvectors: usize,
35        num_centroids: usize,
36    ) -> Self {
37        assert!(num_centroids <= 256, "num_centroids must be <= 256");
38        assert!(dim.is_multiple_of(num_subvectors), "dim must be divisible by num_subvectors");
39        assert!(!vectors.is_empty(), "need at least one vector to train");
40
41        let sub_dim = dim / num_subvectors;
42        let mut codebook = vec![0.0f32; num_subvectors * num_centroids * sub_dim];
43
44        let actual_centroids = num_centroids.min(vectors.len());
45
46        for sv in 0..num_subvectors {
47            let offset = sv * sub_dim;
48            // Extract sub-vectors for this segment
49            let sub_vecs: Vec<&[f32]> = vectors
50                .iter()
51                .map(|v| &v[offset..offset + sub_dim])
52                .collect();
53
54            // Initialize centroids from first `actual_centroids` vectors
55            let cb_offset = sv * num_centroids * sub_dim;
56            for c in 0..actual_centroids {
57                let src = sub_vecs[c % sub_vecs.len()];
58                let dst = &mut codebook[cb_offset + c * sub_dim..cb_offset + (c + 1) * sub_dim];
59                dst.copy_from_slice(src);
60            }
61            // Duplicate if we have fewer vectors than centroids
62            for c in actual_centroids..num_centroids {
63                let src_c = c % actual_centroids;
64                let (src_start, dst_start) = (cb_offset + src_c * sub_dim, cb_offset + c * sub_dim);
65                for d in 0..sub_dim {
66                    codebook[dst_start + d] = codebook[src_start + d];
67                }
68            }
69
70            // K-means iterations
71            let max_iters = 10;
72            let mut assignments = vec![0u8; sub_vecs.len()];
73
74            for _ in 0..max_iters {
75                // Assignment step
76                let mut changed = false;
77                for (vi, sv_data) in sub_vecs.iter().enumerate() {
78                    let mut best_c = 0u8;
79                    let mut best_dist = f32::MAX;
80                    for c in 0..actual_centroids {
81                        let cb_start = cb_offset + c * sub_dim;
82                        let centroid = &codebook[cb_start..cb_start + sub_dim];
83                        let dist = l2_sq(sv_data, centroid);
84                        if dist < best_dist {
85                            best_dist = dist;
86                            best_c = c as u8;
87                        }
88                    }
89                    if assignments[vi] != best_c {
90                        assignments[vi] = best_c;
91                        changed = true;
92                    }
93                }
94                if !changed {
95                    break;
96                }
97
98                // Update step: recompute centroids as mean of assigned vectors
99                let mut counts = vec![0u32; actual_centroids];
100                // Zero out centroids
101                for c in 0..actual_centroids {
102                    let start = cb_offset + c * sub_dim;
103                    for d in 0..sub_dim {
104                        codebook[start + d] = 0.0;
105                    }
106                }
107                for (vi, sv_data) in sub_vecs.iter().enumerate() {
108                    let c = assignments[vi] as usize;
109                    counts[c] += 1;
110                    let start = cb_offset + c * sub_dim;
111                    for d in 0..sub_dim {
112                        codebook[start + d] += sv_data[d];
113                    }
114                }
115                for (c, &count) in counts.iter().enumerate().take(actual_centroids) {
116                    if count > 0 {
117                        let start = cb_offset + c * sub_dim;
118                        let cnt = count as f32;
119                        for d in 0..sub_dim {
120                            codebook[start + d] /= cnt;
121                        }
122                    }
123                }
124            }
125        }
126
127        Self {
128            num_subvectors,
129            num_centroids,
130            dim,
131            sub_dim,
132            codebook,
133        }
134    }
135
136    /// Encode a vector into PQ codes (one u8 per subvector).
137    pub fn encode(&self, vector: &[f32]) -> Vec<u8> {
138        assert_eq!(vector.len(), self.dim);
139        let mut codes = Vec::with_capacity(self.num_subvectors);
140
141        for sv in 0..self.num_subvectors {
142            let v_offset = sv * self.sub_dim;
143            let sub = &vector[v_offset..v_offset + self.sub_dim];
144            let cb_offset = sv * self.num_centroids * self.sub_dim;
145
146            let mut best_c = 0u8;
147            let mut best_dist = f32::MAX;
148            for c in 0..self.num_centroids {
149                let c_start = cb_offset + c * self.sub_dim;
150                let centroid = &self.codebook[c_start..c_start + self.sub_dim];
151                let dist = l2_sq(sub, centroid);
152                if dist < best_dist {
153                    best_dist = dist;
154                    best_c = c as u8;
155                }
156            }
157            codes.push(best_c);
158        }
159        codes
160    }
161
162    /// Decode PQ codes back to an approximate vector.
163    pub fn decode(&self, codes: &[u8]) -> Vec<f32> {
164        assert_eq!(codes.len(), self.num_subvectors);
165        let mut result = Vec::with_capacity(self.dim);
166
167        for (sv, &code) in codes.iter().enumerate() {
168            let cb_offset = sv * self.num_centroids * self.sub_dim;
169            let c_start = cb_offset + code as usize * self.sub_dim;
170            result.extend_from_slice(&self.codebook[c_start..c_start + self.sub_dim]);
171        }
172        result
173    }
174
175    /// Precompute distance table for asymmetric distance computation.
176    ///
177    /// Returns a table of shape `[num_subvectors][num_centroids]` (stored flat)
178    /// containing the squared L2 distance from each query sub-vector to each
179    /// centroid.
180    pub fn precompute_distance_table(&self, query: &[f32]) -> Vec<f32> {
181        assert_eq!(query.len(), self.dim);
182        let mut table = Vec::with_capacity(self.num_subvectors * self.num_centroids);
183
184        for sv in 0..self.num_subvectors {
185            let q_offset = sv * self.sub_dim;
186            let q_sub = &query[q_offset..q_offset + self.sub_dim];
187            let cb_offset = sv * self.num_centroids * self.sub_dim;
188
189            for c in 0..self.num_centroids {
190                let c_start = cb_offset + c * self.sub_dim;
191                let centroid = &self.codebook[c_start..c_start + self.sub_dim];
192                table.push(l2_sq(q_sub, centroid));
193            }
194        }
195        table
196    }
197
198    /// Compute asymmetric distance between query and encoded vector.
199    ///
200    /// Uses a precomputed distance table for speed — this is just
201    /// `num_subvectors` table lookups + additions.
202    pub fn asymmetric_distance_with_table(&self, table: &[f32], codes: &[u8]) -> f32 {
203        let mut dist = 0.0f32;
204        for (sv, &code) in codes.iter().enumerate() {
205            dist += table[sv * self.num_centroids + code as usize];
206        }
207        dist
208    }
209
210    /// Compute asymmetric distance between a query and an encoded vector.
211    pub fn asymmetric_distance(&self, query: &[f32], codes: &[u8]) -> f32 {
212        let table = self.precompute_distance_table(query);
213        self.asymmetric_distance_with_table(&table, codes)
214    }
215
216    /// Search a collection of PQ-encoded vectors and return the top-k nearest
217    /// by asymmetric distance (smallest distance = most similar).
218    ///
219    /// `all_codes` is a flat buffer: `[n_vectors * num_subvectors]`.
220    /// `tombstones` marks deleted vectors.
221    pub fn search(
222        &self,
223        query: &[f32],
224        all_codes: &[u8],
225        tombstones: &[u8],
226        k: usize,
227    ) -> Vec<(usize, f32)> {
228        let table = self.precompute_distance_table(query);
229        let n = all_codes.len() / self.num_subvectors;
230        let mut results: Vec<(usize, f32)> = Vec::with_capacity(n);
231
232        for i in 0..n {
233            if i < tombstones.len() && tombstones[i] != 0 {
234                continue;
235            }
236            let codes = &all_codes[i * self.num_subvectors..(i + 1) * self.num_subvectors];
237            let dist = self.asymmetric_distance_with_table(&table, codes);
238            results.push((i, dist));
239        }
240
241        // Sort by distance ascending (smaller = closer)
242        results.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
243        results.truncate(k);
244        results
245    }
246
247    /// Search with PQ then re-rank top candidates with exact cosine similarity.
248    ///
249    /// Returns `(index, cosine_score)` pairs sorted by score descending.
250    pub fn search_rerank(
251        &self,
252        query: &[f32],
253        all_codes: &[u8],
254        vectors: &[Vec<f32>],
255        tombstones: &[u8],
256        candidates: usize,
257        k: usize,
258    ) -> Vec<(usize, f32)> {
259        let pq_results = self.search(query, all_codes, tombstones, candidates);
260        let query_norm = rustyhdf5_accel::vector_norm(query);
261
262        let mut reranked: Vec<(usize, f32)> = pq_results
263            .iter()
264            .map(|&(idx, _)| {
265                let vec_norm = rustyhdf5_accel::vector_norm(&vectors[idx]);
266                let score = cosine_similarity_prenorm(query, query_norm, &vectors[idx], vec_norm);
267                (idx, score)
268            })
269            .collect();
270
271        reranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
272        reranked.truncate(k);
273        reranked
274    }
275
276    /// Encode all vectors and return flat code buffer.
277    pub fn encode_all(&self, vectors: &[Vec<f32>]) -> Vec<u8> {
278        let mut all_codes = Vec::with_capacity(vectors.len() * self.num_subvectors);
279        for v in vectors {
280            all_codes.extend(self.encode(v));
281        }
282        all_codes
283    }
284
285    /// Serialize the quantizer state to flat data for HDF5 storage.
286    /// Returns (codebook_flat, metadata: [num_subvectors, num_centroids, dim]).
287    pub fn to_hdf5_data(&self) -> (&[f32], [i64; 3]) {
288        (
289            &self.codebook,
290            [
291                self.num_subvectors as i64,
292                self.num_centroids as i64,
293                self.dim as i64,
294            ],
295        )
296    }
297
298    /// Reconstruct from HDF5 data.
299    pub fn from_hdf5_data(codebook: Vec<f32>, metadata: [i64; 3]) -> Self {
300        let num_subvectors = metadata[0] as usize;
301        let num_centroids = metadata[1] as usize;
302        let dim = metadata[2] as usize;
303        let sub_dim = dim / num_subvectors;
304        Self {
305            num_subvectors,
306            num_centroids,
307            dim,
308            sub_dim,
309            codebook,
310        }
311    }
312}
313
314/// Squared L2 distance between two slices.
315#[inline]
316fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
317    let mut sum = 0.0f32;
318    for i in 0..a.len() {
319        let d = a[i] - b[i];
320        sum += d * d;
321    }
322    sum
323}
324
325// ---------------------------------------------------------------------------
326// Tests
327// ---------------------------------------------------------------------------
328
329#[cfg(test)]
330mod tests {
331    use super::*;
332
333    fn make_vectors(n: usize, dim: usize, seed: u32) -> Vec<Vec<f32>> {
334        let mut s = seed;
335        let mut next = || -> f32 {
336            s = s.wrapping_mul(1103515245).wrapping_add(12345);
337            ((s >> 16) as f32) / 65536.0 - 0.5
338        };
339        (0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
340    }
341
342    #[test]
343    fn encode_decode_roundtrip() {
344        let dim = 384;
345        let vectors = make_vectors(200, dim, 42);
346        let pq = ProductQuantizer::train(&vectors, dim, 48, 256);
347
348        // Check reconstruction error
349        let mut total_error = 0.0f32;
350        for v in &vectors {
351            let codes = pq.encode(v);
352            let decoded = pq.decode(&codes);
353            assert_eq!(decoded.len(), dim);
354            let error: f32 = v
355                .iter()
356                .zip(&decoded)
357                .map(|(a, b)| (a - b) * (a - b))
358                .sum();
359            total_error += error;
360        }
361        let avg_error = total_error / vectors.len() as f32 / dim as f32;
362        // Reconstruction error should be reasonable
363        assert!(
364            avg_error < 0.1,
365            "avg per-dim reconstruction error too high: {avg_error}"
366        );
367    }
368
369    #[test]
370    fn pq_code_size() {
371        let dim = 384;
372        let num_sub = 48;
373        let vectors = make_vectors(100, dim, 42);
374        let pq = ProductQuantizer::train(&vectors, dim, num_sub, 256);
375        let codes = pq.encode(&vectors[0]);
376        assert_eq!(codes.len(), num_sub); // 48 bytes per vector
377    }
378
379    #[test]
380    fn asymmetric_distance_basic() {
381        let dim = 16;
382        let vectors = make_vectors(50, dim, 42);
383        let pq = ProductQuantizer::train(&vectors, dim, 4, 16);
384
385        let query = &vectors[0];
386        let codes = pq.encode(&vectors[1]);
387
388        let dist = pq.asymmetric_distance(query, &codes);
389        assert!(dist >= 0.0, "distance should be non-negative");
390    }
391
392    #[test]
393    fn pq_search_returns_closest() {
394        let dim = 32;
395        let mut vectors = make_vectors(100, dim, 42);
396        // Make vectors[0] identical to query
397        let query = vectors[0].clone();
398        vectors[0] = query.clone();
399
400        let pq = ProductQuantizer::train(&vectors, dim, 8, 32);
401        let all_codes = pq.encode_all(&vectors);
402        let tombstones = vec![0u8; 100];
403
404        let results = pq.search(&query, &all_codes, &tombstones, 10);
405        assert!(!results.is_empty());
406        // The query itself (index 0) should be in top results
407        let top_indices: Vec<usize> = results.iter().map(|r| r.0).collect();
408        assert!(top_indices.contains(&0), "query vector should be in top results");
409    }
410
411    #[test]
412    fn pq_search_respects_tombstones() {
413        let dim = 16;
414        let vectors = make_vectors(20, dim, 42);
415        let pq = ProductQuantizer::train(&vectors, dim, 4, 16);
416        let all_codes = pq.encode_all(&vectors);
417        let mut tombstones = vec![0u8; 20];
418        tombstones[0] = 1;
419
420        let results = pq.search(&vectors[0], &all_codes, &tombstones, 20);
421        assert!(results.iter().all(|r| r.0 != 0));
422    }
423
424    #[test]
425    fn pq_search_rerank_improves_quality() {
426        let dim = 64;
427        let vectors = make_vectors(500, dim, 42);
428        let query = vectors[0].clone();
429
430        let pq = ProductQuantizer::train(&vectors, dim, 8, 64);
431        let all_codes = pq.encode_all(&vectors);
432        let tombstones = vec![0u8; 500];
433
434        let reranked = pq.search_rerank(&query, &all_codes, &vectors, &tombstones, 100, 10);
435        assert!(reranked.len() <= 10);
436        // First result should have high cosine similarity (it's the query itself)
437        assert!(reranked[0].1 > 0.9, "top reranked score: {}", reranked[0].1);
438    }
439
440    #[test]
441    fn distance_table_precomputation() {
442        let dim = 16;
443        let vectors = make_vectors(50, dim, 42);
444        let pq = ProductQuantizer::train(&vectors, dim, 4, 16);
445
446        let query = &vectors[0];
447        let codes = pq.encode(&vectors[1]);
448
449        // Distance with table should equal without table
450        let table = pq.precompute_distance_table(query);
451        let dist_table = pq.asymmetric_distance_with_table(&table, &codes);
452        let dist_direct = pq.asymmetric_distance(query, &codes);
453        assert!((dist_table - dist_direct).abs() < 1e-6);
454    }
455
456    #[test]
457    fn pq_hdf5_roundtrip() {
458        let dim = 32;
459        let vectors = make_vectors(50, dim, 42);
460        let pq = ProductQuantizer::train(&vectors, dim, 8, 32);
461
462        let (cb, meta) = pq.to_hdf5_data();
463        let pq2 = ProductQuantizer::from_hdf5_data(cb.to_vec(), meta);
464
465        assert_eq!(pq.num_subvectors, pq2.num_subvectors);
466        assert_eq!(pq.num_centroids, pq2.num_centroids);
467        assert_eq!(pq.dim, pq2.dim);
468        assert_eq!(pq.codebook, pq2.codebook);
469    }
470
471    #[test]
472    fn pq_asymmetric_ranking_reasonable_recall() {
473        // Check that PQ ranking has reasonable overlap with exact ranking
474        let dim = 64;
475        let n = 500;
476        let vectors = make_vectors(n, dim, 42);
477        let query = vectors[0].clone();
478
479        // Exact top-10 by cosine similarity
480        let mut exact: Vec<(usize, f32)> = vectors
481            .iter()
482            .enumerate()
483            .map(|(i, v)| (i, rustyhdf5_accel::cosine_similarity(&query, v)))
484            .collect();
485        exact.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
486        let exact_top10: Vec<usize> = exact.iter().take(10).map(|r| r.0).collect();
487
488        // PQ approximate top-20 then check overlap with exact top-10
489        let pq = ProductQuantizer::train(&vectors, dim, 8, 64);
490        let all_codes = pq.encode_all(&vectors);
491        let tombstones = vec![0u8; n];
492        let pq_top20 = pq.search(&query, &all_codes, &tombstones, 20);
493        let pq_indices: Vec<usize> = pq_top20.iter().map(|r| r.0).collect();
494
495        let overlap = exact_top10
496            .iter()
497            .filter(|i| pq_indices.contains(i))
498            .count();
499        // Recall should be at least 50% (5 out of 10)
500        assert!(
501            overlap >= 5,
502            "PQ recall too low: {overlap}/10 overlap with exact top-10 in PQ top-20"
503        );
504    }
505}