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, time::Instant};
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},
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 = Instant::now();
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        //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(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, //take all the data to compute codebook
282            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)] // num_chunks > dim
398    #[case(128, 0, 256)] // num_chunks == 0
399    #[case(128, 128, 0)] // num_centers == 0
400    fn test_parameter_error_cases(
401        #[case] dim: usize,
402        #[case] num_chunks: usize,
403        #[case] centers: usize,
404    ) {
405        //test the error cases for parameters: num_chunks > dim, num_chunks == 0, num_centers == 0
406        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}