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 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 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(
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, 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)] #[case(128, 0, 256)] #[case(128, 128, 0)] fn test_parameter_error_cases(
407 #[case] dim: usize,
408 #[case] num_chunks: usize,
409 #[case] centers: usize,
410 ) {
411 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}