1use 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 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 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 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#[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 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, &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)] #[case(128, 0, 256)] #[case(128, 128, 0)] fn test_parameter_error_cases(
414 #[case] dim: usize,
415 #[case] num_chunks: usize,
416 #[case] centers: usize,
417 ) {
418 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}