diskann_disk/build/builder/
quantizer.rs1use 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#[derive(Clone)]
26pub enum BuildQuantizer {
27 NoQuant(NoStore),
28 Scalar1Bit(WithBits<1>),
29 PQ(FixedChunkPQTable),
30}
31
32impl BuildQuantizer {
33 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 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 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 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}