Skip to main content

diskann_disk/build/builder/
quantizer.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5//! Disk index quantizer implementation.
6use crate::data_model::GraphDataType;
7use diskann::{ANNError, ANNResult};
8use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
9use diskann_providers::{
10    index::diskann_async::train_pq,
11    model::{
12        graph::provider::async_::{common::NoStore, inmem::WithBits},
13        FixedChunkPQTable, IndexConfiguration, MAX_PQ_TRAINING_SET_SIZE,
14    },
15    storage::{PQStorage, SQStorage},
16    utils::{create_thread_pool, BridgeErr, PQPathNames},
17};
18use diskann_quantization::scalar::train::ScalarQuantizationParameters;
19use diskann_utils::views::MatrixView;
20use tracing::info;
21
22use crate::QuantizationType;
23
24/// Quantizer types used specifically for async disk index building.
25#[derive(Clone)]
26pub enum BuildQuantizer {
27    NoQuant(NoStore),
28    Scalar1Bit(WithBits<1>),
29    PQ(FixedChunkPQTable),
30}
31
32impl BuildQuantizer {
33    /// Train a new quantizer from scratch.
34    pub fn train<Data, StorageProvider>(
35        build_quantization_type: &QuantizationType,
36        index_path_prefix: &str,
37        index_configuration: &IndexConfiguration,
38        pq_storage: &PQStorage,
39        storage_provider: &StorageProvider,
40    ) -> ANNResult<Self>
41    where
42        Data: GraphDataType<VectorIdType = u32>,
43        StorageProvider: StorageReadProvider + StorageWriteProvider,
44    {
45        let num_points = index_configuration.max_points;
46        let p_val = MAX_PQ_TRAINING_SET_SIZE / (num_points as f64);
47        match *build_quantization_type {
48            QuantizationType::FP => Ok(Self::NoQuant(NoStore)),
49            QuantizationType::PQ { num_chunks } => {
50                let table = {
51                    //generate pq pivots.
52                    let seed = index_configuration.random_seed;
53                    let mut rnd =
54                        diskann_providers::utils::create_rnd_provider_from_optional_seed(seed)
55                            .create_rnd();
56                    let (train_data, train_size, train_dim) = pq_storage
57                        .get_random_train_data_slice::<Data::VectorDataType, _>(
58                            p_val,
59                            storage_provider,
60                            &mut rnd,
61                        )?;
62                    train_pq(
63                        MatrixView::try_from(&train_data, train_size, train_dim).bridge_err()?,
64                        num_chunks,
65                        &mut rnd,
66                        create_thread_pool(index_configuration.num_threads)?.as_ref(),
67                    )?
68                };
69                // Save at checkpoint. Note the the compressed data path and pivots path here
70                // are different than the ones used in quant vector generation.
71                let pq_paths = PQPathNames::new(index_path_prefix);
72                let pq_build_storage =
73                    PQStorage::new(&pq_paths.pivots, &pq_paths.compressed_data, None);
74                pq_build_storage.write_pivot_data(
75                    table.get_pq_table(),
76                    None,
77                    table.get_chunk_offsets(),
78                    table.get_num_centers(),
79                    table.get_dim(),
80                    storage_provider,
81                )?;
82                Ok(Self::PQ(table))
83            }
84            QuantizationType::SQ {
85                nbits,
86                standard_deviation,
87            } => {
88                if nbits != 1 {
89                    return Err(ANNError::log_index_config_error(
90                        "build_quantization_type".to_string(),
91                        "SQ quantization is only supported for 1 bit".to_string(),
92                    ));
93                }
94                let rng = diskann_providers::utils::create_rnd_provider_from_optional_seed(
95                    index_configuration.random_seed,
96                );
97                let (train_data_vector, train_size, train_dim) = pq_storage
98                    .get_random_train_data_slice::<Data::VectorDataType, _>(
99                        p_val,
100                        storage_provider,
101                        &mut rng.create_rnd(),
102                    )?;
103
104                let quantizer_params = if let Some(std_dev) = standard_deviation {
105                    ScalarQuantizationParameters::new(std_dev)
106                } else {
107                    ScalarQuantizationParameters::default()
108                };
109
110                let quantizer = quantizer_params.train(
111                    MatrixView::try_from(&train_data_vector, train_size, train_dim).bridge_err()?,
112                );
113
114                info!("Now quantizer is trained and saving to file");
115                let sq_storage = SQStorage::new(index_path_prefix);
116                sq_storage.save_quantizer(&quantizer, storage_provider)?;
117
118                Ok(Self::Scalar1Bit(WithBits::<1>::new(quantizer)))
119            }
120        }
121    }
122
123    /// Load a previously trained quantizer from storage.
124    pub fn load<StorageProvider>(
125        build_quantization_type: &QuantizationType,
126        index_path_prefix: &str,
127        storage_provider: &StorageProvider,
128    ) -> ANNResult<Self>
129    where
130        StorageProvider: StorageReadProvider,
131    {
132        match build_quantization_type {
133            QuantizationType::FP => Ok(Self::NoQuant(NoStore)),
134            QuantizationType::PQ { num_chunks } => {
135                let pq_pivots_paths = PQPathNames::new(index_path_prefix);
136                let pq_build_storage = PQStorage::new(
137                    &pq_pivots_paths.pivots,
138                    &pq_pivots_paths.compressed_data,
139                    None,
140                );
141                let table = pq_build_storage.load_pq_pivots_bin::<StorageProvider>(
142                    &pq_pivots_paths.pivots,
143                    *num_chunks,
144                    storage_provider,
145                )?;
146                Ok(Self::PQ(table))
147            }
148            QuantizationType::SQ { .. } => {
149                let sq_storage = SQStorage::new(index_path_prefix);
150                let sq_quantizer = sq_storage.load_quantizer(storage_provider)?;
151                Ok(Self::Scalar1Bit(WithBits::<1>::new(sq_quantizer)))
152            }
153        }
154    }
155}