From 2d9f1bb25afc87a4307bf791e8b163a7d20cd5b3 Mon Sep 17 00:00:00 2001 From: xiajiang Date: Thu, 30 Jul 2026 15:43:41 +0800 Subject: [PATCH 1/4] Make disk-search PQ pivot generation explicit --- diskann-disk/src/build/builder/build.rs | 2 +- diskann-disk/src/storage/quant/compressor.rs | 4 +- diskann-disk/src/storage/quant/generator.rs | 93 +++++++++----- .../src/storage/quant/pq/pq_generation.rs | 113 +++++++++++------- 4 files changed, 134 insertions(+), 78 deletions(-) diff --git a/diskann-disk/src/build/builder/build.rs b/diskann-disk/src/build/builder/build.rs index c0df5d6b75..5db76bb1a0 100644 --- a/diskann-disk/src/build/builder/build.rs +++ b/diskann-disk/src/build/builder/build.rs @@ -164,7 +164,7 @@ where self.index_writer.get_dataset_file(), self.pq_storage.get_compressed_data_path().into(), &quantizer_context, - )?; + ); generator.generate_data( storage_provider, pool, diff --git a/diskann-disk/src/storage/quant/compressor.rs b/diskann-disk/src/storage/quant/compressor.rs index acddd89940..2de3cdfbfb 100644 --- a/diskann-disk/src/storage/quant/compressor.rs +++ b/diskann-disk/src/storage/quant/compressor.rs @@ -20,7 +20,7 @@ use diskann_utils::views::{MatrixView, MutMatrixView}; /// - [`CompressorContext`]: An overloadable type that provides initialization parameters for the compressor /// /// # Methods -/// - `new`: Constructs a new compressor instance with the provided context. +/// - `prepare`: Performs any setup needed before compression and returns a compressor. /// - `compress`: Compresses a batch of vectors into the output buffer. /// - `compressed_bytes`: Returns the size in bytes of each compressed vector pub trait QuantCompressor: Sized + Sync @@ -29,7 +29,7 @@ where { type CompressorContext; - fn new(context: &Self::CompressorContext) -> ANNResult; + fn prepare(context: &Self::CompressorContext) -> ANNResult; fn compress(&self, vector: MatrixView, output: MutMatrixView) -> ANNResult<()>; fn compressed_bytes(&self) -> usize; } diff --git a/diskann-disk/src/storage/quant/generator.rs b/diskann-disk/src/storage/quant/generator.rs index 2593f9bf31..0ebae5a420 100644 --- a/diskann-disk/src/storage/quant/generator.rs +++ b/diskann-disk/src/storage/quant/generator.rs @@ -25,18 +25,18 @@ use crate::{ /// [`QuantDataGenerator`] orchestrates the process of reading vector data, applying quantization, /// and writing compressed results to storage in batches. -pub struct QuantDataGenerator +pub struct QuantDataGenerator<'a, T, Q> where T: Copy + VectorRepr, Q: QuantCompressor, { - pub quantizer: Q, + quantizer_context: &'a Q::CompressorContext, pub data_path: String, pub compressed_data_path: String, phantom: PhantomData, } -impl QuantDataGenerator +impl<'a, T, Q> QuantDataGenerator<'a, T, Q> where T: Copy + VectorRepr, Q: QuantCompressor, @@ -44,15 +44,14 @@ where pub fn new( data_path: String, compressed_data_path: String, - quantizer_context: &Q::CompressorContext, - ) -> ANNResult { - let quantizer = Q::new(quantizer_context)?; - Ok(Self { + quantizer_context: &'a Q::CompressorContext, + ) -> Self { + Self { data_path, compressed_data_path, - quantizer, + quantizer_context, phantom: PhantomData, - }) + } } /// This method reads the source data file, processes vectors in batches, compresses them @@ -93,6 +92,8 @@ where )); } + let quantizer = Q::prepare(self.quantizer_context)?; + let compressed_size = quantizer.compressed_bytes(); let compressed_path = self.compressed_data_path.as_str(); if storage_provider.exists(compressed_path) { @@ -105,10 +106,8 @@ where data_reader.seek(SeekFrom::Start((std::mem::size_of::() * 2) as u64))?; let mut compressed_data_writer = storage_provider.create_for_write(compressed_path)?; - Metadata::new(num_points, self.quantizer.compressed_bytes())? - .write(&mut compressed_data_writer)?; + Metadata::new(num_points, compressed_size)?.write(&mut compressed_data_writer)?; - let compressed_size = self.quantizer.compressed_bytes(); let block_size = std::cmp::min(num_points, max_block_size); let num_blocks = num_points / block_size + !num_points.is_multiple_of(block_size) as usize; @@ -163,7 +162,7 @@ where base_block .par_window_iter(BATCH_SIZE) .zip_eq(compressed_block.par_window_iter_mut(BATCH_SIZE)) - .try_for_each_in_pool(pool, |(src, dst)| self.quantizer.compress(src, dst))?; + .try_for_each_in_pool(pool, |(src, dst)| quantizer.compress(src, dst))?; let write_offset = start_index * compressed_size + std::mem::size_of::() * 2; compressed_data_writer.seek(SeekFrom::Start(write_offset as u64))?; @@ -191,7 +190,10 @@ where #[cfg(test)] mod generator_tests { - use std::io::BufReader; + use std::{ + io::BufReader, + sync::atomic::{AtomicUsize, Ordering}, + }; use diskann::utils::read_exact_into; use diskann_providers::storage::VirtualStorageProvider; @@ -204,6 +206,24 @@ mod generator_tests { use vfs::{FileSystem, MemoryFS}; use super::*; + pub struct DummyCompressorContext { + pub output_dim: u32, + pub prepare_calls: AtomicUsize, + } + + impl DummyCompressorContext { + pub fn new(output_dim: u32) -> Self { + Self { + output_dim, + prepare_calls: AtomicUsize::new(0), + } + } + + pub fn prepare_calls(&self) -> usize { + self.prepare_calls.load(Ordering::SeqCst) + } + } + pub struct DummyCompressor { pub output_dim: u32, pub code: Vec, @@ -217,10 +237,11 @@ mod generator_tests { } } impl QuantCompressor for DummyCompressor { - type CompressorContext = u32; + type CompressorContext = DummyCompressorContext; - fn new(context: &Self::CompressorContext) -> ANNResult { - Ok(Self::new(*context)) + fn prepare(context: &Self::CompressorContext) -> ANNResult { + context.prepare_calls.fetch_add(1, Ordering::SeqCst); + Ok(Self::new(context.output_dim)) } fn compress( @@ -276,20 +297,16 @@ mod generator_tests { Ok((storage_provider, data_path, compressed_path)) } - fn create_and_call_generator( + fn create_and_call_generator<'a, F: vfs::FileSystem>( compressed_path: String, storage_provider: &VirtualStorageProvider, data_path: String, - output_dim: u32, + context: &'a DummyCompressorContext, max_block_size: usize, - ) -> (QuantDataGenerator, ANNResult<()>) { + ) -> (QuantDataGenerator<'a, f32, DummyCompressor>, ANNResult<()>) { let pool: diskann_providers::utils::RayonThreadPool = create_thread_pool_for_test(); - let generator = QuantDataGenerator::::new( - data_path, - compressed_path, - &output_dim, - ) - .unwrap(); + let generator = + QuantDataGenerator::::new(data_path, compressed_path, context); let result = generator.generate_data(storage_provider, pool.as_ref(), max_block_size); (generator, result) } @@ -304,15 +321,17 @@ mod generator_tests { #[case] output_dim: u32, ) -> ANNResult<()> { let (storage_provider, data_path, compressed_path) = generate_data_files(num_points, dim)?; - let (generator, result) = create_and_call_generator( + let context = DummyCompressorContext::new(output_dim); + let (_generator, result) = create_and_call_generator( compressed_path.clone(), &storage_provider, data_path, - output_dim, + &context, 10_000, ); result?; + assert_eq!(context.prepare_calls(), 1); assert!(storage_provider.exists(&compressed_path)); let expected_size = num_points * output_dim as usize; @@ -328,8 +347,9 @@ mod generator_tests { assert_eq!(metadata.ndims_u32(), output_dim); assert_eq!(metadata.npoints(), num_points); + let expected_code: Vec = (0..output_dim).map(|x| (x % 256) as u8).collect(); data.chunks_exact(output_dim as usize) - .for_each(|chunk| assert_eq!(chunk, generator.quantizer.code.as_slice())); + .for_each(|chunk| assert_eq!(chunk, expected_code.as_slice())); Ok(()) } @@ -345,16 +365,18 @@ mod generator_tests { let data_path = "/test_data/empty.bin".to_string(); let compressed_path = "/test_data/empty_compressed.bin".to_string(); Metadata::new(0, 8)?.write(&mut storage_provider.create_for_write(data_path.as_str())?)?; + let context = DummyCompressorContext::new(4); let (_, result) = create_and_call_generator( compressed_path.clone(), &storage_provider, data_path, - 4, + &context, 10_000, ); assert!(result.is_err()); + assert_eq!(context.prepare_calls(), 0); assert!(!storage_provider.exists(&compressed_path)); Ok(()) } @@ -362,11 +384,18 @@ mod generator_tests { #[test] fn generate_data_rejects_zero_chunk_size() -> ANNResult<()> { let (storage_provider, data_path, compressed_path) = generate_data_files(1, 8)?; + let context = DummyCompressorContext::new(4); - let (_, result) = - create_and_call_generator(compressed_path.clone(), &storage_provider, data_path, 4, 0); + let (_, result) = create_and_call_generator( + compressed_path.clone(), + &storage_provider, + data_path, + &context, + 0, + ); assert!(result.is_err()); + assert_eq!(context.prepare_calls(), 0); assert!(!storage_provider.exists(&compressed_path)); Ok(()) } diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index c5297ae9c2..0ff9f445d1 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -52,14 +52,14 @@ where phantom_storage: PhantomData<&'a Storage>, } -impl<'a, T, Storage> QuantCompressor for PQGeneration<'a, T, Storage> +impl<'a, T, Storage> PQGeneration<'a, T, Storage> where T: VectorRepr, Storage: StorageReadProvider + StorageWriteProvider + 'a, { - type CompressorContext = PQGenerationContext<'a, Storage>; - - fn new(context: &Self::CompressorContext) -> diskann::ANNResult { + pub(crate) fn generate_pivots( + context: &PQGenerationContext<'a, Storage>, + ) -> diskann::ANNResult<()> { // validate that the number of chunks is correct. if context.num_chunks > context.dim { return Err(diskann_error!( @@ -109,6 +109,20 @@ where ); } + Ok(()) + } +} + +impl<'a, T, Storage> QuantCompressor for PQGeneration<'a, T, Storage> +where + T: VectorRepr, + Storage: StorageReadProvider + StorageWriteProvider + 'a, +{ + type CompressorContext = PQGenerationContext<'a, Storage>; + + fn prepare(context: &Self::CompressorContext) -> diskann::ANNResult { + Self::generate_pivots(context)?; + let (_, full_dim) = context .pq_storage .read_existing_pivot_metadata(context.storage_provider)?; @@ -169,9 +183,6 @@ where #[cfg(test)] mod pq_generation_tests { - use diskann::ANNError; - use diskann_providers::model::pq::generate_pq_pivots; - use diskann_providers::model::GeneratePivotArguments; use diskann_providers::storage::{ PQStorage, StorageReadProvider, StorageWriteProvider, VirtualStorageProvider, }; @@ -199,7 +210,7 @@ mod pq_generation_tests { 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, ]; #[allow(clippy::too_many_arguments)] - fn create_new_compressor<'a, F: vfs::FileSystem>( + fn create_context<'a, F: vfs::FileSystem>( provider: &'a VirtualStorageProvider, dim: usize, num_chunks: usize, @@ -210,9 +221,9 @@ mod pq_generation_tests { pivots_path: String, compressed_path: String, data_path: Option<&str>, - ) -> Result>, ANNError> { + ) -> PQGenerationContext<'a, VirtualStorageProvider> { let pq_storage = PQStorage::new(&pivots_path, &compressed_path, data_path); - let context = PQGenerationContext::<'_, _> { + PQGenerationContext::<'_, _> { pq_storage, num_chunks, num_centers, @@ -223,48 +234,31 @@ mod pq_generation_tests { pool, metric: Metric::L2, dim, - }; - PQGeneration::<_, _>::new(&context) + } } #[rstest] - fn test_create_and_load_pivots_file() { + fn explicit_generation_creates_pivots_file() { let storage_provider = VirtualStorageProvider::new_memory(); storage_provider .filesystem() .create_dir("/pq_generation_tests") .expect("Could not create test directory"); - let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin"; - let pivot_file_name_compressor = "/pq_generation_tests/compressor_pivots_test.bin"; + let pivot_file_name = "/pq_generation_tests/pivots_test.bin"; let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin"; let data_path = "/pq_generation_tests/data_path.bin"; - let pq_storage: PQStorage = - PQStorage::new(pivot_file_name, compressed_file_name, Some(data_path)); let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5); - let mut train_data: Vec = VALIDATION_DATA.to_vec(); write_bin( - MatrixView::try_from(train_data.as_slice(), ndata, dim).unwrap(), + MatrixView::try_from(VALIDATION_DATA.as_slice(), ndata, dim).unwrap(), &mut storage_provider.create_for_write(data_path).unwrap(), ) .unwrap(); let pool = create_thread_pool_for_test(); - generate_pq_pivots( - GeneratePivotArguments::new(ndata, dim, num_centers, num_chunks, max_k_means_reps) - .unwrap(), - true, - &mut train_data, - &pq_storage, - &storage_provider, - diskann_providers::utils::create_rnd_provider_from_seed_in_tests(42), - pool.as_ref(), - ) - .unwrap(); - - let compressor = create_new_compressor( + let context = create_context( &storage_provider, dim, num_chunks, @@ -272,12 +266,16 @@ mod pq_generation_tests { num_centers, 1.0, //take all the data to compute codebook pool.as_ref(), - pivot_file_name_compressor.to_string(), + pivot_file_name.to_string(), compressed_file_name.to_string(), Some(data_path), ); + assert!(!storage_provider.exists(pivot_file_name)); + + let compressor = PQGeneration::::prepare(&context); assert!(compressor.is_ok()); + assert!(storage_provider.exists(pivot_file_name)); let compressor = compressor.unwrap(); assert_eq!(compressor.num_chunks, num_chunks); @@ -286,17 +284,44 @@ mod pq_generation_tests { assert_eq!(compressor.table.dim(), dim); assert_eq!(compressor.table.ncenters(), num_centers); assert_eq!(compressor.table.nchunks(), num_chunks); + } - assert!(&storage_provider.exists(pivot_file_name_compressor)); - let compressor_pivots = read_bin::( - &mut storage_provider - .open_reader(pivot_file_name_compressor) - .unwrap(), + #[rstest] + fn prepare_generates_missing_pivots() { + let storage_provider = VirtualStorageProvider::new_memory(); + storage_provider + .filesystem() + .create_dir("/pq_generation_tests") + .expect("Could not create test directory"); + + let pivot_file_name = "/pq_generation_tests/missing_pivots.bin"; + let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin"; + let data_path = "/pq_generation_tests/data_path.bin"; + + write_bin( + MatrixView::try_from(VALIDATION_DATA.as_slice(), 5, 8).unwrap(), + &mut storage_provider.create_for_write(data_path).unwrap(), ) .unwrap(); - let true_pivots = - read_bin::(&mut storage_provider.open_reader(pivot_file_name).unwrap()).unwrap(); - assert_eq!(compressor_pivots, true_pivots); + + let pool = create_thread_pool_for_test(); + let context = create_context( + &storage_provider, + 8, + 2, + 5, + 2, + 1.0, + pool.as_ref(), + pivot_file_name.to_string(), + compressed_file_name.to_string(), + Some(data_path), + ); + + let compressor = PQGeneration::::prepare(&context); + + assert!(compressor.is_ok()); + assert!(storage_provider.exists(pivot_file_name)); } #[rstest] @@ -308,7 +333,7 @@ mod pq_generation_tests { let num_chunks = 1; let max_k_means_reps = 10; - let compressor = create_new_compressor( + let context = create_context( &storage_provider, dim, num_chunks, @@ -320,6 +345,7 @@ mod pq_generation_tests { "".to_string(), None, ); + let compressor = PQGeneration::::prepare(&context); if let Err(x) = compressor.as_ref() { println!("Error creating compressor: {x}"); @@ -359,7 +385,7 @@ mod pq_generation_tests { let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); let pool = create_thread_pool_for_test(); let max_k_means_reps = 10; - let compressor = create_new_compressor( + let context = create_context( &storage_provider, dim, num_chunks, @@ -371,6 +397,7 @@ mod pq_generation_tests { "".to_string(), None, ); + let compressor = PQGeneration::::prepare(&context); assert!(compressor.is_err()); } } From 1ab5f08f34ddc3f7bb50e115cb6a84cd2ce741c1 Mon Sep 17 00:00:00 2001 From: juchen-ms Date: Wed, 12 Aug 2026 15:59:23 +0800 Subject: [PATCH 2/4] Refine quant compressor generation lifecycle (#1314) Follow-up suggestion for #1299. This keeps `new` construction-only and moves the generation work to an explicit instance method: - move the compressor context into the quantizer at construction time - call `generate(&self)` before compression - preserve existing pivot generation/reuse behavior - keep `QuantDataGenerator` responsible for orchestrating generation and compression Validation: - `cargo test -p diskann-disk` - `cargo clippy -p diskann-disk --all-targets -- -D warnings` Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- diskann-disk/src/build/builder/build.rs | 2 +- diskann-disk/src/storage/quant/compressor.rs | 6 +- diskann-disk/src/storage/quant/generator.rs | 87 +++++++------------ .../src/storage/quant/pq/pq_generation.rs | 76 ++++++++++------ 4 files changed, 85 insertions(+), 86 deletions(-) diff --git a/diskann-disk/src/build/builder/build.rs b/diskann-disk/src/build/builder/build.rs index 5db76bb1a0..ec9ade194b 100644 --- a/diskann-disk/src/build/builder/build.rs +++ b/diskann-disk/src/build/builder/build.rs @@ -163,7 +163,7 @@ where >::new( self.index_writer.get_dataset_file(), self.pq_storage.get_compressed_data_path().into(), - &quantizer_context, + quantizer_context, ); generator.generate_data( storage_provider, diff --git a/diskann-disk/src/storage/quant/compressor.rs b/diskann-disk/src/storage/quant/compressor.rs index 2de3cdfbfb..7b55634dff 100644 --- a/diskann-disk/src/storage/quant/compressor.rs +++ b/diskann-disk/src/storage/quant/compressor.rs @@ -20,7 +20,8 @@ use diskann_utils::views::{MatrixView, MutMatrixView}; /// - [`CompressorContext`]: An overloadable type that provides initialization parameters for the compressor /// /// # Methods -/// - `prepare`: Performs any setup needed before compression and returns a compressor. +/// - `new`: Constructs a compressor with the provided context. +/// - `generate`: Generates any data needed before compression. /// - `compress`: Compresses a batch of vectors into the output buffer. /// - `compressed_bytes`: Returns the size in bytes of each compressed vector pub trait QuantCompressor: Sized + Sync @@ -29,7 +30,8 @@ where { type CompressorContext; - fn prepare(context: &Self::CompressorContext) -> ANNResult; + fn new(context: Self::CompressorContext) -> Self; + fn generate(&self) -> ANNResult<()>; fn compress(&self, vector: MatrixView, output: MutMatrixView) -> ANNResult<()>; fn compressed_bytes(&self) -> usize; } diff --git a/diskann-disk/src/storage/quant/generator.rs b/diskann-disk/src/storage/quant/generator.rs index 0ebae5a420..79913fd241 100644 --- a/diskann-disk/src/storage/quant/generator.rs +++ b/diskann-disk/src/storage/quant/generator.rs @@ -25,18 +25,18 @@ use crate::{ /// [`QuantDataGenerator`] orchestrates the process of reading vector data, applying quantization, /// and writing compressed results to storage in batches. -pub struct QuantDataGenerator<'a, T, Q> +pub struct QuantDataGenerator where T: Copy + VectorRepr, Q: QuantCompressor, { - quantizer_context: &'a Q::CompressorContext, + pub quantizer: Q, pub data_path: String, pub compressed_data_path: String, phantom: PhantomData, } -impl<'a, T, Q> QuantDataGenerator<'a, T, Q> +impl QuantDataGenerator where T: Copy + VectorRepr, Q: QuantCompressor, @@ -44,12 +44,13 @@ where pub fn new( data_path: String, compressed_data_path: String, - quantizer_context: &'a Q::CompressorContext, + quantizer_context: Q::CompressorContext, ) -> Self { + let quantizer = Q::new(quantizer_context); Self { data_path, compressed_data_path, - quantizer_context, + quantizer, phantom: PhantomData, } } @@ -92,8 +93,7 @@ where )); } - let quantizer = Q::prepare(self.quantizer_context)?; - let compressed_size = quantizer.compressed_bytes(); + self.quantizer.generate()?; let compressed_path = self.compressed_data_path.as_str(); if storage_provider.exists(compressed_path) { @@ -106,8 +106,10 @@ where data_reader.seek(SeekFrom::Start((std::mem::size_of::() * 2) as u64))?; let mut compressed_data_writer = storage_provider.create_for_write(compressed_path)?; - Metadata::new(num_points, compressed_size)?.write(&mut compressed_data_writer)?; + Metadata::new(num_points, self.quantizer.compressed_bytes())? + .write(&mut compressed_data_writer)?; + let compressed_size = self.quantizer.compressed_bytes(); let block_size = std::cmp::min(num_points, max_block_size); let num_blocks = num_points / block_size + !num_points.is_multiple_of(block_size) as usize; @@ -162,7 +164,7 @@ where base_block .par_window_iter(BATCH_SIZE) .zip_eq(compressed_block.par_window_iter_mut(BATCH_SIZE)) - .try_for_each_in_pool(pool, |(src, dst)| quantizer.compress(src, dst))?; + .try_for_each_in_pool(pool, |(src, dst)| self.quantizer.compress(src, dst))?; let write_offset = start_index * compressed_size + std::mem::size_of::() * 2; compressed_data_writer.seek(SeekFrom::Start(write_offset as u64))?; @@ -190,10 +192,7 @@ where #[cfg(test)] mod generator_tests { - use std::{ - io::BufReader, - sync::atomic::{AtomicUsize, Ordering}, - }; + use std::io::BufReader; use diskann::utils::read_exact_into; use diskann_providers::storage::VirtualStorageProvider; @@ -206,24 +205,6 @@ mod generator_tests { use vfs::{FileSystem, MemoryFS}; use super::*; - pub struct DummyCompressorContext { - pub output_dim: u32, - pub prepare_calls: AtomicUsize, - } - - impl DummyCompressorContext { - pub fn new(output_dim: u32) -> Self { - Self { - output_dim, - prepare_calls: AtomicUsize::new(0), - } - } - - pub fn prepare_calls(&self) -> usize { - self.prepare_calls.load(Ordering::SeqCst) - } - } - pub struct DummyCompressor { pub output_dim: u32, pub code: Vec, @@ -237,11 +218,14 @@ mod generator_tests { } } impl QuantCompressor for DummyCompressor { - type CompressorContext = DummyCompressorContext; + type CompressorContext = u32; + + fn new(context: Self::CompressorContext) -> Self { + Self::new(context) + } - fn prepare(context: &Self::CompressorContext) -> ANNResult { - context.prepare_calls.fetch_add(1, Ordering::SeqCst); - Ok(Self::new(context.output_dim)) + fn generate(&self) -> ANNResult<()> { + Ok(()) } fn compress( @@ -297,16 +281,16 @@ mod generator_tests { Ok((storage_provider, data_path, compressed_path)) } - fn create_and_call_generator<'a, F: vfs::FileSystem>( + fn create_and_call_generator( compressed_path: String, storage_provider: &VirtualStorageProvider, data_path: String, - context: &'a DummyCompressorContext, + output_dim: u32, max_block_size: usize, - ) -> (QuantDataGenerator<'a, f32, DummyCompressor>, ANNResult<()>) { + ) -> (QuantDataGenerator, ANNResult<()>) { let pool: diskann_providers::utils::RayonThreadPool = create_thread_pool_for_test(); let generator = - QuantDataGenerator::::new(data_path, compressed_path, context); + QuantDataGenerator::::new(data_path, compressed_path, output_dim); let result = generator.generate_data(storage_provider, pool.as_ref(), max_block_size); (generator, result) } @@ -321,17 +305,15 @@ mod generator_tests { #[case] output_dim: u32, ) -> ANNResult<()> { let (storage_provider, data_path, compressed_path) = generate_data_files(num_points, dim)?; - let context = DummyCompressorContext::new(output_dim); - let (_generator, result) = create_and_call_generator( + let (generator, result) = create_and_call_generator( compressed_path.clone(), &storage_provider, data_path, - &context, + output_dim, 10_000, ); result?; - assert_eq!(context.prepare_calls(), 1); assert!(storage_provider.exists(&compressed_path)); let expected_size = num_points * output_dim as usize; @@ -347,9 +329,8 @@ mod generator_tests { assert_eq!(metadata.ndims_u32(), output_dim); assert_eq!(metadata.npoints(), num_points); - let expected_code: Vec = (0..output_dim).map(|x| (x % 256) as u8).collect(); data.chunks_exact(output_dim as usize) - .for_each(|chunk| assert_eq!(chunk, expected_code.as_slice())); + .for_each(|chunk| assert_eq!(chunk, generator.quantizer.code.as_slice())); Ok(()) } @@ -365,18 +346,15 @@ mod generator_tests { let data_path = "/test_data/empty.bin".to_string(); let compressed_path = "/test_data/empty_compressed.bin".to_string(); Metadata::new(0, 8)?.write(&mut storage_provider.create_for_write(data_path.as_str())?)?; - let context = DummyCompressorContext::new(4); let (_, result) = create_and_call_generator( compressed_path.clone(), &storage_provider, data_path, - &context, + 4, 10_000, ); - assert!(result.is_err()); - assert_eq!(context.prepare_calls(), 0); assert!(!storage_provider.exists(&compressed_path)); Ok(()) } @@ -384,18 +362,11 @@ mod generator_tests { #[test] fn generate_data_rejects_zero_chunk_size() -> ANNResult<()> { let (storage_provider, data_path, compressed_path) = generate_data_files(1, 8)?; - let context = DummyCompressorContext::new(4); - let (_, result) = create_and_call_generator( - compressed_path.clone(), - &storage_provider, - data_path, - &context, - 0, - ); + let (_, result) = + create_and_call_generator(compressed_path.clone(), &storage_provider, data_path, 4, 0); assert!(result.is_err()); - assert_eq!(context.prepare_calls(), 0); assert!(!storage_provider.exists(&compressed_path)); Ok(()) } diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index 0ff9f445d1..3e64fbf43c 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use std::{marker::PhantomData, time::Instant}; +use std::{marker::PhantomData, sync::OnceLock, time::Instant}; use diskann::utils::VectorRepr; use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider}; @@ -46,10 +46,10 @@ where T: VectorRepr, Storage: StorageReadProvider + StorageWriteProvider + 'a, { - table: TransposedTable, + context: PQGenerationContext<'a, Storage>, + table: OnceLock, num_chunks: usize, phantom_data: PhantomData, - phantom_storage: PhantomData<&'a Storage>, } impl<'a, T, Storage> PQGeneration<'a, T, Storage> @@ -120,9 +120,23 @@ where { type CompressorContext = PQGenerationContext<'a, Storage>; - fn prepare(context: &Self::CompressorContext) -> diskann::ANNResult { - Self::generate_pivots(context)?; + fn new(context: Self::CompressorContext) -> Self { + let num_chunks = context.num_chunks; + Self { + context, + table: OnceLock::new(), + num_chunks, + phantom_data: PhantomData, + } + } + + fn generate(&self) -> diskann::ANNResult<()> { + if self.table.get().is_some() { + return Ok(()); + } + let context = &self.context; + Self::generate_pivots(context)?; let (_, full_dim) = context .pq_storage .read_existing_pivot_metadata(context.storage_provider)?; @@ -154,11 +168,11 @@ where ) .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err)))?; - Ok(Self { - table, - num_chunks, - phantom_data: PhantomData, - phantom_storage: PhantomData, + self.table.set(table).map_err(|_| { + diskann_error!( + ErrorKind::PQError, + "PQ compressor was generated concurrently" + ) }) } @@ -168,6 +182,13 @@ where output: MatrixBase<&mut [u8]>, ) -> Result<(), diskann::ANNError> { self.table + .get() + .ok_or_else(|| { + diskann_error!( + ErrorKind::PQError, + "PQ compressor must be generated before compression" + ) + })? .compress_into(vector, output) .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err))) } @@ -273,21 +294,24 @@ mod pq_generation_tests { assert!(!storage_provider.exists(pivot_file_name)); - let compressor = PQGeneration::::prepare(&context); - assert!(compressor.is_ok()); + let compressor = PQGeneration::::new(context); + assert!(!storage_provider.exists(pivot_file_name)); + + let result = compressor.generate(); + assert!(result.is_ok()); assert!(storage_provider.exists(pivot_file_name)); - let compressor = compressor.unwrap(); assert_eq!(compressor.num_chunks, num_chunks); assert_eq!(compressor.compressed_bytes(), num_chunks); - assert_eq!(compressor.table.dim(), dim); - assert_eq!(compressor.table.ncenters(), num_centers); - assert_eq!(compressor.table.nchunks(), num_chunks); + let table = compressor.table.get().unwrap(); + assert_eq!(table.dim(), dim); + assert_eq!(table.ncenters(), num_centers); + assert_eq!(table.nchunks(), num_chunks); } #[rstest] - fn prepare_generates_missing_pivots() { + fn generate_creates_missing_pivots() { let storage_provider = VirtualStorageProvider::new_memory(); storage_provider .filesystem() @@ -318,9 +342,10 @@ mod pq_generation_tests { Some(data_path), ); - let compressor = PQGeneration::::prepare(&context); + let compressor = PQGeneration::::new(context); + let result = compressor.generate(); - assert!(compressor.is_ok()); + assert!(result.is_ok()); assert!(storage_provider.exists(pivot_file_name)); } @@ -345,19 +370,20 @@ mod pq_generation_tests { "".to_string(), None, ); - let compressor = PQGeneration::::prepare(&context); + let compressor = PQGeneration::::new(context); + let result = compressor.generate(); - if let Err(x) = compressor.as_ref() { + if let Err(x) = result.as_ref() { println!("Error creating compressor: {x}"); }; - assert!(compressor.is_ok()); + assert!(result.is_ok()); let data_matrix = read_bin::(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap(); let npts = data_matrix.nrows(); let mut compressed_mat = vec![0_u8; num_chunks * npts]; - let result = compressor.unwrap().compress( + let result = compressor.compress( data_matrix.as_view(), MutMatrixView::try_from(&mut compressed_mat, npts, num_chunks).unwrap(), ); @@ -397,7 +423,7 @@ mod pq_generation_tests { "".to_string(), None, ); - let compressor = PQGeneration::::prepare(&context); - assert!(compressor.is_err()); + let result = PQGeneration::::new(context).generate(); + assert!(result.is_err()); } } From d226afa81923c1614f5c800297158ed6bd8051bb Mon Sep 17 00:00:00 2001 From: xiajiang Date: Wed, 12 Aug 2026 17:20:56 +0800 Subject: [PATCH 3/4] Preserve quant compressor API for explicit PQ generation --- diskann-disk/src/build/builder/build.rs | 6 +- diskann-disk/src/storage/quant/compressor.rs | 6 +- diskann-disk/src/storage/quant/generator.rs | 27 ++++--- .../src/storage/quant/pq/pq_generation.rs | 74 ++++++------------- 4 files changed, 43 insertions(+), 70 deletions(-) diff --git a/diskann-disk/src/build/builder/build.rs b/diskann-disk/src/build/builder/build.rs index ec9ade194b..b2c0b47bac 100644 --- a/diskann-disk/src/build/builder/build.rs +++ b/diskann-disk/src/build/builder/build.rs @@ -157,14 +157,16 @@ where metric: self.index_configuration.dist_metric, }; + PQGeneration::::generate_pivots(&quantizer_context)?; + let generator = QuantDataGenerator::< Data::VectorDataType, PQGeneration, >::new( self.index_writer.get_dataset_file(), self.pq_storage.get_compressed_data_path().into(), - quantizer_context, - ); + &quantizer_context, + )?; generator.generate_data( storage_provider, pool, diff --git a/diskann-disk/src/storage/quant/compressor.rs b/diskann-disk/src/storage/quant/compressor.rs index 7b55634dff..acddd89940 100644 --- a/diskann-disk/src/storage/quant/compressor.rs +++ b/diskann-disk/src/storage/quant/compressor.rs @@ -20,8 +20,7 @@ use diskann_utils::views::{MatrixView, MutMatrixView}; /// - [`CompressorContext`]: An overloadable type that provides initialization parameters for the compressor /// /// # Methods -/// - `new`: Constructs a compressor with the provided context. -/// - `generate`: Generates any data needed before compression. +/// - `new`: Constructs a new compressor instance with the provided context. /// - `compress`: Compresses a batch of vectors into the output buffer. /// - `compressed_bytes`: Returns the size in bytes of each compressed vector pub trait QuantCompressor: Sized + Sync @@ -30,8 +29,7 @@ where { type CompressorContext; - fn new(context: Self::CompressorContext) -> Self; - fn generate(&self) -> ANNResult<()>; + fn new(context: &Self::CompressorContext) -> ANNResult; fn compress(&self, vector: MatrixView, output: MutMatrixView) -> ANNResult<()>; fn compressed_bytes(&self) -> usize; } diff --git a/diskann-disk/src/storage/quant/generator.rs b/diskann-disk/src/storage/quant/generator.rs index 79913fd241..36af775bfb 100644 --- a/diskann-disk/src/storage/quant/generator.rs +++ b/diskann-disk/src/storage/quant/generator.rs @@ -44,15 +44,15 @@ where pub fn new( data_path: String, compressed_data_path: String, - quantizer_context: Q::CompressorContext, - ) -> Self { - let quantizer = Q::new(quantizer_context); - Self { + quantizer_context: &Q::CompressorContext, + ) -> ANNResult { + let quantizer = Q::new(quantizer_context)?; + Ok(Self { data_path, compressed_data_path, quantizer, phantom: PhantomData, - } + }) } /// This method reads the source data file, processes vectors in batches, compresses them @@ -93,7 +93,6 @@ where )); } - self.quantizer.generate()?; let compressed_path = self.compressed_data_path.as_str(); if storage_provider.exists(compressed_path) { @@ -220,12 +219,8 @@ mod generator_tests { impl QuantCompressor for DummyCompressor { type CompressorContext = u32; - fn new(context: Self::CompressorContext) -> Self { - Self::new(context) - } - - fn generate(&self) -> ANNResult<()> { - Ok(()) + fn new(context: &Self::CompressorContext) -> ANNResult { + Ok(Self::new(*context)) } fn compress( @@ -289,8 +284,12 @@ mod generator_tests { max_block_size: usize, ) -> (QuantDataGenerator, ANNResult<()>) { let pool: diskann_providers::utils::RayonThreadPool = create_thread_pool_for_test(); - let generator = - QuantDataGenerator::::new(data_path, compressed_path, output_dim); + let generator = QuantDataGenerator::::new( + data_path, + compressed_path, + &output_dim, + ) + .unwrap(); let result = generator.generate_data(storage_provider, pool.as_ref(), max_block_size); (generator, result) } diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index 3e64fbf43c..5aabf2254c 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use std::{marker::PhantomData, sync::OnceLock, time::Instant}; +use std::{marker::PhantomData, time::Instant}; use diskann::utils::VectorRepr; use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider}; @@ -46,10 +46,10 @@ where T: VectorRepr, Storage: StorageReadProvider + StorageWriteProvider + 'a, { - context: PQGenerationContext<'a, Storage>, - table: OnceLock, + table: TransposedTable, num_chunks: usize, phantom_data: PhantomData, + phantom_storage: PhantomData<&'a Storage>, } impl<'a, T, Storage> PQGeneration<'a, T, Storage> @@ -120,28 +120,13 @@ where { type CompressorContext = PQGenerationContext<'a, Storage>; - fn new(context: Self::CompressorContext) -> Self { - let num_chunks = context.num_chunks; - Self { - context, - table: OnceLock::new(), - num_chunks, - phantom_data: PhantomData, - } - } - - fn generate(&self) -> diskann::ANNResult<()> { - if self.table.get().is_some() { - return Ok(()); - } - - let context = &self.context; + fn new(context: &Self::CompressorContext) -> diskann::ANNResult { Self::generate_pivots(context)?; + let (_, full_dim) = context .pq_storage .read_existing_pivot_metadata(context.storage_provider)?; - //Load the pivots let num_chunks = context.num_chunks; let (mut full_pivot_data, centroid, chunk_offsets) = context.pq_storage.load_existing_pivot_data( @@ -168,11 +153,11 @@ where ) .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err)))?; - self.table.set(table).map_err(|_| { - diskann_error!( - ErrorKind::PQError, - "PQ compressor was generated concurrently" - ) + Ok(Self { + table, + num_chunks, + phantom_data: PhantomData, + phantom_storage: PhantomData, }) } @@ -182,13 +167,6 @@ where output: MatrixBase<&mut [u8]>, ) -> Result<(), diskann::ANNError> { self.table - .get() - .ok_or_else(|| { - diskann_error!( - ErrorKind::PQError, - "PQ compressor must be generated before compression" - ) - })? .compress_into(vector, output) .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err))) } @@ -294,24 +272,22 @@ mod pq_generation_tests { assert!(!storage_provider.exists(pivot_file_name)); - let compressor = PQGeneration::::new(context); - assert!(!storage_provider.exists(pivot_file_name)); - - let result = compressor.generate(); + let result = PQGeneration::::generate_pivots(&context); assert!(result.is_ok()); assert!(storage_provider.exists(pivot_file_name)); + let compressor = PQGeneration::::new(&context).unwrap(); + assert_eq!(compressor.num_chunks, num_chunks); assert_eq!(compressor.compressed_bytes(), num_chunks); - let table = compressor.table.get().unwrap(); - assert_eq!(table.dim(), dim); - assert_eq!(table.ncenters(), num_centers); - assert_eq!(table.nchunks(), num_chunks); + assert_eq!(compressor.table.dim(), dim); + assert_eq!(compressor.table.ncenters(), num_centers); + assert_eq!(compressor.table.nchunks(), num_chunks); } #[rstest] - fn generate_creates_missing_pivots() { + fn new_preserves_missing_pivot_generation_fallback() { let storage_provider = VirtualStorageProvider::new_memory(); storage_provider .filesystem() @@ -342,10 +318,8 @@ mod pq_generation_tests { Some(data_path), ); - let compressor = PQGeneration::::new(context); - let result = compressor.generate(); - - assert!(result.is_ok()); + let compressor = PQGeneration::::new(&context); + assert!(compressor.is_ok()); assert!(storage_provider.exists(pivot_file_name)); } @@ -370,14 +344,14 @@ mod pq_generation_tests { "".to_string(), None, ); - let compressor = PQGeneration::::new(context); - let result = compressor.generate(); + let compressor = PQGeneration::::new(&context); - if let Err(x) = result.as_ref() { + if let Err(x) = compressor.as_ref() { println!("Error creating compressor: {x}"); }; - assert!(result.is_ok()); + assert!(compressor.is_ok()); + let compressor = compressor.unwrap(); let data_matrix = read_bin::(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap(); @@ -423,7 +397,7 @@ mod pq_generation_tests { "".to_string(), None, ); - let result = PQGeneration::::new(context).generate(); + let result = PQGeneration::::new(&context); assert!(result.is_err()); } } From c84dad248dae4190d2cbb6358ceb510db68ca8d7 Mon Sep 17 00:00:00 2001 From: xiajiang Date: Tue, 18 Aug 2026 20:04:30 +0800 Subject: [PATCH 4/4] Refactor quantization components to improve context handling and clarify compressor interfaces --- diskann-disk/src/build/builder/build.rs | 4 +- diskann-disk/src/storage/quant/compressor.rs | 36 +++- diskann-disk/src/storage/quant/generator.rs | 76 ++++---- diskann-disk/src/storage/quant/mod.rs | 4 +- .../src/storage/quant/pq/pq_generation.rs | 181 ++++++++++++------ 5 files changed, 192 insertions(+), 109 deletions(-) diff --git a/diskann-disk/src/build/builder/build.rs b/diskann-disk/src/build/builder/build.rs index b2c0b47bac..5db76bb1a0 100644 --- a/diskann-disk/src/build/builder/build.rs +++ b/diskann-disk/src/build/builder/build.rs @@ -157,8 +157,6 @@ where metric: self.index_configuration.dist_metric, }; - PQGeneration::::generate_pivots(&quantizer_context)?; - let generator = QuantDataGenerator::< Data::VectorDataType, PQGeneration, @@ -166,7 +164,7 @@ where self.index_writer.get_dataset_file(), self.pq_storage.get_compressed_data_path().into(), &quantizer_context, - )?; + ); generator.generate_data( storage_provider, pool, diff --git a/diskann-disk/src/storage/quant/compressor.rs b/diskann-disk/src/storage/quant/compressor.rs index acddd89940..1539863d01 100644 --- a/diskann-disk/src/storage/quant/compressor.rs +++ b/diskann-disk/src/storage/quant/compressor.rs @@ -6,30 +6,46 @@ use diskann::{utils::VectorRepr, ANNResult}; use diskann_utils::views::{MatrixView, MutMatrixView}; -/// [`QuantCompressor`] defines the interface for quantizer with [`QuantDataGenerator`] +/// [`QuantCompressor`] defines the interface for quantizers used by +/// [`super::QuantDataGenerator`]. /// /// This trait serves as a general wrapper for different quantizers, allowing them to be -/// used interchangeably with QuantDataGenerator. Any type implementing this trait +/// used interchangeably with [`super::QuantDataGenerator`]. Any type implementing this trait /// can be used to compress vector data during the data generation process. /// /// # Type Parameters -/// - `T`: The data type of the input vectors. Must impl Copy + Into + Pod + Sync -/// so that the [`QuantDataGenerator`] can parallelize computation, call compress_into and read from data file. +/// - `T`: The data type of the input vectors. Must implement `Copy + Into + Pod + Sync` +/// so that [`super::QuantDataGenerator`] can parallelize computation, call `compress_into`, +/// and read from the data file. /// /// # Associated Types -/// - [`CompressorContext`]: An overloadable type that provides initialization parameters for the compressor +/// - [`Self::CompressorContext`]: An overloadable type that provides initialization parameters +/// for the compressor. +/// - [`Self::Prepared`]: The ready-to-use compressor produced by [`Self::prepare`]. /// /// # Methods /// - `new`: Constructs a new compressor instance with the provided context. -/// - `compress`: Compresses a batch of vectors into the output buffer. -/// - `compressed_bytes`: Returns the size in bytes of each compressed vector -pub trait QuantCompressor: Sized + Sync +/// - `prepare`: Returns a ready-to-use compressor. +pub trait QuantCompressor<'a, T>: Sized where T: VectorRepr, { - type CompressorContext; + type CompressorContext: 'a; + + type Prepared: PreparedCompressor + Sync; + + fn new(context: &'a Self::CompressorContext) -> Self; + + /// Returns an error if preparation fails. + fn prepare(&self) -> ANNResult; +} - fn new(context: &Self::CompressorContext) -> ANNResult; +/// A quantizer that is ready to compress vectors. +/// +/// # Methods +/// - `compress`: Compresses a batch of vectors into the output buffer. +/// - `compressed_bytes`: Returns the size in bytes of each compressed vector. +pub trait PreparedCompressor { fn compress(&self, vector: MatrixView, output: MutMatrixView) -> ANNResult<()>; fn compressed_bytes(&self) -> usize; } diff --git a/diskann-disk/src/storage/quant/generator.rs b/diskann-disk/src/storage/quant/generator.rs index 36af775bfb..ef08e8a16b 100644 --- a/diskann-disk/src/storage/quant/generator.rs +++ b/diskann-disk/src/storage/quant/generator.rs @@ -20,39 +20,38 @@ use tracing::info; use crate::{ error::{diskann_error, ErrorKind}, - storage::quant::compressor::QuantCompressor, + storage::quant::compressor::{PreparedCompressor, QuantCompressor}, }; /// [`QuantDataGenerator`] orchestrates the process of reading vector data, applying quantization, /// and writing compressed results to storage in batches. -pub struct QuantDataGenerator +pub struct QuantDataGenerator<'a, T, Q> where T: Copy + VectorRepr, - Q: QuantCompressor, + Q: QuantCompressor<'a, T>, { pub quantizer: Q, pub data_path: String, pub compressed_data_path: String, - phantom: PhantomData, + phantom: PhantomData<&'a T>, } -impl QuantDataGenerator +impl<'a, T, Q> QuantDataGenerator<'a, T, Q> where T: Copy + VectorRepr, - Q: QuantCompressor, + Q: QuantCompressor<'a, T>, { pub fn new( data_path: String, compressed_data_path: String, - quantizer_context: &Q::CompressorContext, - ) -> ANNResult { - let quantizer = Q::new(quantizer_context)?; - Ok(Self { + quantizer_context: &'a Q::CompressorContext, + ) -> Self { + Self { data_path, compressed_data_path, - quantizer, + quantizer: Q::new(quantizer_context), phantom: PhantomData, - }) + } } /// This method reads the source data file, processes vectors in batches, compresses them @@ -61,12 +60,13 @@ where /// The implementation is adapted from generate_quantized_data_internal in pq_construction.rs // /// # Processing Flow - /// 1. Opens the source data file and validates its metadata. - /// 2. Deletes any existing output. - /// 3. Creates or opens output compressed file and writes metadata header - [num_points as i32, compressed_vector_size as i32] - /// 4. Processes data in bounded blocks. - /// 5. Compresses each block in small batch sizes in parallel to (potentially) take advantage of batch compression with quantizer - /// 6. Writes compressed blocks to the output file. + /// 1. Prepares the quantizer (training or loading a codebook as needed). + /// 2. Opens the source data file and validates its metadata. + /// 3. Deletes any existing output. + /// 4. Creates or opens output compressed file and writes metadata header - [num_points as i32, compressed_vector_size as i32] + /// 5. Processes data in bounded blocks. + /// 6. Compresses each block in small batch sizes in parallel to (potentially) take advantage of batch compression with quantizer + /// 7. Writes compressed blocks to the output file. pub fn generate_data( &self, storage_provider: &Storage, // Provider for reading source data and writing compressed results @@ -93,6 +93,7 @@ where )); } + let compressor = self.quantizer.prepare()?; let compressed_path = self.compressed_data_path.as_str(); if storage_provider.exists(compressed_path) { @@ -105,10 +106,10 @@ where data_reader.seek(SeekFrom::Start((std::mem::size_of::() * 2) as u64))?; let mut compressed_data_writer = storage_provider.create_for_write(compressed_path)?; - Metadata::new(num_points, self.quantizer.compressed_bytes())? + Metadata::new(num_points, compressor.compressed_bytes())? .write(&mut compressed_data_writer)?; - let compressed_size = self.quantizer.compressed_bytes(); + let compressed_size = compressor.compressed_bytes(); let block_size = std::cmp::min(num_points, max_block_size); let num_blocks = num_points / block_size + !num_points.is_multiple_of(block_size) as usize; @@ -163,7 +164,7 @@ where base_block .par_window_iter(BATCH_SIZE) .zip_eq(compressed_block.par_window_iter_mut(BATCH_SIZE)) - .try_for_each_in_pool(pool, |(src, dst)| self.quantizer.compress(src, dst))?; + .try_for_each_in_pool(pool, |(src, dst)| compressor.compress(src, dst))?; let write_offset = start_index * compressed_size + std::mem::size_of::() * 2; compressed_data_writer.seek(SeekFrom::Start(write_offset as u64))?; @@ -216,13 +217,20 @@ mod generator_tests { } } } - impl QuantCompressor for DummyCompressor { + impl<'a> QuantCompressor<'a, f32> for DummyCompressor { type CompressorContext = u32; + type Prepared = DummyCompressor; + + fn new(context: &'a Self::CompressorContext) -> Self { + Self::new(*context) + } - fn new(context: &Self::CompressorContext) -> ANNResult { - Ok(Self::new(*context)) + fn prepare(&self) -> ANNResult { + Ok(Self::new(self.output_dim)) } + } + impl PreparedCompressor for DummyCompressor { fn compress( &self, _vector: views::MatrixView, @@ -276,20 +284,16 @@ mod generator_tests { Ok((storage_provider, data_path, compressed_path)) } - fn create_and_call_generator( + fn create_and_call_generator<'a, F: vfs::FileSystem>( compressed_path: String, storage_provider: &VirtualStorageProvider, data_path: String, - output_dim: u32, + output_dim: &'a u32, max_block_size: usize, - ) -> (QuantDataGenerator, ANNResult<()>) { + ) -> (QuantDataGenerator<'a, f32, DummyCompressor>, ANNResult<()>) { let pool: diskann_providers::utils::RayonThreadPool = create_thread_pool_for_test(); - let generator = QuantDataGenerator::::new( - data_path, - compressed_path, - &output_dim, - ) - .unwrap(); + let generator = + QuantDataGenerator::::new(data_path, compressed_path, output_dim); let result = generator.generate_data(storage_provider, pool.as_ref(), max_block_size); (generator, result) } @@ -308,7 +312,7 @@ mod generator_tests { compressed_path.clone(), &storage_provider, data_path, - output_dim, + &output_dim, 10_000, ); @@ -350,7 +354,7 @@ mod generator_tests { compressed_path.clone(), &storage_provider, data_path, - 4, + &4, 10_000, ); assert!(result.is_err()); @@ -363,7 +367,7 @@ mod generator_tests { let (storage_provider, data_path, compressed_path) = generate_data_files(1, 8)?; let (_, result) = - create_and_call_generator(compressed_path.clone(), &storage_provider, data_path, 4, 0); + create_and_call_generator(compressed_path.clone(), &storage_provider, data_path, &4, 0); assert!(result.is_err()); assert!(!storage_provider.exists(&compressed_path)); diff --git a/diskann-disk/src/storage/quant/mod.rs b/diskann-disk/src/storage/quant/mod.rs index 9ce442ed1b..9ed4c8bfac 100644 --- a/diskann-disk/src/storage/quant/mod.rs +++ b/diskann-disk/src/storage/quant/mod.rs @@ -7,8 +7,8 @@ mod generator; pub use generator::QuantDataGenerator; pub(crate) mod pq; -pub use pq::pq_generation::{PQGeneration, PQGenerationContext}; +pub use pq::pq_generation::{PQCompressor, PQGeneration, PQGenerationContext}; pub use pq::PQData; mod compressor; -pub use compressor::QuantCompressor; +pub use compressor::{PreparedCompressor, QuantCompressor}; diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index 5aabf2254c..31be0080dc 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -22,7 +22,7 @@ use tracing::info; use crate::{ error::{diskann_error, ErrorKind}, - storage::quant::compressor::QuantCompressor, + storage::quant::compressor::{PreparedCompressor, QuantCompressor}, }; pub struct PQGenerationContext<'a, Storage> @@ -41,25 +41,41 @@ where pub num_centers: usize, } +/// Describes how to obtain a PQ codebook. Construction is cheap; the actual work happens in +/// [`QuantCompressor::prepare`]. pub struct PQGeneration<'a, T, Storage> where T: VectorRepr, Storage: StorageReadProvider + StorageWriteProvider + 'a, { + context: &'a PQGenerationContext<'a, Storage>, + phantom_data: PhantomData, +} + +/// A PQ codebook that is ready to compress vectors. +pub struct PQCompressor { table: TransposedTable, num_chunks: usize, - phantom_data: PhantomData, - phantom_storage: PhantomData<&'a Storage>, } -impl<'a, T, Storage> PQGeneration<'a, T, Storage> +impl<'a, T, Storage> QuantCompressor<'a, T> for PQGeneration<'a, T, Storage> where T: VectorRepr, Storage: StorageReadProvider + StorageWriteProvider + 'a, { - pub(crate) fn generate_pivots( - context: &PQGenerationContext<'a, Storage>, - ) -> diskann::ANNResult<()> { + type CompressorContext = PQGenerationContext<'a, Storage>; + type Prepared = PQCompressor; + + fn new(context: &'a Self::CompressorContext) -> Self { + Self { + context, + phantom_data: PhantomData, + } + } + + fn prepare(&self) -> diskann::ANNResult { + let context = self.context; + // validate that the number of chunks is correct. if context.num_chunks > context.dim { return Err(diskann_error!( @@ -109,20 +125,6 @@ where ); } - Ok(()) - } -} - -impl<'a, T, Storage> QuantCompressor for PQGeneration<'a, T, Storage> -where - T: VectorRepr, - Storage: StorageReadProvider + StorageWriteProvider + 'a, -{ - type CompressorContext = PQGenerationContext<'a, Storage>; - - fn new(context: &Self::CompressorContext) -> diskann::ANNResult { - Self::generate_pivots(context)?; - let (_, full_dim) = context .pq_storage .read_existing_pivot_metadata(context.storage_provider)?; @@ -153,14 +155,11 @@ where ) .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err)))?; - Ok(Self { - table, - num_chunks, - phantom_data: PhantomData, - phantom_storage: PhantomData, - }) + Ok(PQCompressor { table, num_chunks }) } +} +impl PreparedCompressor for PQCompressor { fn compress( &self, vector: MatrixBase<&[f32]>, @@ -182,6 +181,9 @@ where #[cfg(test)] mod pq_generation_tests { + use diskann::ANNError; + use diskann_providers::model::pq::generate_pq_pivots; + use diskann_providers::model::GeneratePivotArguments; use diskann_providers::storage::{ PQStorage, StorageReadProvider, StorageWriteProvider, VirtualStorageProvider, }; @@ -195,8 +197,8 @@ mod pq_generation_tests { use rstest::rstest; use vfs::FileSystem; - use super::{PQGeneration, PQGenerationContext}; - use crate::storage::quant::compressor::QuantCompressor; + use super::{PQCompressor, PQGeneration, PQGenerationContext}; + use crate::storage::quant::compressor::{PreparedCompressor, QuantCompressor}; const TEST_PQ_DATA_PATH: &str = "/sift/siftsmall_learn.bin"; const TEST_PQ_PIVOTS_PATH: &str = "/sift/siftsmall_learn_pq_pivots.bin"; @@ -236,17 +238,47 @@ mod pq_generation_tests { } } + #[allow(clippy::too_many_arguments)] + fn create_new_compressor<'a, F: vfs::FileSystem>( + provider: &'a VirtualStorageProvider, + dim: usize, + num_chunks: usize, + max_kmeans_reps: usize, + num_centers: usize, + p_val: f64, + pool: RayonThreadPoolRef<'a>, + pivots_path: String, + compressed_path: String, + data_path: Option<&str>, + ) -> Result { + let context = create_context( + provider, + dim, + num_chunks, + max_kmeans_reps, + num_centers, + p_val, + pool, + pivots_path, + compressed_path, + data_path, + ); + PQGeneration::::new(&context).prepare() + } + + /// Constructing a [`PQGeneration`] must not touch storage: the pivots file only appears + /// once [`QuantCompressor::prepare`] is called. #[rstest] - fn explicit_generation_creates_pivots_file() { + fn new_is_side_effect_free_and_prepare_creates_pivots() { let storage_provider = VirtualStorageProvider::new_memory(); storage_provider .filesystem() .create_dir("/pq_generation_tests") .expect("Could not create test directory"); - let pivot_file_name = "/pq_generation_tests/pivots_test.bin"; - let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin"; - let data_path = "/pq_generation_tests/data_path.bin"; + let pivot_file_name = "/pq_generation_tests/lazy_pivots_test.bin"; + let compressed_file_name = "/pq_generation_tests/lazy_compressed_not_used.bin"; + let data_path = "/pq_generation_tests/lazy_data_path.bin"; let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5); @@ -272,55 +304,91 @@ mod pq_generation_tests { assert!(!storage_provider.exists(pivot_file_name)); - let result = PQGeneration::::generate_pivots(&context); - assert!(result.is_ok()); - assert!(storage_provider.exists(pivot_file_name)); + let generator = PQGeneration::::new(&context); + assert!( + !storage_provider.exists(pivot_file_name), + "constructing the generator must not write pivots" + ); - let compressor = PQGeneration::::new(&context).unwrap(); + let compressor = generator.prepare().unwrap(); + assert!(storage_provider.exists(pivot_file_name)); - assert_eq!(compressor.num_chunks, num_chunks); assert_eq!(compressor.compressed_bytes(), num_chunks); - assert_eq!(compressor.table.dim(), dim); assert_eq!(compressor.table.ncenters(), num_centers); assert_eq!(compressor.table.nchunks(), num_chunks); } #[rstest] - fn new_preserves_missing_pivot_generation_fallback() { + fn test_create_and_load_pivots_file() { let storage_provider = VirtualStorageProvider::new_memory(); storage_provider .filesystem() .create_dir("/pq_generation_tests") .expect("Could not create test directory"); - let pivot_file_name = "/pq_generation_tests/missing_pivots.bin"; + let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin"; + let pivot_file_name_compressor = "/pq_generation_tests/compressor_pivots_test.bin"; let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin"; let data_path = "/pq_generation_tests/data_path.bin"; + let pq_storage: PQStorage = + PQStorage::new(pivot_file_name, compressed_file_name, Some(data_path)); + + let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5); + let mut train_data: Vec = VALIDATION_DATA.to_vec(); write_bin( - MatrixView::try_from(VALIDATION_DATA.as_slice(), 5, 8).unwrap(), + MatrixView::try_from(train_data.as_slice(), ndata, dim).unwrap(), &mut storage_provider.create_for_write(data_path).unwrap(), ) .unwrap(); let pool = create_thread_pool_for_test(); - let context = create_context( + generate_pq_pivots( + GeneratePivotArguments::new(ndata, dim, num_centers, num_chunks, max_k_means_reps) + .unwrap(), + true, + &mut train_data, + &pq_storage, &storage_provider, - 8, - 2, - 5, - 2, - 1.0, + diskann_providers::utils::create_rnd_provider_from_seed_in_tests(42), pool.as_ref(), - pivot_file_name.to_string(), + ) + .unwrap(); + + let compressor = create_new_compressor( + &storage_provider, + dim, + num_chunks, + max_k_means_reps, + num_centers, + 1.0, //take all the data to compute codebook + pool.as_ref(), + pivot_file_name_compressor.to_string(), compressed_file_name.to_string(), Some(data_path), ); - let compressor = PQGeneration::::new(&context); assert!(compressor.is_ok()); - assert!(storage_provider.exists(pivot_file_name)); + + let compressor = compressor.unwrap(); + assert_eq!(compressor.num_chunks, num_chunks); + assert_eq!(compressor.compressed_bytes(), num_chunks); + + assert_eq!(compressor.table.dim(), dim); + assert_eq!(compressor.table.ncenters(), num_centers); + assert_eq!(compressor.table.nchunks(), num_chunks); + + assert!(&storage_provider.exists(pivot_file_name_compressor)); + let compressor_pivots = read_bin::( + &mut storage_provider + .open_reader(pivot_file_name_compressor) + .unwrap(), + ) + .unwrap(); + let true_pivots = + read_bin::(&mut storage_provider.open_reader(pivot_file_name).unwrap()).unwrap(); + assert_eq!(compressor_pivots, true_pivots); } #[rstest] @@ -332,7 +400,7 @@ mod pq_generation_tests { let num_chunks = 1; let max_k_means_reps = 10; - let context = create_context( + let compressor = create_new_compressor( &storage_provider, dim, num_chunks, @@ -344,20 +412,18 @@ mod pq_generation_tests { "".to_string(), None, ); - let compressor = PQGeneration::::new(&context); if let Err(x) = compressor.as_ref() { println!("Error creating compressor: {x}"); }; assert!(compressor.is_ok()); - let compressor = compressor.unwrap(); let data_matrix = read_bin::(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap(); let npts = data_matrix.nrows(); let mut compressed_mat = vec![0_u8; num_chunks * npts]; - let result = compressor.compress( + let result = compressor.unwrap().compress( data_matrix.as_view(), MutMatrixView::try_from(&mut compressed_mat, npts, num_chunks).unwrap(), ); @@ -385,7 +451,7 @@ mod pq_generation_tests { let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); let pool = create_thread_pool_for_test(); let max_k_means_reps = 10; - let context = create_context( + let compressor = create_new_compressor( &storage_provider, dim, num_chunks, @@ -397,7 +463,6 @@ mod pq_generation_tests { "".to_string(), None, ); - let result = PQGeneration::::new(&context); - assert!(result.is_err()); + assert!(compressor.is_err()); } }