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    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        // validate that the number of chunks is correct.
64        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                //checks for error case when stage is Resume and pivot data doesn't exist.
79                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                    context.metric == Metric::L2,
104                )?,
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        //Load the pivots
123        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//////////////////
173///// Tests /////
174/////////////////
175
176#[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        //sample validation data: npoints=5, dim=8, 5 vectors [1.0;8] [2.0;8] [2.1;8] [2.2;8] [100.0;8]
202        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(
264                ndata,
265                dim,
266                num_centers,
267                num_chunks,
268                max_k_means_reps,
269                true,
270            )
271            .unwrap(),
272            &mut train_data,
273            &pq_storage,
274            &storage_provider,
275            diskann_providers::utils::create_rnd_provider_from_seed_in_tests(42),
276            pool.as_ref(),
277        )
278        .unwrap();
279
280        let compressor = create_new_compressor(
281            CompressionStage::Start,
282            &storage_provider,
283            dim,
284            num_chunks,
285            max_k_means_reps,
286            num_centers,
287            1.0, //take all the data to compute codebook
288            pool.as_ref(),
289            pivot_file_name_compressor.to_string(),
290            compressed_file_name.to_string(),
291            Some(data_path),
292        );
293
294        assert!(compressor.is_ok());
295
296        let compressor = compressor.unwrap();
297        assert_eq!(compressor.num_chunks, num_chunks);
298        assert_eq!(compressor.compressed_bytes(), num_chunks);
299
300        assert_eq!(compressor.table.dim(), dim);
301        assert_eq!(compressor.table.ncenters(), num_centers);
302        assert_eq!(compressor.table.nchunks(), num_chunks);
303
304        assert!(&storage_provider.exists(pivot_file_name_compressor));
305        let compressor_pivots = read_bin::<u8>(
306            &mut storage_provider
307                .open_reader(pivot_file_name_compressor)
308                .unwrap(),
309        )
310        .unwrap();
311        let true_pivots =
312            read_bin::<u8>(&mut storage_provider.open_reader(pivot_file_name).unwrap()).unwrap();
313        assert_eq!(compressor_pivots, true_pivots);
314    }
315
316    #[rstest]
317    fn throw_error_for_resume_and_no_existing_file() {
318        let storage_provider = VirtualStorageProvider::new_memory();
319        storage_provider
320            .filesystem()
321            .create_dir("/pq_generation_tests")
322            .expect("Could not create test directory");
323
324        let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin";
325        let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin";
326        let data_path = "/pq_generation_tests/data_path.bin";
327
328        let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5);
329
330        write_bin(
331            MatrixView::try_from(VALIDATION_DATA.as_slice(), ndata, dim).unwrap(),
332            &mut storage_provider.create_for_write(data_path).unwrap(),
333        )
334        .unwrap();
335        let pool = create_thread_pool_for_test();
336
337        let compressor = create_new_compressor(
338            CompressionStage::Resume,
339            &storage_provider,
340            dim,
341            num_chunks,
342            max_k_means_reps,
343            num_centers,
344            1.0,
345            pool.as_ref(),
346            pivot_file_name.to_string(),
347            compressed_file_name.to_string(),
348            Some(data_path),
349        );
350
351        assert!(compressor.is_err());
352    }
353
354    #[rstest]
355    fn test_pq_end_to_end_with_codebook() {
356        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
357
358        let pool = create_thread_pool_for_test();
359        let dim = 128;
360        let num_chunks = 1;
361        let max_k_means_reps = 10;
362
363        let compressor = create_new_compressor(
364            CompressionStage::Resume,
365            &storage_provider,
366            dim,
367            num_chunks,
368            max_k_means_reps,
369            256,
370            1.0,
371            pool.as_ref(),
372            TEST_PQ_PIVOTS_PATH.to_string(),
373            "".to_string(),
374            None,
375        );
376
377        if let Err(x) = compressor.as_ref() {
378            println!("Error creating compressor: {x}");
379        };
380
381        assert!(compressor.is_ok());
382
383        let data_matrix =
384            read_bin::<f32>(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap();
385        let npts = data_matrix.nrows();
386        let mut compressed_mat = vec![0_u8; num_chunks * npts];
387        let result = compressor.unwrap().compress(
388            data_matrix.as_view(),
389            MutMatrixView::try_from(&mut compressed_mat, npts, num_chunks).unwrap(),
390        );
391        assert!(result.is_ok());
392
393        let compressed_gt = read_bin::<u8>(
394            &mut storage_provider
395                .open_reader(TEST_PQ_COMPRESSED_PATH)
396                .unwrap(),
397        )
398        .unwrap();
399        assert_eq!(compressed_gt.as_slice(), &compressed_mat);
400    }
401
402    #[rstest]
403    #[case(129, 128, 256)] // num_chunks > dim
404    #[case(128, 0, 256)] // num_chunks == 0
405    #[case(128, 128, 0)] // num_centers == 0
406    fn test_parameter_error_cases(
407        #[case] dim: usize,
408        #[case] num_chunks: usize,
409        #[case] centers: usize,
410    ) {
411        //test the error cases for parameters: num_chunks > dim, num_chunks == 0, num_centers == 0
412        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
413        let pool = create_thread_pool_for_test();
414        let max_k_means_reps = 10;
415        let compressor = create_new_compressor(
416            CompressionStage::Start,
417            &storage_provider,
418            dim,
419            num_chunks,
420            max_k_means_reps,
421            centers,
422            1.0,
423            pool.as_ref(),
424            TEST_PQ_PIVOTS_PATH.to_string(),
425            "".to_string(),
426            None,
427        );
428        assert!(compressor.is_err());
429    }
430}