1use std::marker::PhantomData;
7
8use diskann::{utils::VectorRepr, ANNError};
9use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
10use diskann_providers::{
11 model::{
12 pq::{accum_row_inplace, generate_pq_pivots},
13 GeneratePivotArguments,
14 },
15 storage::PQStorage,
16 utils::{BridgeErr, RayonThreadPoolRef, Timer},
17};
18use diskann_quantization::{product::TransposedTable, CompressInto};
19use diskann_utils::views::MatrixBase;
20use diskann_vector::distance::Metric;
21use tracing::info;
22
23use crate::storage::quant::compressor::{CompressionStage, QuantCompressor};
24
25pub struct PQGenerationContext<'a, Storage>
26where
27 Storage: StorageReadProvider + StorageWriteProvider,
28{
29 pub pq_storage: PQStorage,
30 pub num_chunks: usize,
31 pub seed: Option<u64>,
32 pub p_val: f64,
33 pub storage_provider: &'a Storage,
34 pub pool: RayonThreadPoolRef<'a>,
35 pub metric: Metric,
36 pub dim: usize,
37 pub max_kmeans_reps: usize,
38 pub num_centers: usize,
39}
40
41pub struct PQGeneration<'a, T, Storage>
42where
43 T: VectorRepr,
44 Storage: StorageReadProvider + StorageWriteProvider + 'a,
45{
46 table: TransposedTable,
47 num_chunks: usize,
48 phantom_data: PhantomData<T>,
49 phantom_storage: PhantomData<&'a Storage>,
50}
51
52impl<'a, T, Storage> QuantCompressor<T> for PQGeneration<'a, T, Storage>
53where
54 T: VectorRepr,
55 Storage: StorageReadProvider + StorageWriteProvider + 'a,
56{
57 type CompressorContext = PQGenerationContext<'a, Storage>;
58
59 fn new_at_stage(
60 stage: CompressionStage,
61 context: &Self::CompressorContext,
62 ) -> diskann::ANNResult<Self> {
63 if context.num_chunks > context.dim {
65 return Err(ANNError::log_pq_error(
66 "Error: number of chunks more than dimension.",
67 ));
68 }
69
70 let pivots_exists = context
71 .pq_storage
72 .pivot_data_exist(context.storage_provider);
73
74 let pool = context.pool;
75
76 if !pivots_exists {
77 if stage == CompressionStage::Resume {
78 return Err(ANNError::log_pq_error(
80 "Error: Pivot data does not exist when start_vertex_id is not 0.",
81 ));
82 }
83
84 let timer = Timer::new();
85
86 let rng =
87 diskann_providers::utils::create_rnd_provider_from_optional_seed(context.seed);
88 let (mut train_data, train_size, train_dim) = context
89 .pq_storage
90 .get_random_train_data_slice::<T, Storage>(
91 context.p_val,
92 context.storage_provider,
93 &mut rng.create_rnd(),
94 )?;
95
96 generate_pq_pivots(
97 GeneratePivotArguments::new(
98 train_size,
99 train_dim,
100 context.num_centers,
101 context.num_chunks,
102 context.max_kmeans_reps,
103 )?,
104 context.metric == Metric::L2,
105 &mut train_data,
106 &context.pq_storage,
107 context.storage_provider,
108 rng,
109 pool,
110 )?;
111
112 info!(
113 "PQ pivot generation took {} seconds",
114 timer.elapsed().as_secs_f64()
115 );
116 }
117
118 let (_, full_dim) = context
119 .pq_storage
120 .read_existing_pivot_metadata(context.storage_provider)?;
121
122 let num_chunks = context.num_chunks;
124 let (mut full_pivot_data, centroid, chunk_offsets) =
125 context.pq_storage.load_existing_pivot_data(
126 &num_chunks,
127 &context.num_centers,
128 &full_dim,
129 context.storage_provider,
130 )?;
131
132 let mut full_pivot_data_mat = diskann_utils::views::MutMatrixView::try_from(
133 full_pivot_data.as_mut_slice(),
134 context.num_centers,
135 full_dim,
136 )
137 .bridge_err()?;
138
139 accum_row_inplace(full_pivot_data_mat.as_mut_view(), centroid.as_slice());
140
141 let table = TransposedTable::from_parts(
142 full_pivot_data_mat.as_view(),
143 diskann_quantization::views::ChunkOffsetsView::new(&chunk_offsets)
144 .bridge_err()?
145 .to_owned(),
146 )
147 .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?;
148
149 Ok(Self {
150 table,
151 num_chunks,
152 phantom_data: PhantomData,
153 phantom_storage: PhantomData,
154 })
155 }
156
157 fn compress(
158 &self,
159 vector: MatrixBase<&[f32]>,
160 output: MatrixBase<&mut [u8]>,
161 ) -> Result<(), diskann::ANNError> {
162 self.table
163 .compress_into(vector, output)
164 .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))
165 }
166
167 fn compressed_bytes(&self) -> usize {
168 self.num_chunks
169 }
170}
171
172#[cfg(test)]
177mod pq_generation_tests {
178 use diskann::ANNError;
179 use diskann_providers::model::pq::generate_pq_pivots;
180 use diskann_providers::model::GeneratePivotArguments;
181 use diskann_providers::storage::{
182 PQStorage, StorageReadProvider, StorageWriteProvider, VirtualStorageProvider,
183 };
184 use diskann_providers::utils::{create_thread_pool_for_test, RayonThreadPoolRef};
185 use diskann_utils::{
186 io::{read_bin, write_bin},
187 test_data_root,
188 views::{MatrixView, MutMatrixView},
189 };
190 use diskann_vector::distance::Metric;
191 use rstest::rstest;
192 use vfs::FileSystem;
193
194 use super::{CompressionStage, PQGeneration, PQGenerationContext};
195 use crate::storage::quant::compressor::QuantCompressor;
196
197 const TEST_PQ_DATA_PATH: &str = "/sift/siftsmall_learn.bin";
198 const TEST_PQ_PIVOTS_PATH: &str = "/sift/siftsmall_learn_pq_pivots.bin";
199 const TEST_PQ_COMPRESSED_PATH: &str = "/sift/siftsmall_learn_pq_compressed.bin";
200 const VALIDATION_DATA: [f32; 40] = [
201 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 2.0f32, 2.0f32, 2.0f32,
203 2.0f32, 2.0f32, 2.0f32, 2.0f32, 2.0f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32,
204 2.1f32, 2.1f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 100.0f32,
205 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32,
206 ];
207 #[allow(clippy::too_many_arguments)]
208 fn create_new_compressor<'a, F: vfs::FileSystem>(
209 stage: CompressionStage,
210 provider: &'a VirtualStorageProvider<F>,
211 dim: usize,
212 num_chunks: usize,
213 max_kmeans_reps: usize,
214 num_centers: usize,
215 p_val: f64,
216 pool: RayonThreadPoolRef<'a>,
217 pivots_path: String,
218 compressed_path: String,
219 data_path: Option<&str>,
220 ) -> Result<PQGeneration<'a, f32, VirtualStorageProvider<F>>, ANNError> {
221 let pq_storage = PQStorage::new(&pivots_path, &compressed_path, data_path);
222 let context = PQGenerationContext::<'_, _> {
223 pq_storage,
224 num_chunks,
225 num_centers,
226 seed: Some(42),
227 p_val,
228 max_kmeans_reps,
229 storage_provider: provider,
230 pool,
231 metric: Metric::L2,
232 dim,
233 };
234 PQGeneration::<_, _>::new_at_stage(stage, &context)
235 }
236
237 #[rstest]
238 fn test_create_and_load_pivots_file() {
239 let storage_provider = VirtualStorageProvider::new_memory();
240 storage_provider
241 .filesystem()
242 .create_dir("/pq_generation_tests")
243 .expect("Could not create test directory");
244
245 let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin";
246 let pivot_file_name_compressor = "/pq_generation_tests/compressor_pivots_test.bin";
247 let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin";
248 let data_path = "/pq_generation_tests/data_path.bin";
249 let pq_storage: PQStorage =
250 PQStorage::new(pivot_file_name, compressed_file_name, Some(data_path));
251
252 let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5);
253 let mut train_data: Vec<f32> = VALIDATION_DATA.to_vec();
254
255 write_bin(
256 MatrixView::try_from(train_data.as_slice(), ndata, dim).unwrap(),
257 &mut storage_provider.create_for_write(data_path).unwrap(),
258 )
259 .unwrap();
260
261 let pool = create_thread_pool_for_test();
262 generate_pq_pivots(
263 GeneratePivotArguments::new(ndata, dim, num_centers, num_chunks, max_k_means_reps)
264 .unwrap(),
265 true,
266 &mut train_data,
267 &pq_storage,
268 &storage_provider,
269 diskann_providers::utils::create_rnd_provider_from_seed_in_tests(42),
270 pool.as_ref(),
271 )
272 .unwrap();
273
274 let compressor = create_new_compressor(
275 CompressionStage::Start,
276 &storage_provider,
277 dim,
278 num_chunks,
279 max_k_means_reps,
280 num_centers,
281 1.0, pool.as_ref(),
283 pivot_file_name_compressor.to_string(),
284 compressed_file_name.to_string(),
285 Some(data_path),
286 );
287
288 assert!(compressor.is_ok());
289
290 let compressor = compressor.unwrap();
291 assert_eq!(compressor.num_chunks, num_chunks);
292 assert_eq!(compressor.compressed_bytes(), num_chunks);
293
294 assert_eq!(compressor.table.dim(), dim);
295 assert_eq!(compressor.table.ncenters(), num_centers);
296 assert_eq!(compressor.table.nchunks(), num_chunks);
297
298 assert!(&storage_provider.exists(pivot_file_name_compressor));
299 let compressor_pivots = read_bin::<u8>(
300 &mut storage_provider
301 .open_reader(pivot_file_name_compressor)
302 .unwrap(),
303 )
304 .unwrap();
305 let true_pivots =
306 read_bin::<u8>(&mut storage_provider.open_reader(pivot_file_name).unwrap()).unwrap();
307 assert_eq!(compressor_pivots, true_pivots);
308 }
309
310 #[rstest]
311 fn throw_error_for_resume_and_no_existing_file() {
312 let storage_provider = VirtualStorageProvider::new_memory();
313 storage_provider
314 .filesystem()
315 .create_dir("/pq_generation_tests")
316 .expect("Could not create test directory");
317
318 let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin";
319 let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin";
320 let data_path = "/pq_generation_tests/data_path.bin";
321
322 let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5);
323
324 write_bin(
325 MatrixView::try_from(VALIDATION_DATA.as_slice(), ndata, dim).unwrap(),
326 &mut storage_provider.create_for_write(data_path).unwrap(),
327 )
328 .unwrap();
329 let pool = create_thread_pool_for_test();
330
331 let compressor = create_new_compressor(
332 CompressionStage::Resume,
333 &storage_provider,
334 dim,
335 num_chunks,
336 max_k_means_reps,
337 num_centers,
338 1.0,
339 pool.as_ref(),
340 pivot_file_name.to_string(),
341 compressed_file_name.to_string(),
342 Some(data_path),
343 );
344
345 assert!(compressor.is_err());
346 }
347
348 #[rstest]
349 fn test_pq_end_to_end_with_codebook() {
350 let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
351
352 let pool = create_thread_pool_for_test();
353 let dim = 128;
354 let num_chunks = 1;
355 let max_k_means_reps = 10;
356
357 let compressor = create_new_compressor(
358 CompressionStage::Resume,
359 &storage_provider,
360 dim,
361 num_chunks,
362 max_k_means_reps,
363 256,
364 1.0,
365 pool.as_ref(),
366 TEST_PQ_PIVOTS_PATH.to_string(),
367 "".to_string(),
368 None,
369 );
370
371 if let Err(x) = compressor.as_ref() {
372 println!("Error creating compressor: {x}");
373 };
374
375 assert!(compressor.is_ok());
376
377 let data_matrix =
378 read_bin::<f32>(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap();
379 let npts = data_matrix.nrows();
380 let mut compressed_mat = vec![0_u8; num_chunks * npts];
381 let result = compressor.unwrap().compress(
382 data_matrix.as_view(),
383 MutMatrixView::try_from(&mut compressed_mat, npts, num_chunks).unwrap(),
384 );
385 assert!(result.is_ok());
386
387 let compressed_gt = read_bin::<u8>(
388 &mut storage_provider
389 .open_reader(TEST_PQ_COMPRESSED_PATH)
390 .unwrap(),
391 )
392 .unwrap();
393 assert_eq!(compressed_gt.as_slice(), &compressed_mat);
394 }
395
396 #[rstest]
397 #[case(129, 128, 256)] #[case(128, 0, 256)] #[case(128, 128, 0)] fn test_parameter_error_cases(
401 #[case] dim: usize,
402 #[case] num_chunks: usize,
403 #[case] centers: usize,
404 ) {
405 let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
407 let pool = create_thread_pool_for_test();
408 let max_k_means_reps = 10;
409 let compressor = create_new_compressor(
410 CompressionStage::Start,
411 &storage_provider,
412 dim,
413 num_chunks,
414 max_k_means_reps,
415 centers,
416 1.0,
417 pool.as_ref(),
418 TEST_PQ_PIVOTS_PATH.to_string(),
419 "".to_string(),
420 None,
421 );
422 assert!(compressor.is_err());
423 }
424}