1use crate::cosine_similarity_prenorm;
10
11pub struct ProductQuantizer {
13 pub num_subvectors: usize,
15 pub num_centroids: usize,
17 pub dim: usize,
19 pub sub_dim: usize,
21 pub codebook: Vec<f32>,
24}
25
26impl ProductQuantizer {
27 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 let sub_vecs: Vec<&[f32]> = vectors
50 .iter()
51 .map(|v| &v[offset..offset + sub_dim])
52 .collect();
53
54 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 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 let max_iters = 10;
72 let mut assignments = vec![0u8; sub_vecs.len()];
73
74 for _ in 0..max_iters {
75 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 let mut counts = vec![0u32; actual_centroids];
100 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 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 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 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 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 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 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 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 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 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 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 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#[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#[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 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 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); }
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 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 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 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 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 let dim = 64;
475 let n = 500;
476 let vectors = make_vectors(n, dim, 42);
477 let query = vectors[0].clone();
478
479 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 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 assert!(
501 overlap >= 5,
502 "PQ recall too low: {overlap}/10 overlap with exact top-10 in PQ top-20"
503 );
504 }
505}