Skip to main content

diskann_disk/storage/quant/pq/
pq_generation.rs

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