From 2d9f1bb25afc87a4307bf791e8b163a7d20cd5b3 Mon Sep 17 00:00:00 2001 From: xiajiang Date: Thu, 30 Jul 2026 15:43:41 +0800 Subject: [PATCH] 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()); } }