1use crate::prolly::builder::SortedBatchBuilder;
2use crate::prolly::cid::Cid;
3use crate::prolly::config::Config;
4use crate::prolly::encoding::Encoding;
5use crate::prolly::error::Error;
6use crate::prolly::proximity::distance::{prepare_vector, query_score};
7use crate::prolly::proximity::search::{
8 retained_candidate_bytes, EligibilityCardinality, PreparedFilter, RerankCandidate,
9};
10use crate::prolly::proximity::storage::codec::{
11 put_cid, put_f32, put_f64, put_varint, Reader, MAX_OBJECT_ENTRIES,
12};
13use crate::prolly::proximity::storage::StoredRecord;
14use crate::prolly::proximity::{
15 BuildParallelism, DistanceMetric, ProximityMap, ProximitySearchStats, SearchBackend,
16 SearchCompletion, SearchPolicy, SearchRequest, SearchResult,
17};
18use crate::prolly::store::{NodePublication, PublicationOrigin, Store};
19use crate::prolly::tree::Tree;
20use crate::prolly::Prolly;
21use rayon::prelude::*;
22use std::cmp::Ordering;
23use std::collections::{BinaryHeap, HashSet};
24use xxhash_rust::xxh64::xxh64;
25
26const MAGIC: &[u8; 4] = b"PQPQ";
27const PQ_FORMAT_VERSION: u8 = 2;
28const SAMPLING_HASH_ALGORITHM_XXH64: u8 = 1;
29const SAMPLING_HASH_VERSION: u8 = 1;
30type Codebooks = Vec<Vec<Vec<f32>>>;
31type TrainingOutput = (Codebooks, usize);
32
33#[derive(Clone, Debug, PartialEq, Eq)]
35pub struct ProductQuantizationConfig {
36 pub subquantizers: u32,
37 pub centroids_per_subquantizer: u16,
38 pub training_iterations: u16,
39 pub rerank_multiplier: u32,
40 pub seed: u64,
41 pub max_training_vectors: usize,
42}
43
44impl Default for ProductQuantizationConfig {
45 fn default() -> Self {
46 Self {
47 subquantizers: 8,
48 centroids_per_subquantizer: 256,
49 training_iterations: 12,
50 rerank_multiplier: 8,
51 seed: 0,
52 max_training_vectors: 65_536,
53 }
54 }
55}
56
57impl ProductQuantizationConfig {
58 pub(crate) fn validate(&self, dimensions: u32, records: usize) -> Result<(), Error> {
59 if self.subquantizers == 0 || self.subquantizers > dimensions {
60 return Err(invalid_config("subquantizers must be in 1..=dimensions"));
61 }
62 let centroids = usize::from(self.centroids_per_subquantizer);
63 if centroids == 0 || centroids > 256 || centroids > records {
64 return Err(invalid_config(
65 "centroids_per_subquantizer must be in 1..=min(256, record_count)",
66 ));
67 }
68 if self.training_iterations == 0
69 || self.rerank_multiplier == 0
70 || self.max_training_vectors < centroids
71 {
72 return Err(invalid_config(
73 "training_iterations/rerank_multiplier must be positive and max_training_vectors must cover all centroids",
74 ));
75 }
76 Ok(())
77 }
78}
79
80#[derive(Clone, Copy, Debug, Default, PartialEq)]
82pub struct ProductQuantizationQuality {
83 pub mean_squared_error: f64,
84 pub maximum_squared_error: f64,
85}
86
87#[derive(Clone, Debug, Default, PartialEq, Eq)]
89pub struct ProductQuantizationBuildStats {
90 pub training_distance_evaluations: usize,
91 pub encoding_distance_evaluations: usize,
92 pub encoded_vectors: usize,
93 pub training_vectors: usize,
94 pub training_bytes: usize,
95 pub encoded_output_bytes: usize,
96}
97
98#[derive(Clone, Debug, Default, PartialEq, Eq)]
99pub struct ProductQuantizationBuildLimits {
100 pub max_training_vectors: Option<usize>,
101 pub max_training_bytes: Option<usize>,
102 pub max_temporary_code_bytes: Option<usize>,
103 pub max_distance_evaluations: Option<usize>,
104 pub max_encoded_output_bytes: Option<usize>,
105 pub max_worker_threads: Option<usize>,
106}
107
108impl ProductQuantizationBuildLimits {
109 fn validate(&self) -> Result<(), Error> {
110 for (name, value) in [
111 ("max_training_vectors", self.max_training_vectors),
112 ("max_training_bytes", self.max_training_bytes),
113 ("max_temporary_code_bytes", self.max_temporary_code_bytes),
114 ("max_distance_evaluations", self.max_distance_evaluations),
115 ("max_encoded_output_bytes", self.max_encoded_output_bytes),
116 ("max_worker_threads", self.max_worker_threads),
117 ] {
118 if value == Some(0) {
119 return Err(invalid_config(format!("PQ {name} must be positive")));
120 }
121 }
122 Ok(())
123 }
124}
125
126pub struct ProductQuantizer<S: Store> {
128 codes: Prolly<S>,
129 code_tree: Tree,
130 manifest: Cid,
131 pub(super) source: Cid,
132 pub(super) dimensions: u32,
133 pub(super) metric: DistanceMetric,
134 pub(super) count: u64,
135 config: ProductQuantizationConfig,
136 codebooks: Codebooks,
137 quality: ProductQuantizationQuality,
138}
139
140impl<S> ProductQuantizer<S>
141where
142 S: Store + Clone + Send + Sync,
143 S::Error: Send + Sync,
144{
145 pub fn build(
147 map: &ProximityMap<S>,
148 config: ProductQuantizationConfig,
149 parallelism: BuildParallelism,
150 ) -> Result<(Self, ProductQuantizationBuildStats), Error> {
151 Self::build_with_limits(
152 map,
153 config,
154 parallelism,
155 ProductQuantizationBuildLimits::default(),
156 )
157 }
158
159 pub fn build_with_limits(
160 map: &ProximityMap<S>,
161 config: ProductQuantizationConfig,
162 parallelism: BuildParallelism,
163 limits: ProductQuantizationBuildLimits,
164 ) -> Result<(Self, ProductQuantizationBuildStats), Error> {
165 limits.validate()?;
166 let record_count = usize::try_from(map.tree().count)
167 .map_err(|_| resource_limit("records", usize::MAX, usize::MAX))?;
168 config.validate(map.tree().config.dimensions, record_count)?;
169 if let Some(limit) = limits.max_worker_threads {
170 enforce_resource("worker_threads", Some(limit), parallelism.threads())?;
171 }
172 let sample_target = config.max_training_vectors.min(record_count);
173 enforce_resource(
174 "training_vectors",
175 limits.max_training_vectors,
176 sample_target,
177 )?;
178 enforce_resource(
179 "temporary_code_bytes",
180 limits.max_temporary_code_bytes,
181 config.subquantizers as usize,
182 )?;
183
184 let mut samples = BinaryHeap::<TrainingSample>::with_capacity(sample_target);
185 for entry in map
186 .directory_manager()
187 .range(&map.tree().directory, &[], None)?
188 {
189 let (key, bytes) = entry?;
190 let stored = StoredRecord::decode(&bytes, map.tree().config.dimensions)?;
191 let sample = TrainingSample {
192 hash: xxh64(&key, config.seed),
193 key,
194 vector: stored.vector,
195 };
196 if samples.len() < sample_target {
197 samples.push(sample);
198 } else if samples.peek().is_some_and(|worst| sample < *worst) {
199 samples.pop();
200 samples.push(sample);
201 }
202 }
203 let mut samples = samples.into_vec();
204 samples.sort_by(|left, right| left.key.cmp(&right.key));
205 let sample_bytes = samples.iter().try_fold(0usize, |total, sample| {
206 total
207 .checked_add(sample.key.len())
208 .and_then(|value| value.checked_add(sample.vector.len().checked_mul(4)?))
209 .and_then(|value| value.checked_add(std::mem::size_of::<TrainingSample>()))
210 .ok_or_else(|| resource_limit("training_bytes", usize::MAX, usize::MAX))
211 })?;
212 let assignment_bytes = samples
213 .len()
214 .checked_mul(config.subquantizers as usize)
215 .ok_or_else(|| resource_limit("training_bytes", usize::MAX, usize::MAX))?;
216 let centroid_components = (map.tree().config.dimensions as usize)
217 .checked_mul(usize::from(config.centroids_per_subquantizer))
218 .ok_or_else(|| resource_limit("training_bytes", usize::MAX, usize::MAX))?;
219 let centroid_bytes = centroid_components
220 .checked_mul(4 + 8)
221 .and_then(|value| {
222 value.checked_add(
223 (config.subquantizers as usize)
224 .checked_mul(usize::from(config.centroids_per_subquantizer))?
225 .checked_mul(std::mem::size_of::<usize>())?,
226 )
227 })
228 .ok_or_else(|| resource_limit("training_bytes", usize::MAX, usize::MAX))?;
229 let training_bytes = sample_bytes
230 .checked_add(assignment_bytes)
231 .and_then(|value| value.checked_add(centroid_bytes))
232 .ok_or_else(|| resource_limit("training_bytes", usize::MAX, usize::MAX))?;
233 enforce_resource("training_bytes", limits.max_training_bytes, training_bytes)?;
234 let expected_training_evaluations = samples
235 .len()
236 .checked_mul(config.subquantizers as usize)
237 .and_then(|value| value.checked_mul(usize::from(config.centroids_per_subquantizer)))
238 .and_then(|value| value.checked_mul(usize::from(config.training_iterations)))
239 .ok_or_else(|| resource_limit("distance_evaluations", usize::MAX, usize::MAX))?;
240 enforce_resource(
241 "distance_evaluations",
242 limits.max_distance_evaluations,
243 expected_training_evaluations,
244 )?;
245 let vectors: Vec<_> = samples
246 .iter()
247 .map(|sample| sample.vector.as_slice())
248 .collect();
249 let (codebooks, training_evaluations) =
250 train(&vectors, map.tree().config.dimensions, &config, parallelism)?;
251
252 let store = map.store_clone();
253 let code_config = code_tree_config();
254 let mut builder = SortedBatchBuilder::new_with_origin(
255 store.clone(),
256 code_config.clone(),
257 PublicationOrigin::Maintenance,
258 );
259 let layout = subspace_layout(
260 map.tree().config.dimensions as usize,
261 config.subquantizers as usize,
262 );
263 let encoding_per_vector = layout
264 .len()
265 .checked_mul(usize::from(config.centroids_per_subquantizer))
266 .ok_or_else(|| resource_limit("distance_evaluations", usize::MAX, usize::MAX))?;
267 let mut encoding_evaluations = 0usize;
268 let mut encoded_output_bytes = 0usize;
269 let mut quality_sum = 0.0f64;
270 let mut quality_maximum = 0.0f64;
271 let mut encoded_vectors = 0usize;
272 for entry in map
273 .directory_manager()
274 .range(&map.tree().directory, &[], None)?
275 {
276 let (key, bytes) = entry?;
277 let stored = StoredRecord::decode(&bytes, map.tree().config.dimensions)?;
278 encoding_evaluations = encoding_evaluations
279 .checked_add(encoding_per_vector)
280 .ok_or_else(|| resource_limit("distance_evaluations", usize::MAX, usize::MAX))?;
281 let total_evaluations = training_evaluations
282 .checked_add(encoding_evaluations)
283 .ok_or_else(|| resource_limit("distance_evaluations", usize::MAX, usize::MAX))?;
284 enforce_resource(
285 "distance_evaluations",
286 limits.max_distance_evaluations,
287 total_evaluations,
288 )?;
289 let code = encode_vector(&stored.vector, &layout, &codebooks);
290 let error = reconstruction_error(&stored.vector, &code, &codebooks);
291 quality_sum += error;
292 quality_maximum = quality_maximum.max(error);
293 encoded_output_bytes = encoded_output_bytes
294 .checked_add(key.len())
295 .and_then(|value| value.checked_add(code.len()))
296 .ok_or_else(|| resource_limit("encoded_output_bytes", usize::MAX, usize::MAX))?;
297 enforce_resource(
298 "encoded_output_bytes",
299 limits.max_encoded_output_bytes,
300 encoded_output_bytes,
301 )?;
302 builder.add(key, code)?;
303 encoded_vectors += 1;
304 }
305 let code_tree = builder.build()?;
306 let code_root = code_tree
307 .root
308 .clone()
309 .ok_or_else(|| invalid_object("product quantization requires a non-empty code tree"))?;
310 let quality = ProductQuantizationQuality {
311 mean_squared_error: quality_sum / encoded_vectors as f64,
312 maximum_squared_error: quality_maximum,
313 };
314 let manifest_object = Manifest {
315 source: map.tree().descriptor.clone(),
316 dimensions: map.tree().config.dimensions,
317 metric: map.tree().config.metric,
318 count: map.tree().count,
319 config: config.clone(),
320 code_root,
321 codebooks: codebooks.clone(),
322 quality,
323 sampling_hash_algorithm: SAMPLING_HASH_ALGORITHM_XXH64,
324 sampling_hash_version: SAMPLING_HASH_VERSION,
325 training_sample_count: samples.len() as u64,
326 };
327 let manifest_bytes = manifest_object.encode()?;
328 let manifest = Cid::from_bytes(&manifest_bytes);
329 let existed = store
330 .get(manifest.as_bytes())
331 .map_err(|error| Error::Store(Box::new(error)))?;
332 if let Some(bytes) = existed {
333 let actual = Cid::from_bytes(&bytes);
334 if actual != manifest {
335 return Err(Error::CidMismatch {
336 expected: manifest,
337 actual,
338 });
339 }
340 } else {
341 let entries = [(manifest.as_bytes(), manifest_bytes.as_slice())];
342 store
343 .publish_nodes(NodePublication::new(
344 &entries,
345 PublicationOrigin::Maintenance,
346 ))
347 .map_err(|error| Error::Store(Box::new(error)))?;
348 }
349 let stats = ProductQuantizationBuildStats {
350 training_distance_evaluations: training_evaluations,
351 encoding_distance_evaluations: encoding_evaluations,
352 encoded_vectors,
353 training_vectors: samples.len(),
354 training_bytes,
355 encoded_output_bytes,
356 };
357 Ok((
358 Self {
359 codes: Prolly::new(store, code_config),
360 code_tree,
361 manifest,
362 source: manifest_object.source,
363 dimensions: manifest_object.dimensions,
364 metric: manifest_object.metric,
365 count: manifest_object.count,
366 config,
367 codebooks,
368 quality,
369 },
370 stats,
371 ))
372 }
373
374 pub fn load(store: S, manifest: Cid) -> Result<Self, Error> {
376 let bytes = load_content(&store, &manifest)?;
377 let object = Manifest::decode(&bytes)?;
378 object.config.validate(
379 object.dimensions,
380 usize::from(object.config.centroids_per_subquantizer),
381 )?;
382 let code_tree = Tree {
383 root: Some(object.code_root),
384 config: code_tree_config(),
385 };
386 let root = code_tree.root.as_ref().expect("manifest code root");
388 load_content(&store, root)?;
389 Ok(Self {
390 codes: Prolly::new(store.clone(), code_tree.config.clone()),
391 code_tree,
392 manifest,
393 source: object.source,
394 dimensions: object.dimensions,
395 metric: object.metric,
396 count: object.count,
397 config: object.config,
398 codebooks: object.codebooks,
399 quality: object.quality,
400 })
401 }
402
403 pub fn manifest_cid(&self) -> &Cid {
404 &self.manifest
405 }
406
407 pub fn source_descriptor(&self) -> &Cid {
408 &self.source
409 }
410
411 pub fn config(&self) -> &ProductQuantizationConfig {
412 &self.config
413 }
414
415 pub fn quality(&self) -> ProductQuantizationQuality {
416 self.quality
417 }
418
419 pub(crate) fn rebind<T: Store>(&self, store: T) -> ProductQuantizer<T> {
420 ProductQuantizer {
421 codes: Prolly::new(store, self.code_tree.config.clone()),
422 code_tree: self.code_tree.clone(),
423 manifest: self.manifest.clone(),
424 source: self.source.clone(),
425 dimensions: self.dimensions,
426 metric: self.metric,
427 count: self.count,
428 config: self.config.clone(),
429 codebooks: self.codebooks.clone(),
430 quality: self.quality,
431 }
432 }
433
434 pub fn search(
436 &self,
437 map: &ProximityMap<S>,
438 request: SearchRequest<'_>,
439 ) -> Result<SearchResult, Error> {
440 request.validate()?;
441 if request.policy == SearchPolicy::Exact {
442 return Err(invalid_search(
443 "product quantization cannot satisfy exact search",
444 ));
445 }
446 if !matches!(
447 request.options.backend,
448 SearchBackend::ProductQuantized | SearchBackend::Auto
449 ) {
450 return Err(invalid_search(
451 "product quantizer requires ProductQuantized or Auto backend",
452 ));
453 }
454 let filter = PreparedFilter::new(request.filter.clone(), &map.tree().directory)?;
455 let eligible_limit = match filter.cardinality(map.tree().count) {
456 EligibilityCardinality::Known(count) => count as usize,
457 EligibilityCardinality::Unknown => map.tree().count as usize,
458 };
459 let multiplier = request
460 .options
461 .pq
462 .rerank_multiplier
463 .map(usize::from)
464 .unwrap_or(self.config.rerank_multiplier as usize);
465 let plan = crate::prolly::proximity::search::SearchPlan::ProductQuantized {
466 rerank_target: request
467 .k
468 .saturating_mul(multiplier)
469 .max(request.k)
470 .min(eligible_limit),
471 direct_lookup: filter.sorted_keys().is_some()
472 && eligible_limit <= request.options.planner.eligible_exact_max_records,
473 };
474 self.search_planned(map, request, &plan)
475 }
476
477 pub(crate) fn search_planned(
478 &self,
479 map: &ProximityMap<S>,
480 request: SearchRequest<'_>,
481 plan: &crate::prolly::proximity::search::SearchPlan,
482 ) -> Result<SearchResult, Error> {
483 self.search_planned_with_exclusion(map, &map.tree().descriptor, request, plan, |_| {
484 Ok(false)
485 })
486 }
487
488 pub(crate) fn search_planned_with_exclusion<F>(
489 &self,
490 map: &ProximityMap<S>,
491 expected_source: &Cid,
492 request: SearchRequest<'_>,
493 plan: &crate::prolly::proximity::search::SearchPlan,
494 mut excluded: F,
495 ) -> Result<SearchResult, Error>
496 where
497 F: FnMut(&[u8]) -> Result<bool, Error>,
498 {
499 let crate::prolly::proximity::search::SearchPlan::ProductQuantized {
500 rerank_target,
501 direct_lookup,
502 } = plan
503 else {
504 return Err(invalid_search(
505 "product quantization executor requires a PQ search plan",
506 ));
507 };
508 request.validate()?;
509 if request.policy == SearchPolicy::Exact {
510 return Err(invalid_search(
511 "product quantization cannot satisfy exact search",
512 ));
513 }
514 if &self.source != expected_source
515 || self.dimensions != map.tree().config.dimensions
516 || self.metric != map.tree().config.metric
517 {
518 return Err(invalid_search(
519 "product quantizer is bound to a different source descriptor",
520 ));
521 }
522
523 let query = prepare_vector(self.metric, request.query, self.dimensions)?;
524 let filter = PreparedFilter::new(request.filter.clone(), &map.tree().directory)?;
525 let lookup = build_lookup(&query, self.metric, &self.codebooks);
526 let mut stats = ProximitySearchStats::default();
527 let mut approximate = BinaryHeap::<PqRanked>::new();
528 let mut completion = SearchCompletion::ApproximatePolicySatisfied;
529 if *direct_lookup {
530 let Some((keys, source_bound)) = filter.sorted_keys() else {
531 return Err(invalid_search(
532 "PQ direct-lookup plan requires sorted eligible keys",
533 ));
534 };
535 for key in keys {
536 if excluded(key)? {
537 continue;
538 }
539 let code = self.codes.get(&self.code_tree, key)?;
540 let Some(code) = code else {
541 if source_bound {
542 return Err(invalid_object("source-bound eligible key has no PQ code"));
543 }
544 continue;
545 };
546 if !admit_code(
547 key.clone(),
548 code,
549 &lookup,
550 self.metric,
551 &self.codebooks,
552 *rerank_target,
553 &request,
554 &mut stats,
555 &mut approximate,
556 )? {
557 completion = SearchCompletion::BudgetExhausted;
558 break;
559 }
560 }
561 } else {
562 for entry in self.codes.range(&self.code_tree, &[], None)? {
563 let (key, code) = entry?;
564 if !filter.contains(&key) || excluded(&key)? {
565 continue;
566 }
567 if !admit_code(
568 key,
569 code,
570 &lookup,
571 self.metric,
572 &self.codebooks,
573 *rerank_target,
574 &request,
575 &mut stats,
576 &mut approximate,
577 )? {
578 completion = SearchCompletion::BudgetExhausted;
579 break;
580 }
581 }
582 }
583 let mut approximate = approximate.into_vec();
584 approximate.sort();
585 let shortlist = approximate.len();
586
587 let mut reranked = Vec::<RerankCandidate>::with_capacity(shortlist);
588 let mut vector_scratch = vec![0.0f32; map.tree().config.dimensions as usize];
589 let mut directory = map.directory_manager().read(&map.tree().directory)?;
590 for candidate in approximate {
591 if budget_exhausted(&request, &stats)
592 || request
593 .budget
594 .max_nodes
595 .is_some_and(|limit| stats.nodes_read >= limit)
596 {
597 completion = SearchCompletion::BudgetExhausted;
598 break;
599 }
600 let Some(handle) = directory.get_handle(&candidate.key)? else {
601 return Err(invalid_object(
602 "PQ code key is absent from authoritative directory",
603 ));
604 };
605 let bytes = handle.value()?.len();
606 if request
607 .budget
608 .max_committed_bytes
609 .is_some_and(|limit| stats.committed_bytes.saturating_add(bytes) > limit)
610 {
611 completion = SearchCompletion::BudgetExhausted;
612 break;
613 }
614 let record = crate::prolly::proximity::storage::StoredRecordRef::decode(
615 handle.value()?,
616 map.tree().config.dimensions,
617 )?;
618 crate::prolly::proximity::ProximityVectorRef::from_encoded(record.vector)
619 .copy_to_slice(&mut vector_scratch)?;
620 let distance = query_score(request.kernel, self.metric, &query, &vector_scratch);
621 stats.nodes_read += 1;
622 stats.bytes_read = stats.bytes_read.saturating_add(bytes);
623 stats.committed_bytes = stats.committed_bytes.saturating_add(bytes);
624 stats.distance_evaluations += 1;
625 reranked.push(RerankCandidate::new(handle, &candidate.key, distance)?);
626 }
627 stats.reranked_candidates = reranked.len();
628 stats.candidate_handles_peak = reranked.len();
629 stats.candidate_retained_bytes_peak = retained_candidate_bytes(&reranked);
630 reranked.sort_by(|left, right| {
631 left.distance
632 .total_cmp(&right.distance)
633 .then_with(|| left.key().cmp(right.key()))
634 });
635 let neighbors = reranked
636 .into_iter()
637 .take(request.k)
638 .map(|candidate| candidate.into_neighbor(map.tree().config.dimensions))
639 .collect::<Result<Vec<_>, Error>>()?;
640 Ok(SearchResult {
641 neighbors,
642 stats,
643 completion,
644 plan: plan.summary(),
645 })
646 }
647}
648
649#[derive(Clone, Debug)]
650struct PqRanked {
651 distance: f64,
652 key: Vec<u8>,
653}
654
655impl PartialEq for PqRanked {
656 fn eq(&self, other: &Self) -> bool {
657 self.distance.to_bits() == other.distance.to_bits() && self.key == other.key
658 }
659}
660
661impl Eq for PqRanked {}
662
663impl PartialOrd for PqRanked {
664 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
665 Some(self.cmp(other))
666 }
667}
668
669impl Ord for PqRanked {
670 fn cmp(&self, other: &Self) -> Ordering {
671 self.distance
672 .total_cmp(&other.distance)
673 .then_with(|| self.key.cmp(&other.key))
674 }
675}
676
677#[allow(clippy::too_many_arguments)]
678fn admit_code(
679 key: Vec<u8>,
680 code: Vec<u8>,
681 lookup: &[Vec<f64>],
682 metric: DistanceMetric,
683 codebooks: &Codebooks,
684 target: usize,
685 request: &SearchRequest<'_>,
686 stats: &mut ProximitySearchStats,
687 approximate: &mut BinaryHeap<PqRanked>,
688) -> Result<bool, Error> {
689 if request
690 .budget
691 .max_nodes
692 .is_some_and(|limit| stats.nodes_read >= limit)
693 || request
694 .budget
695 .max_committed_bytes
696 .is_some_and(|limit| stats.committed_bytes.saturating_add(code.len()) > limit)
697 || request
698 .budget
699 .max_distance_evaluations
700 .is_some_and(|limit| {
701 stats
702 .distance_evaluations
703 .saturating_add(stats.quantized_distance_evaluations)
704 >= limit
705 })
706 || request
707 .budget
708 .max_frontier_entries
709 .is_some_and(|limit| approximate.len().saturating_add(1) > limit)
710 {
711 return Ok(false);
712 }
713 validate_code(&code, codebooks)?;
714 stats.nodes_read += 1;
715 stats.bytes_read = stats.bytes_read.saturating_add(code.len());
716 stats.committed_bytes = stats.committed_bytes.saturating_add(code.len());
717 stats.quantized_distance_evaluations += 1;
718 approximate.push(PqRanked {
719 distance: score_code(metric, lookup, &code),
720 key,
721 });
722 if approximate.len() > target {
723 approximate.pop();
724 }
725 stats.frontier_peak = stats.frontier_peak.max(approximate.len());
726 Ok(true)
727}
728
729#[derive(Clone)]
730pub(crate) struct Manifest {
731 pub(crate) source: Cid,
732 pub(crate) dimensions: u32,
733 pub(crate) metric: DistanceMetric,
734 pub(crate) count: u64,
735 pub(crate) config: ProductQuantizationConfig,
736 pub(crate) code_root: Cid,
737 pub(crate) codebooks: Codebooks,
738 pub(crate) quality: ProductQuantizationQuality,
739 pub(crate) sampling_hash_algorithm: u8,
740 pub(crate) sampling_hash_version: u8,
741 pub(crate) training_sample_count: u64,
742}
743
744impl Manifest {
745 fn encode(&self) -> Result<Vec<u8>, Error> {
746 let mut bytes = Vec::new();
747 bytes.extend_from_slice(MAGIC);
748 bytes.push(PQ_FORMAT_VERSION);
749 bytes.push(0);
750 put_cid(&self.source, &mut bytes);
751 put_varint(u64::from(self.dimensions), &mut bytes);
752 bytes.push(self.metric.id());
753 put_varint(self.count, &mut bytes);
754 bytes.push(self.sampling_hash_algorithm);
755 bytes.push(self.sampling_hash_version);
756 put_varint(self.training_sample_count, &mut bytes);
757 encode_config(&self.config, &mut bytes);
758 put_cid(&self.code_root, &mut bytes);
759 put_varint(self.codebooks.len() as u64, &mut bytes);
760 for subspace in &self.codebooks {
761 put_varint(subspace.len() as u64, &mut bytes);
762 let width = subspace.first().map_or(0, Vec::len);
763 put_varint(width as u64, &mut bytes);
764 for centroid in subspace {
765 if centroid.len() != width {
766 return Err(invalid_object("inconsistent PQ centroid width"));
767 }
768 for &component in centroid {
769 put_f32(component, &mut bytes)?;
770 }
771 }
772 }
773 put_f64(self.quality.mean_squared_error, &mut bytes)?;
774 put_f64(self.quality.maximum_squared_error, &mut bytes)?;
775 put_cid(&config_fingerprint(&self.config), &mut bytes);
776 Ok(bytes)
777 }
778
779 pub(crate) fn decode(bytes: &[u8]) -> Result<Self, Error> {
780 let mut reader = Reader::new(bytes, "product quantizer");
781 reader.exact(MAGIC)?;
782 require_pq_version(reader.u8()?)?;
783 if reader.u8()? != 0 {
784 return Err(reader.invalid("unknown flags"));
785 }
786 let source = reader.cid()?;
787 let dimensions =
788 u32::try_from(reader.varint()?).map_err(|_| reader.invalid("dimensions exceed u32"))?;
789 let metric = DistanceMetric::from_id(reader.u8()?)?;
790 let count = reader.varint()?;
791 if count == 0 {
792 return Err(reader.invalid("PQ source count must be positive"));
793 }
794 let sampling_hash_algorithm = reader.u8()?;
795 let sampling_hash_version = reader.u8()?;
796 let training_sample_count = reader.varint()?;
797 if sampling_hash_algorithm != SAMPLING_HASH_ALGORITHM_XXH64
798 || sampling_hash_version != SAMPLING_HASH_VERSION
799 || training_sample_count == 0
800 || training_sample_count > count
801 {
802 return Err(reader.invalid("unsupported or invalid PQ sampling policy"));
803 }
804 let config = decode_config(&mut reader)?;
805 if training_sample_count != count.min(config.max_training_vectors as u64) {
806 return Err(reader.invalid("PQ training sample count disagrees with configuration"));
807 }
808 let code_root = reader.cid()?;
809 let subspaces = reader.bounded_usize(MAX_OBJECT_ENTRIES)?;
810 if subspaces != config.subquantizers as usize {
811 return Err(reader.invalid("subquantizer count mismatch"));
812 }
813 let mut codebooks = Vec::with_capacity(subspaces);
814 let mut total_width = 0usize;
815 for _ in 0..subspaces {
816 let centroids = reader.bounded_usize(256)?;
817 let width = reader.bounded_usize(dimensions as usize)?;
818 if centroids != usize::from(config.centroids_per_subquantizer) || width == 0 {
819 return Err(reader.invalid("PQ codebook shape mismatch"));
820 }
821 let count = centroids
822 .checked_mul(width)
823 .ok_or_else(|| reader.invalid("PQ codebook length overflow"))?;
824 if count
825 .checked_mul(4)
826 .map_or(true, |len| len > reader.remaining())
827 {
828 return Err(reader.invalid("impossible PQ codebook length"));
829 }
830 let mut subspace = Vec::with_capacity(centroids);
831 for _ in 0..centroids {
832 let mut centroid = Vec::with_capacity(width);
833 for _ in 0..width {
834 centroid.push(reader.f32()?);
835 }
836 subspace.push(centroid);
837 }
838 total_width = total_width
839 .checked_add(width)
840 .ok_or_else(|| reader.invalid("PQ dimension overflow"))?;
841 codebooks.push(subspace);
842 }
843 if total_width != dimensions as usize {
844 return Err(reader.invalid("PQ subspaces do not cover dimensions"));
845 }
846 let quality = ProductQuantizationQuality {
847 mean_squared_error: reader.f64()?,
848 maximum_squared_error: reader.f64()?,
849 };
850 if quality.mean_squared_error < 0.0 || quality.maximum_squared_error < 0.0 {
851 return Err(reader.invalid("negative PQ quality measurement"));
852 }
853 if reader.cid()? != config_fingerprint(&config) {
854 return Err(reader.invalid("PQ configuration fingerprint mismatch"));
855 }
856 reader.finish()?;
857 Ok(Self {
858 source,
859 dimensions,
860 metric,
861 count,
862 config,
863 code_root,
864 codebooks,
865 quality,
866 sampling_hash_algorithm,
867 sampling_hash_version,
868 training_sample_count,
869 })
870 }
871}
872
873#[derive(Clone, Debug)]
874struct TrainingSample {
875 hash: u64,
876 key: Vec<u8>,
877 vector: Vec<f32>,
878}
879
880impl PartialEq for TrainingSample {
881 fn eq(&self, other: &Self) -> bool {
882 self.hash == other.hash && self.key == other.key
883 }
884}
885
886impl Eq for TrainingSample {}
887
888impl PartialOrd for TrainingSample {
889 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
890 Some(self.cmp(other))
891 }
892}
893
894impl Ord for TrainingSample {
895 fn cmp(&self, other: &Self) -> Ordering {
896 self.hash
897 .cmp(&other.hash)
898 .then_with(|| self.key.cmp(&other.key))
899 }
900}
901
902fn train(
903 vectors: &[&[f32]],
904 dimensions: u32,
905 config: &ProductQuantizationConfig,
906 parallelism: BuildParallelism,
907) -> Result<TrainingOutput, Error> {
908 let layout = subspace_layout(dimensions as usize, config.subquantizers as usize);
909 let centroid_count = usize::from(config.centroids_per_subquantizer);
910 let mut codebooks = Vec::with_capacity(layout.len());
911 for (subspace, &(start, end)) in layout.iter().enumerate() {
912 let mut used = HashSet::new();
913 let mut centroids = Vec::with_capacity(centroid_count);
914 for centroid in 0..centroid_count {
915 let mut identity = [0u8; 16];
916 identity[..8].copy_from_slice(&(subspace as u64).to_le_bytes());
917 identity[8..].copy_from_slice(&(centroid as u64).to_le_bytes());
918 let initial = (xxh64(&identity, config.seed) as usize) % vectors.len();
919 let selected = (0..vectors.len())
920 .map(|offset| (initial + offset) % vectors.len())
921 .find(|candidate| used.insert(*candidate))
922 .expect("centroid count does not exceed records");
923 centroids.push(vectors[selected][start..end].to_vec());
924 }
925 codebooks.push(centroids);
926 }
927
928 let pool = (parallelism.threads() > 1)
929 .then(|| {
930 rayon::ThreadPoolBuilder::new()
931 .num_threads(parallelism.threads())
932 .build()
933 })
934 .transpose()
935 .map_err(|error| invalid_config(format!("cannot create PQ worker pool: {error}")))?;
936 let mut evaluations = 0usize;
937 for _ in 0..config.training_iterations {
938 let assignments = assign_all(vectors, &layout, &codebooks, pool.as_ref());
939 evaluations = evaluations.saturating_add(
940 vectors
941 .len()
942 .saturating_mul(layout.len())
943 .saturating_mul(centroid_count),
944 );
945 for (subspace, &(start, end)) in layout.iter().enumerate() {
946 let width = end - start;
947 let mut sums = vec![vec![0.0f64; width]; centroid_count];
948 let mut counts = vec![0usize; centroid_count];
949 for (vector_index, vector) in vectors.iter().enumerate() {
950 let centroid = assignments[vector_index][subspace] as usize;
951 counts[centroid] += 1;
952 for (offset, &component) in vector[start..end].iter().enumerate() {
953 sums[centroid][offset] += f64::from(component);
954 }
955 }
956 for centroid in 0..centroid_count {
957 if counts[centroid] == 0 {
958 continue;
959 }
960 for offset in 0..width {
961 let value = (sums[centroid][offset] / counts[centroid] as f64) as f32;
962 codebooks[subspace][centroid][offset] = if value == 0.0 { 0.0 } else { value };
963 }
964 }
965 }
966 }
967 Ok((codebooks, evaluations))
968}
969
970fn assign_all(
971 vectors: &[&[f32]],
972 layout: &[(usize, usize)],
973 codebooks: &[Vec<Vec<f32>>],
974 pool: Option<&rayon::ThreadPool>,
975) -> Vec<Vec<u8>> {
976 let compute = || {
977 vectors
978 .par_iter()
979 .map(|vector| {
980 layout
981 .iter()
982 .zip(codebooks)
983 .map(|(&(start, end), centroids)| {
984 nearest_centroid(&vector[start..end], centroids)
985 })
986 .collect()
987 })
988 .collect()
989 };
990 if let Some(pool) = pool {
991 pool.install(compute)
992 } else {
993 vectors
994 .iter()
995 .map(|vector| {
996 layout
997 .iter()
998 .zip(codebooks)
999 .map(|(&(start, end), centroids)| {
1000 nearest_centroid(&vector[start..end], centroids)
1001 })
1002 .collect()
1003 })
1004 .collect()
1005 }
1006}
1007
1008fn nearest_centroid(vector: &[f32], centroids: &[Vec<f32>]) -> u8 {
1009 let mut best = (0usize, f64::INFINITY);
1010 for (index, centroid) in centroids.iter().enumerate() {
1011 let distance = vector.iter().zip(centroid).fold(0.0, |sum, (&a, &b)| {
1012 let delta = f64::from(a) - f64::from(b);
1013 sum + delta * delta
1014 });
1015 if distance
1016 .total_cmp(&best.1)
1017 .then_with(|| index.cmp(&best.0))
1018 .is_lt()
1019 {
1020 best = (index, distance);
1021 }
1022 }
1023 best.0 as u8
1024}
1025
1026fn encode_vector(
1027 vector: &[f32],
1028 layout: &[(usize, usize)],
1029 codebooks: &[Vec<Vec<f32>>],
1030) -> Vec<u8> {
1031 layout
1032 .iter()
1033 .zip(codebooks)
1034 .map(|(&(start, end), centroids)| nearest_centroid(&vector[start..end], centroids))
1035 .collect()
1036}
1037
1038fn subspace_layout(dimensions: usize, subquantizers: usize) -> Vec<(usize, usize)> {
1039 (0..subquantizers)
1040 .map(|index| {
1041 (
1042 index * dimensions / subquantizers,
1043 (index + 1) * dimensions / subquantizers,
1044 )
1045 })
1046 .collect()
1047}
1048
1049fn reconstruction_error(vector: &[f32], code: &[u8], codebooks: &[Vec<Vec<f32>>]) -> f64 {
1050 let mut offset = 0usize;
1051 let mut error = 0.0f64;
1052 for (subspace, ¢roid) in codebooks.iter().zip(code) {
1053 for &component in &subspace[centroid as usize] {
1054 let delta = f64::from(vector[offset]) - f64::from(component);
1055 error += delta * delta;
1056 offset += 1;
1057 }
1058 }
1059 error
1060}
1061
1062pub(crate) fn build_lookup(
1063 query: &[f32],
1064 metric: DistanceMetric,
1065 codebooks: &[Vec<Vec<f32>>],
1066) -> Vec<Vec<f64>> {
1067 let mut offset = 0usize;
1068 codebooks
1069 .iter()
1070 .map(|subspace| {
1071 let width = subspace.first().map_or(0, Vec::len);
1072 let query = &query[offset..offset + width];
1073 offset += width;
1074 subspace
1075 .iter()
1076 .map(|centroid| match metric {
1077 DistanceMetric::L2Squared => {
1078 query.iter().zip(centroid).fold(0.0, |sum, (&a, &b)| {
1079 let delta = f64::from(a) - f64::from(b);
1080 sum + delta * delta
1081 })
1082 }
1083 DistanceMetric::Cosine | DistanceMetric::InnerProduct => query
1084 .iter()
1085 .zip(centroid)
1086 .fold(0.0, |sum, (&a, &b)| sum + f64::from(a) * f64::from(b)),
1087 })
1088 .collect()
1089 })
1090 .collect()
1091}
1092
1093pub(crate) fn score_code(metric: DistanceMetric, lookup: &[Vec<f64>], code: &[u8]) -> f64 {
1094 let reduced = lookup
1095 .iter()
1096 .zip(code)
1097 .fold(0.0, |sum, (subspace, ¢roid)| {
1098 sum + subspace[centroid as usize]
1099 });
1100 let result = match metric {
1101 DistanceMetric::L2Squared => reduced,
1102 DistanceMetric::Cosine => 1.0 - reduced.clamp(-1.0, 1.0),
1103 DistanceMetric::InnerProduct => -reduced,
1104 };
1105 if result == 0.0 {
1106 0.0
1107 } else {
1108 result
1109 }
1110}
1111
1112pub(crate) fn validate_code(code: &[u8], codebooks: &[Vec<Vec<f32>>]) -> Result<(), Error> {
1113 if code.len() != codebooks.len()
1114 || code
1115 .iter()
1116 .zip(codebooks)
1117 .any(|(¢roid, subspace)| centroid as usize >= subspace.len())
1118 {
1119 return Err(invalid_object("invalid PQ vector code"));
1120 }
1121 Ok(())
1122}
1123
1124fn budget_exhausted(request: &SearchRequest<'_>, stats: &ProximitySearchStats) -> bool {
1125 request
1126 .budget
1127 .max_distance_evaluations
1128 .is_some_and(|maximum| {
1129 stats
1130 .distance_evaluations
1131 .saturating_add(stats.quantized_distance_evaluations)
1132 >= maximum
1133 })
1134}
1135
1136fn encode_config(config: &ProductQuantizationConfig, bytes: &mut Vec<u8>) {
1137 put_varint(u64::from(config.subquantizers), bytes);
1138 put_varint(u64::from(config.centroids_per_subquantizer), bytes);
1139 put_varint(u64::from(config.training_iterations), bytes);
1140 put_varint(u64::from(config.rerank_multiplier), bytes);
1141 bytes.extend_from_slice(&config.seed.to_le_bytes());
1142 put_varint(config.max_training_vectors as u64, bytes);
1143}
1144
1145fn decode_config(reader: &mut Reader<'_>) -> Result<ProductQuantizationConfig, Error> {
1146 Ok(ProductQuantizationConfig {
1147 subquantizers: u32::try_from(reader.varint()?)
1148 .map_err(|_| reader.invalid("subquantizers exceed u32"))?,
1149 centroids_per_subquantizer: u16::try_from(reader.varint()?)
1150 .map_err(|_| reader.invalid("centroid count exceeds u16"))?,
1151 training_iterations: u16::try_from(reader.varint()?)
1152 .map_err(|_| reader.invalid("training iterations exceed u16"))?,
1153 rerank_multiplier: u32::try_from(reader.varint()?)
1154 .map_err(|_| reader.invalid("rerank multiplier exceeds u32"))?,
1155 seed: reader.u64_le()?,
1156 max_training_vectors: usize::try_from(reader.varint()?)
1157 .map_err(|_| reader.invalid("max training vectors exceed usize"))?,
1158 })
1159}
1160
1161fn require_pq_version(found: u8) -> Result<(), Error> {
1162 if found == PQ_FORMAT_VERSION {
1163 Ok(())
1164 } else {
1165 Err(Error::UnsupportedProximityVersion {
1166 found,
1167 required: PQ_FORMAT_VERSION,
1168 })
1169 }
1170}
1171
1172fn enforce_resource(
1173 resource: &'static str,
1174 limit: Option<usize>,
1175 actual: usize,
1176) -> Result<(), Error> {
1177 if let Some(limit) = limit {
1178 if actual > limit {
1179 return Err(resource_limit(resource, limit, actual));
1180 }
1181 }
1182 Ok(())
1183}
1184
1185fn resource_limit(resource: &'static str, limit: usize, actual: usize) -> Error {
1186 Error::ProximityResourceLimitExceeded {
1187 resource,
1188 limit,
1189 actual,
1190 }
1191}
1192
1193pub(crate) fn config_fingerprint(config: &ProductQuantizationConfig) -> Cid {
1194 let mut bytes = Vec::new();
1195 encode_config(config, &mut bytes);
1196 Cid::from_bytes(&bytes)
1197}
1198
1199pub(crate) fn code_tree_config() -> Config {
1200 Config::builder()
1203 .min_chunk_size(4)
1204 .max_chunk_size(1024 * 1024)
1205 .chunking_factor(128)
1206 .hash_seed(0)
1207 .encoding(Encoding::Raw)
1208 .build()
1209}
1210
1211fn load_content<S: Store>(store: &S, cid: &Cid) -> Result<Vec<u8>, Error> {
1212 let bytes = store
1213 .get(cid.as_bytes())
1214 .map_err(|error| Error::Store(Box::new(error)))?
1215 .ok_or_else(|| Error::NotFound(cid.clone()))?;
1216 let actual = Cid::from_bytes(&bytes);
1217 if actual != *cid {
1218 return Err(Error::CidMismatch {
1219 expected: cid.clone(),
1220 actual,
1221 });
1222 }
1223 Ok(bytes)
1224}
1225
1226fn invalid_config(reason: impl Into<String>) -> Error {
1227 Error::InvalidProximityConfig {
1228 reason: reason.into(),
1229 }
1230}
1231
1232fn invalid_object(reason: impl Into<String>) -> Error {
1233 Error::InvalidProximityObject {
1234 kind: "product quantizer",
1235 reason: reason.into(),
1236 }
1237}
1238
1239fn invalid_search(reason: impl Into<String>) -> Error {
1240 Error::InvalidProximitySearch {
1241 reason: reason.into(),
1242 }
1243}
1244
1245#[cfg(test)]
1246mod tests {
1247 use super::*;
1248
1249 #[test]
1250 fn stable_ties_choose_the_lowest_centroid() {
1251 assert_eq!(nearest_centroid(&[0.0], &[vec![-1.0], vec![1.0]]), 0);
1252 }
1253
1254 #[test]
1255 fn uneven_subspaces_cover_every_dimension_once() {
1256 assert_eq!(subspace_layout(7, 3), vec![(0, 2), (2, 4), (4, 7)]);
1257 }
1258
1259 #[test]
1260 fn v1_manifest_requires_rebuild() {
1261 assert!(matches!(
1262 Manifest::decode(b"PQPQ\x01"),
1263 Err(Error::UnsupportedProximityVersion {
1264 found: 1,
1265 required: 2
1266 })
1267 ));
1268 }
1269}