Skip to main content

prolly/prolly/proximity/accelerator/
pq.rs

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/// Deterministic offline product-quantization training and serving policy.
34#[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/// Reconstruction measurements committed into a PQ manifest.
81#[derive(Clone, Copy, Debug, Default, PartialEq)]
82pub struct ProductQuantizationQuality {
83    pub mean_squared_error: f64,
84    pub maximum_squared_error: f64,
85}
86
87/// Canonical logical work performed by one PQ build.
88#[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
126/// Source-bound persisted product-quantization sidecar.
127pub 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    /// Train in key order, persist encoded vectors, and publish the manifest last.
146    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    /// Load and validate a persisted PQ manifest and its encoded-vector root.
375    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        // Validate the root eagerly; individual codes are checked during search.
387        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    /// Search the PQ code tree, then rerank the deterministic shortlist using full vectors.
435    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, &centroid) 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, &centroid)| {
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(|(&centroid, 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    // This is a wire-level PQ constant. Do not inherit future changes to
1201    // the general ordered-tree defaults when loading an existing sidecar.
1202    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}