From f9591bd1f0a31f79c5304a30a6e8df8e71ba7053 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 27 Jul 2026 14:55:43 +0800 Subject: [PATCH 01/13] Converge PQ pivot loading --- .../src/model/pq/fixed_chunk_pq_table.rs | 99 ++-------- diskann-providers/src/storage/pq_storage.rs | 179 +++++++++++------- 2 files changed, 119 insertions(+), 159 deletions(-) diff --git a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs index 6dbeeb89c..babebf2eb 100644 --- a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs +++ b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs @@ -675,9 +675,8 @@ pub fn compute_pq_distance_for_pq_coordinates( mod fixed_chunk_pq_table_test { use core::ops::Range; - use crate::storage::{StorageReadProvider, VirtualStorageProvider}; + use crate::storage::{PQStorage, VirtualStorageProvider}; use approx::assert_relative_eq; - use diskann::error::ErrorContext; use diskann_utils::test_data_root; use diskann_vector::{ PureDistanceFunction, @@ -686,9 +685,17 @@ mod fixed_chunk_pq_table_test { use itertools::iproduct; use super::*; - use crate::{model::NUM_PQ_CENTROIDS, utils::read_bin_from}; + use crate::model::NUM_PQ_CENTROIDS; const DIM: usize = 128; + const PQ_PIVOTS_PATH: &str = "/sift/siftsmall_learn_pq_pivots.bin"; + + fn load_test_pivots() -> FixedChunkPQTable { + let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); + PQStorage::new(PQ_PIVOTS_PATH, "", None) + .load_pq_pivots_bin(PQ_PIVOTS_PATH, 1, &storage_provider) + .unwrap() + } #[test] fn constructor_errors() { @@ -811,14 +818,8 @@ mod fixed_chunk_pq_table_test { #[test] fn load_pivot_test() { - let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); - let pq_pivots_path: &str = "/sift/siftsmall_learn_pq_pivots.bin"; - let (dim, pq_table, chunk_offsets) = - load_pq_pivots_bin(pq_pivots_path, &1, &storage_provider).unwrap(); - let fixed_chunk_pq_table = - FixedChunkPQTable::new(dim, pq_table.into(), chunk_offsets.into()).unwrap(); + let fixed_chunk_pq_table = load_test_pivots(); - assert_eq!(dim, DIM); assert_eq!(fixed_chunk_pq_table.table.dim(), DIM); assert_eq!(fixed_chunk_pq_table.table.ncenters(), NUM_PQ_CENTROIDS); @@ -838,14 +839,7 @@ mod fixed_chunk_pq_table_test { #[test] fn calculate_distances_tests() { - let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); - - let pq_pivots_path: &str = "/sift/siftsmall_learn_pq_pivots.bin"; - - let (dim, pq_table, chunk_offsets) = - load_pq_pivots_bin(pq_pivots_path, &1, &storage_provider).unwrap(); - let fixed_chunk_pq_table = - FixedChunkPQTable::new(dim, pq_table.into(), chunk_offsets.into()).unwrap(); + let fixed_chunk_pq_table = load_test_pivots(); let query_vec: Vec = vec![ 32.39f32, 78.57f32, 50.32f32, 80.46f32, 6.47f32, 69.76f32, 94.2f32, 83.36f32, 5.8f32, @@ -994,75 +988,6 @@ mod fixed_chunk_pq_table_test { } } - type LoadPQPivotResult = (usize, Vec, Vec); - fn load_pq_pivots_bin( - pq_pivots_path: &str, - num_pq_chunks: &usize, - storage_provider: &StorageProvider, - ) -> ANNResult { - let mut reader = storage_provider - .open_reader(pq_pivots_path) - .with_context(|| format!("ERROR: Opening PQ k-means pivot file {}", pq_pivots_path))?; - - let offsets = read_bin_from::(&mut reader, 0)?; - if offsets.nrows() != 4 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. \ - Offsets don't contain correct metadata, \ - # offsets = {}, but expecting 4.", - pq_pivots_path, - offsets.nrows() - ))); - } - let file_offset_data = offsets.map(|x| x.into_usize()); - - let mut pivots = read_bin_from::(&mut reader, file_offset_data[(0, 0)])?; - - if pivots.nrows() != NUM_PQ_CENTROIDS { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", - pq_pivots_path, - pivots.nrows(), - NUM_PQ_CENTROIDS - ))); - } - let dim = pivots.ncols(); - - let centroids = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; - if centroids.nrows() != dim || centroids.ncols() != 1 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_dim = {}, \ - file_cols = {} but expecting {} entries in 1 dimension.", - pq_pivots_path, - centroids.nrows(), - centroids.ncols(), - dim - ))); - } - - pivots.row_iter_mut().for_each(|row| { - std::iter::zip(row.iter_mut(), centroids.as_slice().iter()).for_each(|(p, c)| *p += *c); - }); - - let chunk_offsets_m = read_bin_from::(&mut reader, file_offset_data[(2, 0)])?; - if chunk_offsets_m.nrows() != num_pq_chunks + 1 || chunk_offsets_m.ncols() != 1 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file at chunk offsets; \ - file has nr={}, nc={} but expecting nr={} and nc=1.", - chunk_offsets_m.nrows(), - chunk_offsets_m.ncols(), - num_pq_chunks + 1 - ))); - } - let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); - - Ok(( - dim, - pivots.into_inner().into_vec(), - chunk_offsets.into_inner().into_vec(), - )) - } - #[test] fn test_populate_chunk_distances() { let dim = 8; diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 5b6449ef8..59d0cfa40 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -29,6 +29,19 @@ type FullPivotDataType = Vec; type CentroidType = Vec; type ChunkOffsetsType = Vec; +#[derive(Debug)] +struct PivotFileParts { + pivots: Matrix, + centroid: Matrix, + chunk_offsets: Matrix, +} + +impl PivotFileParts { + fn dim(&self) -> usize { + self.pivots.ncols() + } +} + #[derive(Debug, Clone)] pub struct PQStorage { /// Pivot table path @@ -179,63 +192,18 @@ impl PQStorage { where Storage: StorageReadProvider, { - // Load file offset data. File layout: offset table(4*1) -> pivot data(num_centers*dim) -> centroid(dim*1) -> chunk offsets(num_chunks+1*1) - let reader = &mut storage_provider.open_reader(&self.pivot_data_path)?; - - let offsets = read_bin_from::(reader, 0)?; - if offsets.nrows() != 4 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. Offsets don't contain correct \ - metadata, # offsets = {}, but expecting 4.", - &self.pivot_data_path, - offsets.nrows() - ))); - } - let file_offset_data = offsets.map(|x| x.into_usize()); - - info!(" Offset data: {:?}", file_offset_data.as_slice()); - - let pivots = read_bin_from::(reader, file_offset_data[(0, 0)])?; - if pivots.nrows() != *num_centers || pivots.ncols() != *dim { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_num_centers = {}, \ - file_dim = {} but expecting {} centers in {} dimensions.", - &self.pivot_data_path, - pivots.nrows(), - pivots.ncols(), - num_centers, - dim - ))); - } - - let centroid_m = read_bin_from::(reader, file_offset_data[(1, 0)])?; - if centroid_m.nrows() != *dim || centroid_m.ncols() != 1 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_dim = {}, \ - file_cols = {} but expecting {} entries in 1 dimension.", - &self.pivot_data_path, - centroid_m.nrows(), - centroid_m.ncols(), - dim - ))); - } - - let chunk_offsets_m = read_bin_from::(reader, file_offset_data[(2, 0)])?; - if chunk_offsets_m.nrows() != *num_pq_chunks + 1 || chunk_offsets_m.ncols() != 1 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file at chunk offsets; \ - file has nr={}, nc={} but expecting nr={} and nc=1.", - chunk_offsets_m.nrows(), - chunk_offsets_m.ncols(), - num_pq_chunks + 1 - ))); - } - let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); + let parts = self.load_pivot_file_parts( + &self.pivot_data_path, + Some(*num_pq_chunks), + Some(*num_centers), + Some(*dim), + storage_provider, + )?; Ok(( - pivots.into_inner().into_vec(), - centroid_m.into_inner().into_vec(), - chunk_offsets.into_inner().into_vec(), + parts.pivots.into_inner().into_vec(), + parts.centroid.into_inner().into_vec(), + parts.chunk_offsets.into_inner().into_vec(), )) } @@ -287,7 +255,43 @@ impl PQStorage { info!("Loading PQ pivots from {}...", pq_pivots); + let PivotFileParts { + mut pivots, + centroid, + chunk_offsets, + } = self.load_pivot_file_parts( + pq_pivots, + (num_pq_chunks != 0).then_some(num_pq_chunks), + None, + None, + storage_provider, + )?; + + // If the centroid is non-zero, we need to add it to the pivots to restore the + // numeric behavior. + if centroid.as_slice().iter().any(|c| *c != 0.0) { + accum_row_inplace(pivots.as_mut_view(), centroid.as_slice()) + } + + FixedChunkPQTable::new( + pivots.ncols(), + pivots.into_inner(), + chunk_offsets.into_inner(), + ) + } + + fn load_pivot_file_parts( + &self, + pq_pivots: &str, + expected_num_pq_chunks: Option, + expected_num_centers: Option, + expected_dim: Option, + storage_provider: &Storage, + ) -> ANNResult { let mut reader = storage_provider.open_reader(pq_pivots)?; + + // File layout: offset table(4*1) -> pivot data(num_centers*dim) -> + // centroid(dim*1) -> chunk offsets(num_chunks+1*1). let offsets = read_bin_from::(&mut reader, 0)?; if offsets.nrows() != 4 { return Err(ANNError::message(format!( @@ -299,31 +303,53 @@ impl PQStorage { } let file_offset_data = offsets.map(|x| x.into_usize()); - let mut pivots = read_bin_from::(&mut reader, file_offset_data[(0, 0)])?; - if pivots.nrows() > NUM_PQ_CENTROIDS { - return Err(ANNError::message(format!( + info!(" Offset data: {:?}", file_offset_data.as_slice()); + + let pivots = read_bin_from::(&mut reader, file_offset_data[(0, 0)])?; + if let Some(num_centers) = expected_num_centers { + if pivots.nrows() != num_centers { + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", + pq_pivots, + pivots.nrows(), + num_centers + ))); + } + } else if pivots.nrows() > NUM_PQ_CENTROIDS { + return Err(ANNError::log_pq_error(format_args!( "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", pq_pivots, pivots.nrows(), NUM_PQ_CENTROIDS ))); } - let dim = pivots.ncols(); - let centroids = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; - if centroids.nrows() != dim || centroids.ncols() != 1 { - return Err(ANNError::message(format!( + if let Some(dim) = expected_dim { + if pivots.ncols() != dim { + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file {}. file_dim = {} but expecting {} dimensions.", + pq_pivots, + pivots.ncols(), + dim + ))); + } + } + + let centroid = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; + if centroid.nrows() != pivots.ncols() || centroid.ncols() != 1 { + return Err(ANNError::log_pq_error(format_args!( "Error reading pq_pivots file {}. file_dim = {}, file_cols = {} \ but expecting {} entries in 1 dimension.", pq_pivots, - centroids.nrows(), - centroids.ncols(), - dim + centroid.nrows(), + centroid.ncols(), + pivots.ncols() ))); } let chunk_offsets_m = read_bin_from::(&mut reader, file_offset_data[(2, 0)])?; - if (chunk_offsets_m.nrows() != num_pq_chunks + 1 && num_pq_chunks as u32 != 0) + if expected_num_pq_chunks + .is_some_and(|num_pq_chunks| chunk_offsets_m.nrows() != num_pq_chunks + 1) || chunk_offsets_m.ncols() != 1 { return Err(ANNError::message(format!( @@ -332,18 +358,27 @@ impl PQStorage { passed as 0 if we want to infer.", chunk_offsets_m.nrows(), chunk_offsets_m.ncols(), - num_pq_chunks + 1 + expected_num_pq_chunks.map_or(0, |num_pq_chunks| num_pq_chunks + 1) ))); } let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); - // If the centroid is non-zero, we need to add it to the pivots to restore the - // numeric behavior. - if centroids.as_slice().iter().any(|c| *c != 0.0) { - accum_row_inplace(pivots.as_mut_view(), centroids.as_slice()) + let parts = PivotFileParts { + pivots, + centroid, + chunk_offsets, + }; + if parts.chunk_offsets.nrows() < 2 + || parts.chunk_offsets[(0, 0)] != 0 + || parts.chunk_offsets[(parts.chunk_offsets.nrows() - 1, 0)] != parts.dim() + { + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file at chunk offsets; chunk offsets must start at 0, end at dim {}, and contain at least two entries.", + parts.dim() + ))); } - FixedChunkPQTable::new(dim, pivots.into_inner(), chunk_offsets.into_inner()) + Ok(parts) } /// streams data from the file, and samples each vector with probability p_val From d3e570d90cd5adb7d28f3b1336566f94004d6a81 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 27 Jul 2026 17:14:07 +0800 Subject: [PATCH 02/13] Fix clippy warning in PQ pivot loading --- diskann-providers/src/storage/pq_storage.rs | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 59d0cfa40..545bd45e8 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -324,15 +324,15 @@ impl PQStorage { ))); } - if let Some(dim) = expected_dim { - if pivots.ncols() != dim { - return Err(ANNError::log_pq_error(format_args!( - "Error reading pq_pivots file {}. file_dim = {} but expecting {} dimensions.", - pq_pivots, - pivots.ncols(), - dim - ))); - } + if let Some(dim) = expected_dim + && pivots.ncols() != dim + { + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file {}. file_dim = {} but expecting {} dimensions.", + pq_pivots, + pivots.ncols(), + dim + ))); } let centroid = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; From 053897c7f1a58a175f3bafbb8393ff10479d3c28 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 27 Jul 2026 23:02:33 +0800 Subject: [PATCH 03/13] Cover PQ pivot validation paths --- diskann-providers/src/storage/pq_storage.rs | 147 +++++++++++++++++++- 1 file changed, 146 insertions(+), 1 deletion(-) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 545bd45e8..356d6b8fb 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -194,7 +194,7 @@ impl PQStorage { { let parts = self.load_pivot_file_parts( &self.pivot_data_path, - Some(*num_pq_chunks), + (*num_pq_chunks != 0).then_some(*num_pq_chunks), Some(*num_centers), Some(*dim), storage_provider, @@ -428,6 +428,27 @@ mod pq_storage_tests { const PQ_PIVOT_PATH: &str = "/sift/siftsmall_learn_pq_pivots.bin"; const PQ_COMPRESSED_PATH: &str = "/sift/empty_pq_compressed.bin"; + fn write_test_pivots( + storage_provider: &VirtualStorageProvider, + pivot_path: &str, + num_centers: usize, + dim: usize, + centroid: Option<&[f32]>, + chunk_offsets: &[usize], + ) { + let pivots: Vec = (0..num_centers * dim).map(|i| i as f32).collect(); + PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .write_pivot_data( + &pivots, + centroid, + chunk_offsets, + num_centers, + dim, + storage_provider, + ) + .unwrap(); + } + #[test] fn new_test() { PQStorage::new(PQ_PIVOT_PATH, PQ_COMPRESSED_PATH, Some(DATA_FILE)); @@ -545,6 +566,130 @@ mod pq_storage_tests { assert_eq!(loaded_pivots, table.view_pivots().as_slice()); } + #[test] + fn load_pivot_data_infers_chunk_count_when_zero() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/infer_chunk_count_pivots.bin"; + + let num_centers = 3; + let dim = 4; + let pivots: Vec = (0..num_centers * dim).map(|i| i as f32).collect(); + let chunk_offsets = vec![0, 2, dim]; + + let pq_storage = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None); + pq_storage + .write_pivot_data( + &pivots, + None, + &chunk_offsets, + num_centers, + dim, + &storage_provider, + ) + .unwrap(); + + let (_, _, loaded_offsets) = pq_storage + .load_existing_pivot_data(&0, &num_centers, &dim, &storage_provider) + .unwrap(); + + assert_eq!(loaded_offsets, chunk_offsets); + } + + #[test] + fn load_pivot_data_rejects_mismatched_shape() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/mismatched_shape_pivots.bin"; + + write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 4]); + let pq_storage = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None); + + assert!( + pq_storage + .load_existing_pivot_data(&2, &4, &4, &storage_provider) + .is_err() + ); + assert!( + pq_storage + .load_existing_pivot_data(&2, &3, &5, &storage_provider) + .is_err() + ); + } + + #[test] + fn load_pivot_data_rejects_invalid_centroid_and_offsets() { + let storage_provider = VirtualStorageProvider::new_memory(); + + let wrong_centroid_path = "/wrong_centroid_pivots.bin"; + write_test_pivots( + &storage_provider, + wrong_centroid_path, + 3, + 4, + Some(&[1.0, 2.0, 3.0]), + &[0, 2, 4], + ); + assert!( + PQStorage::new(wrong_centroid_path, PQ_COMPRESSED_PATH, None) + .load_existing_pivot_data(&2, &3, &4, &storage_provider) + .is_err() + ); + + let wrong_count_path = "/wrong_chunk_count_pivots.bin"; + write_test_pivots(&storage_provider, wrong_count_path, 3, 4, None, &[0, 4]); + assert!( + PQStorage::new(wrong_count_path, PQ_COMPRESSED_PATH, None) + .load_existing_pivot_data(&2, &3, &4, &storage_provider) + .is_err() + ); + + let wrong_bounds_path = "/wrong_chunk_bounds_pivots.bin"; + write_test_pivots(&storage_provider, wrong_bounds_path, 3, 4, None, &[1, 4]); + assert!( + PQStorage::new(wrong_bounds_path, PQ_COMPRESSED_PATH, None) + .load_existing_pivot_data(&1, &3, &4, &storage_provider) + .is_err() + ); + } + + #[test] + fn load_pq_pivots_rejects_too_many_centers() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/too_many_centers_pivots.bin"; + + write_test_pivots( + &storage_provider, + pivot_path, + NUM_PQ_CENTROIDS + 1, + 1, + None, + &[0, 1], + ); + + assert!( + PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .is_err() + ); + } + + #[test] + fn load_pivot_data_rejects_malformed_offset_table() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/malformed_offsets_pivots.bin"; + + { + let mut writer = storage_provider.create_for_write(pivot_path).unwrap(); + let offsets = [METADATA_SIZE as u64, 0, 0]; + write_bin(MatrixView::column_vector(offsets.as_slice()), &mut writer).unwrap(); + } + + assert!( + PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .is_err() + ); + } + /// Write pivot data with a non-zero centroid, read it back, and verify that /// folding the centroid via `accum_row_inplace` produces the expected /// adjusted pivots. From d7528adfb26edffe861921e47d322b8bcd1a22f8 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 27 Jul 2026 23:33:45 +0800 Subject: [PATCH 04/13] Keep PQ pivot chunk inference scoped --- diskann-providers/src/storage/pq_storage.rs | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 356d6b8fb..14566fef9 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -194,7 +194,7 @@ impl PQStorage { { let parts = self.load_pivot_file_parts( &self.pivot_data_path, - (*num_pq_chunks != 0).then_some(*num_pq_chunks), + Some(*num_pq_chunks), Some(*num_centers), Some(*dim), storage_provider, @@ -567,7 +567,7 @@ mod pq_storage_tests { } #[test] - fn load_pivot_data_infers_chunk_count_when_zero() { + fn load_pq_pivots_infers_chunk_count_when_zero() { let storage_provider = VirtualStorageProvider::new_memory(); let pivot_path = "/infer_chunk_count_pivots.bin"; @@ -588,11 +588,11 @@ mod pq_storage_tests { ) .unwrap(); - let (_, _, loaded_offsets) = pq_storage - .load_existing_pivot_data(&0, &num_centers, &dim, &storage_provider) + let table = pq_storage + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) .unwrap(); - assert_eq!(loaded_offsets, chunk_offsets); + assert_eq!(table.view_pivots().as_slice(), pivots); } #[test] From fd65e4c7d9638ddbb5e1e5e80590afd3610385b8 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Tue, 28 Jul 2026 10:24:02 +0800 Subject: [PATCH 05/13] Clarify PQ chunk offset shape errors --- diskann-providers/src/storage/pq_storage.rs | 70 ++++++++++++++++++--- 1 file changed, 61 insertions(+), 9 deletions(-) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 14566fef9..8fde25e0c 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -348,17 +348,19 @@ impl PQStorage { } let chunk_offsets_m = read_bin_from::(&mut reader, file_offset_data[(2, 0)])?; - if expected_num_pq_chunks - .is_some_and(|num_pq_chunks| chunk_offsets_m.nrows() != num_pq_chunks + 1) - || chunk_offsets_m.ncols() != 1 + if let Some(num_pq_chunks) = expected_num_pq_chunks + && chunk_offsets_m.nrows() != num_pq_chunks + 1 { - return Err(ANNError::message(format!( - "Error reading pq_pivots file at chunk offsets; file has nr={}, nc={} \ - but expecting nr={} and nc=1. The expected num_pq_chunks should be \ - passed as 0 if we want to infer.", + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file at chunk offsets; file has nr={}, but expecting nr={}.", chunk_offsets_m.nrows(), - chunk_offsets_m.ncols(), - expected_num_pq_chunks.map_or(0, |num_pq_chunks| num_pq_chunks + 1) + num_pq_chunks + 1 + ))); + } + if chunk_offsets_m.ncols() != 1 { + return Err(ANNError::log_pq_error(format_args!( + "Error reading pq_pivots file at chunk offsets; file has nc={}, but expecting nc=1.", + chunk_offsets_m.ncols() ))); } let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); @@ -651,6 +653,56 @@ mod pq_storage_tests { ); } + #[test] + fn load_pq_pivots_reports_chunk_offset_column_mismatch() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/wrong_chunk_offset_columns_pivots.bin"; + + { + let mut writer = storage_provider.create_for_write(pivot_path).unwrap(); + let mut cumul_bytes = [0usize; 4]; + cumul_bytes[0] = METADATA_SIZE; + + writer.seek(SeekFrom::Start(cumul_bytes[0] as u64)).unwrap(); + + let pivots = [0.0, 1.0, 2.0, 3.0]; + cumul_bytes[1] = cumul_bytes[0] + + write_bin( + MatrixView::try_from(pivots.as_slice(), 2, 2).unwrap(), + &mut writer, + ) + .unwrap(); + + let centroid = [0.0, 0.0]; + cumul_bytes[2] = cumul_bytes[1] + + write_bin(MatrixView::column_vector(centroid.as_slice()), &mut writer).unwrap(); + + let chunk_offsets = [0_u32, 2_u32]; + cumul_bytes[3] = cumul_bytes[2] + + write_bin( + MatrixView::try_from(chunk_offsets.as_slice(), 1, 2).unwrap(), + &mut writer, + ) + .unwrap(); + + let offsets: Vec = cumul_bytes.iter().map(|&offset| offset as u64).collect(); + write_bin_from( + MatrixView::column_vector(offsets.as_slice()), + &mut writer, + 0, + ) + .unwrap(); + } + + let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .unwrap_err(); + let message = err.to_string(); + + assert!(message.contains("file has nc=2, but expecting nc=1")); + assert!(!message.contains("expecting nr=0")); + } + #[test] fn load_pq_pivots_rejects_too_many_centers() { let storage_provider = VirtualStorageProvider::new_memory(); From 976f1fc3bbe9dcb6258c64456c71cdcc4dc8fa15 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Tue, 28 Jul 2026 16:20:37 +0800 Subject: [PATCH 06/13] Limit PQ chunk offset bounds checks to inference --- diskann-providers/src/storage/pq_storage.rs | 56 +++++++++++++++++---- 1 file changed, 45 insertions(+), 11 deletions(-) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 8fde25e0c..c92bd8bf5 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -42,6 +42,12 @@ impl PivotFileParts { } } +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum ChunkOffsetValidation { + ShapeOnly, + ShapeAndBounds, +} + #[derive(Debug, Clone)] pub struct PQStorage { /// Pivot table path @@ -197,6 +203,7 @@ impl PQStorage { Some(*num_pq_chunks), Some(*num_centers), Some(*dim), + ChunkOffsetValidation::ShapeOnly, storage_provider, )?; @@ -264,6 +271,11 @@ impl PQStorage { (num_pq_chunks != 0).then_some(num_pq_chunks), None, None, + if num_pq_chunks == 0 { + ChunkOffsetValidation::ShapeAndBounds + } else { + ChunkOffsetValidation::ShapeOnly + }, storage_provider, )?; @@ -286,6 +298,7 @@ impl PQStorage { expected_num_pq_chunks: Option, expected_num_centers: Option, expected_dim: Option, + chunk_offset_validation: ChunkOffsetValidation, storage_provider: &Storage, ) -> ANNResult { let mut reader = storage_provider.open_reader(pq_pivots)?; @@ -370,9 +383,10 @@ impl PQStorage { centroid, chunk_offsets, }; - if parts.chunk_offsets.nrows() < 2 - || parts.chunk_offsets[(0, 0)] != 0 - || parts.chunk_offsets[(parts.chunk_offsets.nrows() - 1, 0)] != parts.dim() + if chunk_offset_validation == ChunkOffsetValidation::ShapeAndBounds + && (parts.chunk_offsets.nrows() < 2 + || parts.chunk_offsets[(0, 0)] != 0 + || parts.chunk_offsets[(parts.chunk_offsets.nrows() - 1, 0)] != parts.dim()) { return Err(ANNError::log_pq_error(format_args!( "Error reading pq_pivots file at chunk offsets; chunk offsets must start at 0, end at dim {}, and contain at least two entries.", @@ -618,7 +632,7 @@ mod pq_storage_tests { } #[test] - fn load_pivot_data_rejects_invalid_centroid_and_offsets() { + fn load_pivot_data_rejects_invalid_centroid_and_chunk_count() { let storage_provider = VirtualStorageProvider::new_memory(); let wrong_centroid_path = "/wrong_centroid_pivots.bin"; @@ -643,14 +657,34 @@ mod pq_storage_tests { .load_existing_pivot_data(&2, &3, &4, &storage_provider) .is_err() ); + } - let wrong_bounds_path = "/wrong_chunk_bounds_pivots.bin"; - write_test_pivots(&storage_provider, wrong_bounds_path, 3, 4, None, &[1, 4]); - assert!( - PQStorage::new(wrong_bounds_path, PQ_COMPRESSED_PATH, None) - .load_existing_pivot_data(&1, &3, &4, &storage_provider) - .is_err() - ); + #[test] + fn load_pivot_data_allows_legacy_chunk_bounds() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/legacy_chunk_bounds_pivots.bin"; + + write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); + + let (_, _, chunk_offsets) = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_existing_pivot_data(&1, &3, &4, &storage_provider) + .unwrap(); + + assert_eq!(chunk_offsets, vec![1, 4]); + } + + #[test] + fn load_pq_pivots_infer_rejects_invalid_chunk_bounds() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/infer_wrong_chunk_bounds_pivots.bin"; + + write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); + + let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .unwrap_err(); + + assert!(err.to_string().contains("chunk offsets must start at 0")); } #[test] From 2da540192b5890b264eab2cdca7ceb76da6591c4 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Thu, 30 Jul 2026 16:01:13 +0800 Subject: [PATCH 07/13] Refine PQ pivot loading invariants --- .../src/model/pq/fixed_chunk_pq_table.rs | 5 + .../src/model/pq/pq_construction.rs | 60 +++--- diskann-providers/src/storage/pq_storage.rs | 200 ++++++++++++------ .../src/product/tables/basic.rs | 15 +- diskann-quantization/src/views.rs | 5 + 5 files changed, 187 insertions(+), 98 deletions(-) diff --git a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs index babebf2eb..e12e1bc0a 100644 --- a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs +++ b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs @@ -144,6 +144,11 @@ impl FixedChunkPQTable { Ok(Self { table }) } + /// Wrap an already-validated basic PQ table. + pub fn from_basic_table(table: BasicTable) -> Self { + Self { table } + } + /// Get chunk number. pub fn get_num_chunks(&self) -> usize { self.table.nchunks() diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index 68fc1fed4..3c4238b60 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -340,29 +340,25 @@ where let (num_points, dim) = Metadata::read(uncompressed_data_reader)?.into_dims(); - let mut full_pivot_data: Vec; - let centroid: Vec; - let chunk_offsets: Vec; let full_dim: usize; + let mut table; if !pq_storage.pivot_data_exist(storage_provider) { return Err(ANNError::message("ERROR: PQ k-means pivot file not found.")); } else { (_, full_dim) = pq_storage.read_existing_pivot_metadata(storage_provider)?; - (full_pivot_data, centroid, chunk_offsets) = pq_storage.load_existing_pivot_data( - &num_pq_chunks, - &num_centers, - &full_dim, + let (loaded_table, centroid) = pq_storage.load_existing_pivot_table( + num_pq_chunks, + num_centers, + full_dim, storage_provider, )?; - } + table = loaded_table; - // Instead of subtracting the center from each data set component, we instead - // add it to each center. - let mut full_pivot_data_mat = - MutMatrixView::try_from(full_pivot_data.as_mut_slice(), num_centers, full_dim) - .bridge_err()?; - accum_row_inplace(full_pivot_data_mat.as_mut_view(), centroid.as_slice()); + // Instead of subtracting the center from each data set component, we instead + // add it to each center. + accum_row_inplace(table.view_pivots_mut(), centroid.as_slice()); + } pq_storage.write_compressed_pivot_metadata::( num_points, @@ -387,13 +383,8 @@ where ))?; // The compression table. - let table = TransposedTable::from_parts( - full_pivot_data_mat.as_view(), - diskann_quantization::views::ChunkOffsetsView::new(&chunk_offsets) - .bridge_err()? - .to_owned(), - ) - .map_err(ANNError::new)?; + let table = TransposedTable::from_parts(table.view_pivots(), table.view_offsets().to_owned()) + .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?; let mut buffer = vec![0.0; full_dim * block_size]; @@ -503,6 +494,19 @@ pub fn generate_pq_data_from_pivots_from_membuf>( ) .map_err(ANNError::new)?; + generate_pq_data_from_pivots_table(vector_data, &table, pq_out) +} + +fn generate_pq_data_from_pivots_table( + vector_data: &[T], + table: &diskann_quantization::product::BasicTableBase, + pq_out: &mut [u8], +) -> ANNResult<()> +where + T: Copy + Into, + U: diskann_utils::views::DenseData, + V: diskann_utils::views::DenseData, +{ let data = vector_data .iter() .map(|x| (*x).into()) @@ -546,17 +550,17 @@ pub fn generate_pq_data_from_pivots_from_membuf_batch return Err(ANNError::message("Error: Invalid PQ buffer input size.")); } + let table = BasicTableView::new( + MatrixView::try_from(pivot_data, parameters.num_centers(), dim).bridge_err()?, + ChunkOffsetsView::new(offsets).bridge_err()?, + ) + .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?; + pq_out .par_chunks_mut(num_pq_chunks) .zip(vector_data.par_chunks(dim)) .try_for_each_in_pool(pool, |(pq_slice, vector_slice)| { - generate_pq_data_from_pivots_from_membuf( - vector_slice, - pivot_data, - parameters.num_centers(), - offsets, - pq_slice, - ) + generate_pq_data_from_pivots_table(vector_slice, &table, pq_slice) }) } diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index c92bd8bf5..1479d9450 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -9,6 +9,7 @@ use diskann::{ ANNError, ANNResult, utils::{IntoUsize, VectorRepr}, }; +use diskann_quantization::{product::BasicTable, views::ChunkOffsetsBase}; use diskann_utils::{ io::{Metadata, read_bin, write_bin}, views::{Matrix, MatrixView}, @@ -36,18 +37,6 @@ struct PivotFileParts { chunk_offsets: Matrix, } -impl PivotFileParts { - fn dim(&self) -> usize { - self.pivots.ncols() - } -} - -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -enum ChunkOffsetValidation { - ShapeOnly, - ShapeAndBounds, -} - #[derive(Debug, Clone)] pub struct PQStorage { /// Pivot table path @@ -195,23 +184,43 @@ impl PQStorage { dim: &usize, storage_provider: &Storage, ) -> ANNResult<(FullPivotDataType, CentroidType, ChunkOffsetsType)> + where + Storage: StorageReadProvider, + { + let (table, centroid) = + self.load_existing_pivot_table(*num_pq_chunks, *num_centers, *dim, storage_provider)?; + let (pivots, chunk_offsets) = table.into_parts(); + + Ok(( + pivots.into_inner().into_vec(), + centroid.into_inner().into_vec(), + chunk_offsets.into_inner().into_vec(), + )) + } + + /// Load the raw pivot table and centroid from the configured pivot file. + /// + /// The returned table has not had the centroid folded into its pivots. + pub fn load_existing_pivot_table( + &self, + num_pq_chunks: usize, + num_centers: usize, + dim: usize, + storage_provider: &Storage, + ) -> ANNResult<(BasicTable, Matrix)> where Storage: StorageReadProvider, { let parts = self.load_pivot_file_parts( &self.pivot_data_path, - Some(*num_pq_chunks), - Some(*num_centers), - Some(*dim), - ChunkOffsetValidation::ShapeOnly, + Some(num_pq_chunks), + Some(num_centers), + Some(dim), storage_provider, )?; - - Ok(( - parts.pivots.into_inner().into_vec(), - parts.centroid.into_inner().into_vec(), - parts.chunk_offsets.into_inner().into_vec(), - )) + let centroid = parts.centroid.clone(); + let table = Self::pivot_file_parts_into_basic_table(&self.pivot_data_path, parts)?; + Ok((table, centroid)) } /// Load the compressed pq dataset from file. @@ -255,6 +264,30 @@ impl PQStorage { pq_pivots: &str, num_pq_chunks: usize, storage_provider: &Storage, + ) -> ANNResult { + if num_pq_chunks == 0 { + return Err(ANNError::log_pq_error( + "num_pq_chunks must be non-zero; use load_pq_pivots_bin_infer_chunks to infer from the file.", + )); + } + + self.load_pq_pivots_bin_impl(pq_pivots, Some(num_pq_chunks), storage_provider) + } + + /// Load pre-trained pivot table, inferring the number of chunks from the file. + pub fn load_pq_pivots_bin_infer_chunks( + &self, + pq_pivots: &str, + storage_provider: &Storage, + ) -> ANNResult { + self.load_pq_pivots_bin_impl(pq_pivots, None, storage_provider) + } + + fn load_pq_pivots_bin_impl( + &self, + pq_pivots: &str, + expected_num_pq_chunks: Option, + storage_provider: &Storage, ) -> ANNResult { if !storage_provider.exists(pq_pivots) { return Err(ANNError::message("ERROR: PQ k-means pivot file not found.")); @@ -262,34 +295,22 @@ impl PQStorage { info!("Loading PQ pivots from {}...", pq_pivots); - let PivotFileParts { - mut pivots, - centroid, - chunk_offsets, - } = self.load_pivot_file_parts( + let mut parts = self.load_pivot_file_parts( pq_pivots, - (num_pq_chunks != 0).then_some(num_pq_chunks), + expected_num_pq_chunks, None, None, - if num_pq_chunks == 0 { - ChunkOffsetValidation::ShapeAndBounds - } else { - ChunkOffsetValidation::ShapeOnly - }, storage_provider, )?; // If the centroid is non-zero, we need to add it to the pivots to restore the // numeric behavior. - if centroid.as_slice().iter().any(|c| *c != 0.0) { - accum_row_inplace(pivots.as_mut_view(), centroid.as_slice()) + if parts.centroid.as_slice().iter().any(|c| *c != 0.0) { + accum_row_inplace(parts.pivots.as_mut_view(), parts.centroid.as_slice()) } - FixedChunkPQTable::new( - pivots.ncols(), - pivots.into_inner(), - chunk_offsets.into_inner(), - ) + let table = Self::pivot_file_parts_into_basic_table(pq_pivots, parts)?; + Ok(FixedChunkPQTable::from_basic_table(table)) } fn load_pivot_file_parts( @@ -298,20 +319,24 @@ impl PQStorage { expected_num_pq_chunks: Option, expected_num_centers: Option, expected_dim: Option, - chunk_offset_validation: ChunkOffsetValidation, storage_provider: &Storage, ) -> ANNResult { let mut reader = storage_provider.open_reader(pq_pivots)?; // File layout: offset table(4*1) -> pivot data(num_centers*dim) -> // centroid(dim*1) -> chunk offsets(num_chunks+1*1). + // + // The expected values here are file-format checks only. Structural PQ table + // invariants, such as chunk-offset monotonicity and pivot/offset dimension + // agreement, are validated by `ChunkOffsetsBase` and `BasicTable`. let offsets = read_bin_from::(&mut reader, 0)?; - if offsets.nrows() != 4 { - return Err(ANNError::message(format!( + if offsets.nrows() != 4 || offsets.ncols() != 1 { + return Err(ANNError::log_pq_error(format_args!( "Error reading pq_pivots file {}. Offsets don't contain correct metadata, \ - # offsets = {}, but expecting 4.", + file has nr={}, nc={}, but expecting nr=4 and nc=1.", pq_pivots, - offsets.nrows() + offsets.nrows(), + offsets.ncols() ))); } let file_offset_data = offsets.map(|x| x.into_usize()); @@ -378,23 +403,32 @@ impl PQStorage { } let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); - let parts = PivotFileParts { + Ok(PivotFileParts { pivots, centroid, chunk_offsets, - }; - if chunk_offset_validation == ChunkOffsetValidation::ShapeAndBounds - && (parts.chunk_offsets.nrows() < 2 - || parts.chunk_offsets[(0, 0)] != 0 - || parts.chunk_offsets[(parts.chunk_offsets.nrows() - 1, 0)] != parts.dim()) - { - return Err(ANNError::log_pq_error(format_args!( - "Error reading pq_pivots file at chunk offsets; chunk offsets must start at 0, end at dim {}, and contain at least two entries.", - parts.dim() - ))); - } + }) + } - Ok(parts) + fn pivot_file_parts_into_basic_table( + pq_pivots: &str, + parts: PivotFileParts, + ) -> ANNResult { + let offsets = ChunkOffsetsBase::new(parts.chunk_offsets.into_inner()).map_err(|err| { + ANNError::log_pq_error(format_args!( + "Error constructing chunk offsets from pq_pivots file {}: {}", + pq_pivots, + diskann_quantization::error::format(&err) + )) + })?; + + BasicTable::new(parts.pivots, offsets).map_err(|err| { + ANNError::log_pq_error(format_args!( + "Error constructing PQ table from pq_pivots file {}: {}", + pq_pivots, + diskann_quantization::error::format(&err) + )) + }) } /// streams data from the file, and samples each vector with probability p_val @@ -583,7 +617,7 @@ mod pq_storage_tests { } #[test] - fn load_pq_pivots_infers_chunk_count_when_zero() { + fn load_pq_pivots_infer_chunks_loads_without_expected_count() { let storage_provider = VirtualStorageProvider::new_memory(); let pivot_path = "/infer_chunk_count_pivots.bin"; @@ -605,12 +639,26 @@ mod pq_storage_tests { .unwrap(); let table = pq_storage - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) .unwrap(); assert_eq!(table.view_pivots().as_slice(), pivots); } + #[test] + fn load_pq_pivots_rejects_zero_chunk_count() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/zero_chunk_count_pivots.bin"; + + write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 4]); + + let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .unwrap_err(); + + assert!(err.to_string().contains("num_pq_chunks must be non-zero")); + } + #[test] fn load_pivot_data_rejects_mismatched_shape() { let storage_provider = VirtualStorageProvider::new_memory(); @@ -660,17 +708,31 @@ mod pq_storage_tests { } #[test] - fn load_pivot_data_allows_legacy_chunk_bounds() { + fn load_pivot_data_rejects_invalid_chunk_bounds() { let storage_provider = VirtualStorageProvider::new_memory(); let pivot_path = "/legacy_chunk_bounds_pivots.bin"; write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); - let (_, _, chunk_offsets) = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) .load_existing_pivot_data(&1, &3, &4, &storage_provider) - .unwrap(); + .unwrap_err(); - assert_eq!(chunk_offsets, vec![1, 4]); + assert!(err.to_string().contains("offsets must begin at 0")); + } + + #[test] + fn load_pivot_data_rejects_chunk_offsets_dim_mismatch() { + let storage_provider = VirtualStorageProvider::new_memory(); + let pivot_path = "/chunk_offsets_dim_mismatch_pivots.bin"; + + write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 3]); + + let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .load_existing_pivot_data(&2, &3, &4, &storage_provider) + .unwrap_err(); + + assert!(err.to_string().contains("offsets expect 3")); } #[test] @@ -681,10 +743,10 @@ mod pq_storage_tests { write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) .unwrap_err(); - assert!(err.to_string().contains("chunk offsets must start at 0")); + assert!(err.to_string().contains("offsets must begin at 0")); } #[test] @@ -729,7 +791,7 @@ mod pq_storage_tests { } let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) .unwrap_err(); let message = err.to_string(); @@ -753,7 +815,7 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) .is_err() ); } @@ -771,7 +833,7 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) .is_err() ); } diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index 9469fa58d..822940645 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -5,7 +5,7 @@ use crate::traits::CompressInto; use crate::views::{ChunkOffsetsBase, ChunkOffsetsView}; -use diskann_utils::views::{DenseData, MatrixBase, MatrixView}; +use diskann_utils::views::{DenseData, MatrixBase, MatrixView, MutDenseData, MutMatrixView}; use diskann_vector::{PureDistanceFunction, distance::SquaredL2}; use thiserror::Error; @@ -86,6 +86,14 @@ where self.pivots.as_view() } + /// Return a mutable view over the pivot table. + pub fn view_pivots_mut(&mut self) -> MutMatrixView<'_, f32> + where + T: MutDenseData, + { + self.pivots.as_mut_view() + } + /// Return a view over the schema offsets. pub fn view_offsets(&self) -> ChunkOffsetsView<'_> { self.offsets.as_view() @@ -105,6 +113,11 @@ where pub fn dim(&self) -> usize { self.pivots.ncols() } + + /// Consume this table and return the underlying pivots and offsets. + pub fn into_parts(self) -> (MatrixBase, ChunkOffsetsBase) { + (self.pivots, self.offsets) + } } #[derive(Error, Debug)] diff --git a/diskann-quantization/src/views.rs b/diskann-quantization/src/views.rs index 04c4a0953..03ed01223 100644 --- a/diskann-quantization/src/views.rs +++ b/diskann-quantization/src/views.rs @@ -211,6 +211,11 @@ where pub fn as_slice(&self) -> &[usize] { self.offsets.as_slice() } + + /// Consume the offsets, returning the inner representation. + pub fn into_inner(self) -> T { + self.offsets + } } pub type ChunkOffsetsView<'a> = ChunkOffsetsBase<&'a [usize]>; From 98888fae036ddc387a73fc5bb5885f3d3911caea Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Thu, 30 Jul 2026 17:19:38 +0800 Subject: [PATCH 08/13] Adapt PQ pivot loading after ANNError cleanup --- .../src/model/pq/pq_construction.rs | 4 +- diskann-providers/src/storage/pq_storage.rs | 43 ++++++++++--------- 2 files changed, 24 insertions(+), 23 deletions(-) diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index 3c4238b60..1aef866aa 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -384,7 +384,7 @@ where // The compression table. let table = TransposedTable::from_parts(table.view_pivots(), table.view_offsets().to_owned()) - .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?; + .map_err(|err| ANNError::message(diskann_quantization::error::format(&err)))?; let mut buffer = vec![0.0; full_dim * block_size]; @@ -554,7 +554,7 @@ pub fn generate_pq_data_from_pivots_from_membuf_batch MatrixView::try_from(pivot_data, parameters.num_centers(), dim).bridge_err()?, ChunkOffsetsView::new(offsets).bridge_err()?, ) - .map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?; + .map_err(|err| ANNError::message(diskann_quantization::error::format(&err)))?; pq_out .par_chunks_mut(num_pq_chunks) diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 1479d9450..3e6daf121 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -265,13 +265,11 @@ impl PQStorage { num_pq_chunks: usize, storage_provider: &Storage, ) -> ANNResult { - if num_pq_chunks == 0 { - return Err(ANNError::log_pq_error( - "num_pq_chunks must be non-zero; use load_pq_pivots_bin_infer_chunks to infer from the file.", - )); - } - - self.load_pq_pivots_bin_impl(pq_pivots, Some(num_pq_chunks), storage_provider) + self.load_pq_pivots_bin_impl( + pq_pivots, + (num_pq_chunks != 0).then_some(num_pq_chunks), + storage_provider, + ) } /// Load pre-trained pivot table, inferring the number of chunks from the file. @@ -331,7 +329,7 @@ impl PQStorage { // agreement, are validated by `ChunkOffsetsBase` and `BasicTable`. let offsets = read_bin_from::(&mut reader, 0)?; if offsets.nrows() != 4 || offsets.ncols() != 1 { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file {}. Offsets don't contain correct metadata, \ file has nr={}, nc={}, but expecting nr=4 and nc=1.", pq_pivots, @@ -346,7 +344,7 @@ impl PQStorage { let pivots = read_bin_from::(&mut reader, file_offset_data[(0, 0)])?; if let Some(num_centers) = expected_num_centers { if pivots.nrows() != num_centers { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", pq_pivots, pivots.nrows(), @@ -354,7 +352,7 @@ impl PQStorage { ))); } } else if pivots.nrows() > NUM_PQ_CENTROIDS { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", pq_pivots, pivots.nrows(), @@ -365,7 +363,7 @@ impl PQStorage { if let Some(dim) = expected_dim && pivots.ncols() != dim { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file {}. file_dim = {} but expecting {} dimensions.", pq_pivots, pivots.ncols(), @@ -375,7 +373,7 @@ impl PQStorage { let centroid = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; if centroid.nrows() != pivots.ncols() || centroid.ncols() != 1 { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file {}. file_dim = {}, file_cols = {} \ but expecting {} entries in 1 dimension.", pq_pivots, @@ -389,14 +387,14 @@ impl PQStorage { if let Some(num_pq_chunks) = expected_num_pq_chunks && chunk_offsets_m.nrows() != num_pq_chunks + 1 { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file at chunk offsets; file has nr={}, but expecting nr={}.", chunk_offsets_m.nrows(), num_pq_chunks + 1 ))); } if chunk_offsets_m.ncols() != 1 { - return Err(ANNError::log_pq_error(format_args!( + return Err(ANNError::message(format!( "Error reading pq_pivots file at chunk offsets; file has nc={}, but expecting nc=1.", chunk_offsets_m.ncols() ))); @@ -415,7 +413,7 @@ impl PQStorage { parts: PivotFileParts, ) -> ANNResult { let offsets = ChunkOffsetsBase::new(parts.chunk_offsets.into_inner()).map_err(|err| { - ANNError::log_pq_error(format_args!( + ANNError::message(format!( "Error constructing chunk offsets from pq_pivots file {}: {}", pq_pivots, diskann_quantization::error::format(&err) @@ -423,7 +421,7 @@ impl PQStorage { })?; BasicTable::new(parts.pivots, offsets).map_err(|err| { - ANNError::log_pq_error(format_args!( + ANNError::message(format!( "Error constructing PQ table from pq_pivots file {}: {}", pq_pivots, diskann_quantization::error::format(&err) @@ -646,17 +644,20 @@ mod pq_storage_tests { } #[test] - fn load_pq_pivots_rejects_zero_chunk_count() { + fn load_pq_pivots_zero_chunk_count_infers_from_file() { let storage_provider = VirtualStorageProvider::new_memory(); let pivot_path = "/zero_chunk_count_pivots.bin"; + let pivots: Vec = (0..12).map(|i| i as f32).collect(); - write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 4]); + PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + .write_pivot_data(&pivots, None, &[0, 2, 4], 3, 4, &storage_provider) + .unwrap(); - let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) + let table = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) .load_pq_pivots_bin(pivot_path, 0, &storage_provider) - .unwrap_err(); + .unwrap(); - assert!(err.to_string().contains("num_pq_chunks must be non-zero")); + assert_eq!(table.view_pivots().as_slice(), pivots); } #[test] From 4318a5683ef3de839dff643f1c2a64d0cb0d547c Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Fri, 31 Jul 2026 10:19:14 +0800 Subject: [PATCH 09/13] Simplify PQ pivot loading --- .../src/storage/quant/pq/pq_generation.rs | 48 +++---- .../async_/experimental/multi_pq_async.rs | 21 +-- .../src/model/pq/pq_construction.rs | 42 +++--- diskann-providers/src/storage/pq_storage.rs | 135 +++++++----------- 4 files changed, 99 insertions(+), 147 deletions(-) diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index c5297ae9c..132c88026 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -8,12 +8,9 @@ use std::{marker::PhantomData, time::Instant}; use diskann::utils::VectorRepr; use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider}; use diskann_providers::{ - model::{ - pq::{accum_row_inplace, generate_pq_pivots}, - GeneratePivotArguments, - }, + model::{pq::generate_pq_pivots, GeneratePivotArguments}, storage::PQStorage, - utils::{BridgeErr, RayonThreadPoolRef}, + utils::RayonThreadPoolRef, }; use diskann_quantization::{error::Format, product::TransposedTable, CompressInto}; use diskann_utils::views::MatrixBase; @@ -113,32 +110,27 @@ where .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( - &num_chunks, - &context.num_centers, - &full_dim, - context.storage_provider, - )?; + let table = context.pq_storage.load_pivots( + context.pq_storage.get_pivot_data_path(), + Some(num_chunks), + context.storage_provider, + )?; - let mut full_pivot_data_mat = diskann_utils::views::MutMatrixView::try_from( - full_pivot_data.as_mut_slice(), - context.num_centers, - full_dim, - ) - .bridge_err()?; - - accum_row_inplace(full_pivot_data_mat.as_mut_view(), centroid.as_slice()); + if table.ncenters() != context.num_centers || table.dim() != full_dim { + return Err(diskann_error!( + ErrorKind::PQError, + "PQ pivot table mismatch: file has {} centers in {} dimensions but expected {} centers in {} dimensions.", + table.ncenters(), + table.dim(), + context.num_centers, + full_dim + )); + } - let table = TransposedTable::from_parts( - full_pivot_data_mat.as_view(), - diskann_quantization::views::ChunkOffsetsView::new(&chunk_offsets) - .bridge_err()? - .to_owned(), - ) - .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err)))?; + let table = + TransposedTable::from_parts(table.view_pivots(), table.view_offsets().to_owned()) + .map_err(|err| diskann_error!(ErrorKind::PQError, "{}", Format(err)))?; Ok(Self { table, diff --git a/diskann-providers/src/model/graph/provider/async_/experimental/multi_pq_async.rs b/diskann-providers/src/model/graph/provider/async_/experimental/multi_pq_async.rs index 012253507..d22b77802 100644 --- a/diskann-providers/src/model/graph/provider/async_/experimental/multi_pq_async.rs +++ b/diskann-providers/src/model/graph/provider/async_/experimental/multi_pq_async.rs @@ -7,13 +7,14 @@ use std::sync::{Arc, Mutex}; use arc_swap::{ArcSwap, Guard}; use diskann::{ANNError, ANNResult, error::IntoANNResult, utils::VectorRepr}; +use diskann_quantization::CompressInto; use diskann_utils::lazy_format; use diskann_vector::{DistanceFunction, PreprocessedDistanceFunction, distance::Metric}; use rand::{Rng, SeedableRng, rngs::StdRng}; -use crate::model::{ - FixedChunkPQTable, - pq::{distance::multi, generate_pq_data_from_pivots_from_membuf}, +use crate::{ + model::{FixedChunkPQTable, pq::distance::multi}, + utils::BridgeErr, }; /// The discriminant type for PQ vector versions. @@ -155,17 +156,9 @@ impl TestMultiPQProviderAsync { }; let mut quant_vector: Vec = vec![0; table.get_num_chunks()]; - if generate_pq_data_from_pivots_from_membuf( - &vector_f32, - table.get_pq_table(), - table.get_num_centers(), - table.get_chunk_offsets(), - &mut quant_vector, - ) - .is_err() - { - return Err(ANNError::message("Error in generating PQ data.")); - } + table + .compress_into(vector_f32.as_slice(), &mut quant_vector) + .bridge_err()?; let new = Arc::new(VersionedPQVector::new(quant_vector, version)); self.quant_vectors[id].swap(new); diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index 1aef866aa..a95e43fa3 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -341,23 +341,27 @@ where let (num_points, dim) = Metadata::read(uncompressed_data_reader)?.into_dims(); let full_dim: usize; - let mut table; + let table; if !pq_storage.pivot_data_exist(storage_provider) { return Err(ANNError::message("ERROR: PQ k-means pivot file not found.")); } else { (_, full_dim) = pq_storage.read_existing_pivot_metadata(storage_provider)?; - let (loaded_table, centroid) = pq_storage.load_existing_pivot_table( - num_pq_chunks, - num_centers, - full_dim, + table = pq_storage.load_pivots( + pq_storage.get_pivot_data_path(), + Some(num_pq_chunks), storage_provider, )?; - table = loaded_table; - // Instead of subtracting the center from each data set component, we instead - // add it to each center. - accum_row_inplace(table.view_pivots_mut(), centroid.as_slice()); + if table.ncenters() != num_centers || table.dim() != full_dim { + return Err(ANNError::message(format!( + "PQ pivot table mismatch: file has {} centers in {} dimensions but expected {} centers in {} dimensions.", + table.ncenters(), + table.dim(), + num_centers, + full_dim + ))); + } } pq_storage.write_compressed_pivot_metadata::( @@ -494,19 +498,6 @@ pub fn generate_pq_data_from_pivots_from_membuf>( ) .map_err(ANNError::new)?; - generate_pq_data_from_pivots_table(vector_data, &table, pq_out) -} - -fn generate_pq_data_from_pivots_table( - vector_data: &[T], - table: &diskann_quantization::product::BasicTableBase, - pq_out: &mut [u8], -) -> ANNResult<()> -where - T: Copy + Into, - U: diskann_utils::views::DenseData, - V: diskann_utils::views::DenseData, -{ let data = vector_data .iter() .map(|x| (*x).into()) @@ -559,8 +550,11 @@ pub fn generate_pq_data_from_pivots_from_membuf_batch pq_out .par_chunks_mut(num_pq_chunks) .zip(vector_data.par_chunks(dim)) - .try_for_each_in_pool(pool, |(pq_slice, vector_slice)| { - generate_pq_data_from_pivots_table(vector_slice, &table, pq_slice) + .try_for_each_in_pool(pool, |(pq_slice, vector)| { + let data = vector.iter().map(|x| (*x).into()).collect::>(); + table + .compress_into(data.as_slice(), pq_slice) + .map_err(ANNError::new) }) } diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 3e6daf121..7bc06adf0 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -30,13 +30,6 @@ type FullPivotDataType = Vec; type CentroidType = Vec; type ChunkOffsetsType = Vec; -#[derive(Debug)] -struct PivotFileParts { - pivots: Matrix, - centroid: Matrix, - chunk_offsets: Matrix, -} - #[derive(Debug, Clone)] pub struct PQStorage { /// Pivot table path @@ -187,8 +180,15 @@ impl PQStorage { where Storage: StorageReadProvider, { - let (table, centroid) = - self.load_existing_pivot_table(*num_pq_chunks, *num_centers, *dim, storage_provider)?; + let (pivots, centroid, chunk_offsets) = self.read_pivot_file( + &self.pivot_data_path, + Some(*num_pq_chunks), + Some(*num_centers), + Some(*dim), + storage_provider, + )?; + let table = + Self::pivot_data_into_basic_table(&self.pivot_data_path, pivots, chunk_offsets)?; let (pivots, chunk_offsets) = table.into_parts(); Ok(( @@ -198,29 +198,31 @@ impl PQStorage { )) } - /// Load the raw pivot table and centroid from the configured pivot file. + /// Load the effective PQ pivot table from a pivot file. /// - /// The returned table has not had the centroid folded into its pivots. - pub fn load_existing_pivot_table( + /// If `expected_num_pq_chunks` is `None`, the chunk count is inferred from the + /// file. The loader verifies the pivot file layout and folds any stored legacy + /// centroid into the pivots. `BasicTable::new` validates the resulting table + /// invariants, including chunk-offset bounds and monotonicity. + pub fn load_pivots( &self, - num_pq_chunks: usize, - num_centers: usize, - dim: usize, + pq_pivots: &str, + expected_num_pq_chunks: Option, storage_provider: &Storage, - ) -> ANNResult<(BasicTable, Matrix)> - where - Storage: StorageReadProvider, - { - let parts = self.load_pivot_file_parts( - &self.pivot_data_path, - Some(num_pq_chunks), - Some(num_centers), - Some(dim), + ) -> ANNResult { + let (mut pivots, centroid, chunk_offsets) = self.read_pivot_file( + pq_pivots, + expected_num_pq_chunks, + None, + None, storage_provider, )?; - let centroid = parts.centroid.clone(); - let table = Self::pivot_file_parts_into_basic_table(&self.pivot_data_path, parts)?; - Ok((table, centroid)) + + if centroid.as_slice().iter().any(|c| *c != 0.0) { + accum_row_inplace(pivots.as_mut_view(), centroid.as_slice()) + } + + Self::pivot_data_into_basic_table(pq_pivots, pivots, chunk_offsets) } /// Load the compressed pq dataset from file. @@ -265,60 +267,30 @@ impl PQStorage { num_pq_chunks: usize, storage_provider: &Storage, ) -> ANNResult { - self.load_pq_pivots_bin_impl( + let table = self.load_pivots( pq_pivots, (num_pq_chunks != 0).then_some(num_pq_chunks), storage_provider, - ) - } - - /// Load pre-trained pivot table, inferring the number of chunks from the file. - pub fn load_pq_pivots_bin_infer_chunks( - &self, - pq_pivots: &str, - storage_provider: &Storage, - ) -> ANNResult { - self.load_pq_pivots_bin_impl(pq_pivots, None, storage_provider) + )?; + Ok(FixedChunkPQTable::from_basic_table(table)) } - fn load_pq_pivots_bin_impl( + fn read_pivot_file( &self, pq_pivots: &str, expected_num_pq_chunks: Option, + expected_num_centers: Option, + expected_dim: Option, storage_provider: &Storage, - ) -> ANNResult { + ) -> ANNResult<(Matrix, Matrix, Matrix)> { if !storage_provider.exists(pq_pivots) { - return Err(ANNError::message("ERROR: PQ k-means pivot file not found.")); + return Err(ANNError::message(format!( + "ERROR: PQ k-means pivot file not found: {pq_pivots}." + ))); } info!("Loading PQ pivots from {}...", pq_pivots); - let mut parts = self.load_pivot_file_parts( - pq_pivots, - expected_num_pq_chunks, - None, - None, - storage_provider, - )?; - - // If the centroid is non-zero, we need to add it to the pivots to restore the - // numeric behavior. - if parts.centroid.as_slice().iter().any(|c| *c != 0.0) { - accum_row_inplace(parts.pivots.as_mut_view(), parts.centroid.as_slice()) - } - - let table = Self::pivot_file_parts_into_basic_table(pq_pivots, parts)?; - Ok(FixedChunkPQTable::from_basic_table(table)) - } - - fn load_pivot_file_parts( - &self, - pq_pivots: &str, - expected_num_pq_chunks: Option, - expected_num_centers: Option, - expected_dim: Option, - storage_provider: &Storage, - ) -> ANNResult { let mut reader = storage_provider.open_reader(pq_pivots)?; // File layout: offset table(4*1) -> pivot data(num_centers*dim) -> @@ -401,18 +373,15 @@ impl PQStorage { } let chunk_offsets = chunk_offsets_m.map(|x| x.into_usize()); - Ok(PivotFileParts { - pivots, - centroid, - chunk_offsets, - }) + Ok((pivots, centroid, chunk_offsets)) } - fn pivot_file_parts_into_basic_table( + fn pivot_data_into_basic_table( pq_pivots: &str, - parts: PivotFileParts, + pivots: Matrix, + chunk_offsets: Matrix, ) -> ANNResult { - let offsets = ChunkOffsetsBase::new(parts.chunk_offsets.into_inner()).map_err(|err| { + let offsets = ChunkOffsetsBase::new(chunk_offsets.into_inner()).map_err(|err| { ANNError::message(format!( "Error constructing chunk offsets from pq_pivots file {}: {}", pq_pivots, @@ -420,7 +389,7 @@ impl PQStorage { )) })?; - BasicTable::new(parts.pivots, offsets).map_err(|err| { + BasicTable::new(pivots, offsets).map_err(|err| { ANNError::message(format!( "Error constructing PQ table from pq_pivots file {}: {}", pq_pivots, @@ -460,6 +429,10 @@ impl PQStorage { pub fn get_compressed_data_path(&self) -> &str { &self.compressed_data_path } + + pub fn get_pivot_data_path(&self) -> &str { + &self.pivot_data_path + } } #[cfg(test)] @@ -637,7 +610,7 @@ mod pq_storage_tests { .unwrap(); let table = pq_storage - .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) + .load_pivots(pivot_path, None, &storage_provider) .unwrap(); assert_eq!(table.view_pivots().as_slice(), pivots); @@ -744,7 +717,7 @@ mod pq_storage_tests { write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) + .load_pivots(pivot_path, None, &storage_provider) .unwrap_err(); assert!(err.to_string().contains("offsets must begin at 0")); @@ -792,7 +765,7 @@ mod pq_storage_tests { } let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) + .load_pivots(pivot_path, None, &storage_provider) .unwrap_err(); let message = err.to_string(); @@ -816,7 +789,7 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) + .load_pivots(pivot_path, None, &storage_provider) .is_err() ); } @@ -834,7 +807,7 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin_infer_chunks(pivot_path, &storage_provider) + .load_pivots(pivot_path, None, &storage_provider) .is_err() ); } From 684ae61a41b7c3a7ed268e8dcc0af778c59c0fcd Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 3 Aug 2026 13:54:31 +0800 Subject: [PATCH 10/13] Consolidate PQ pivot loading APIs --- diskann-disk/src/storage/disk_index_reader.rs | 10 +- .../src/storage/quant/pq/pq_generation.rs | 15 +- .../fast_memory_quant_vector_provider.rs | 12 +- .../async_/memory_quant_vector_provider.rs | 12 +- .../src/model/pq/fixed_chunk_pq_table.rs | 14 +- .../src/model/pq/pq_construction.rs | 37 +-- diskann-providers/src/storage/pq_storage.rs | 223 +++--------------- .../src/product/tables/basic.rs | 10 +- 8 files changed, 96 insertions(+), 237 deletions(-) diff --git a/diskann-disk/src/storage/disk_index_reader.rs b/diskann-disk/src/storage/disk_index_reader.rs index 593fe7280..e776bf056 100644 --- a/diskann-disk/src/storage/disk_index_reader.rs +++ b/diskann-disk/src/storage/disk_index_reader.rs @@ -6,7 +6,9 @@ use std::sync::Arc; use diskann::ANNResult; use diskann_providers::storage::StorageReadProvider; -use diskann_providers::{storage::PQStorage, utils::load_metadata_from_file}; +use diskann_providers::{ + model::FixedChunkPQTable, storage::PQStorage, utils::load_metadata_from_file, +}; use crate::search::pq::PQData; use tracing::info; @@ -29,11 +31,7 @@ impl DiskIndexReader { storage_provider: &Storage, ) -> ANNResult { let pq_storage = PQStorage::new(&pq_pivot_path, &pq_compressed_data_path, None); - let pq_pivot_table = pq_storage.load_pq_pivots_bin::( - &pq_pivot_path, - 0, // Use 0 to infer num_pq_chunks from the file - storage_provider, - )?; + let pq_pivot_table: FixedChunkPQTable = pq_storage.load_pivots(storage_provider)?.into(); // Auto-detect number of points from compressed PQ file metadata let metadata = load_metadata_from_file(storage_provider, &pq_compressed_data_path)?; diff --git a/diskann-disk/src/storage/quant/pq/pq_generation.rs b/diskann-disk/src/storage/quant/pq/pq_generation.rs index 132c88026..a51351692 100644 --- a/diskann-disk/src/storage/quant/pq/pq_generation.rs +++ b/diskann-disk/src/storage/quant/pq/pq_generation.rs @@ -111,18 +111,19 @@ where .read_existing_pivot_metadata(context.storage_provider)?; let num_chunks = context.num_chunks; - let table = context.pq_storage.load_pivots( - context.pq_storage.get_pivot_data_path(), - Some(num_chunks), - context.storage_provider, - )?; + let table = context.pq_storage.load_pivots(context.storage_provider)?; - if table.ncenters() != context.num_centers || table.dim() != full_dim { + if table.nchunks() != num_chunks + || table.ncenters() != context.num_centers + || table.dim() != full_dim + { return Err(diskann_error!( ErrorKind::PQError, - "PQ pivot table mismatch: file has {} centers in {} dimensions but expected {} centers in {} dimensions.", + "PQ pivot table mismatch: file has {} chunks, {} centers in {} dimensions but expected {} chunks, {} centers in {} dimensions.", + table.nchunks(), table.ncenters(), table.dim(), + num_chunks, context.num_centers, full_dim )); diff --git a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs index d01261dcd..2235f418c 100644 --- a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs @@ -229,7 +229,7 @@ impl FastMemoryQuantVectorProviderAsync { /// Load `self` from a pivots file and data file. /// - /// The pivots file follows the format in [`storage::PQStorage::load_pq_pivots_bin`] and + /// The pivots file follows the format in [`storage::PQStorage::load_pivots`] and /// the compressed code is saved in a canonical `.bin` format. /// /// See also: [`storage::bin::load_from_bin`]. @@ -245,7 +245,15 @@ impl FastMemoryQuantVectorProviderAsync { // We can use that information to load the pivots, then finish the rest // of initialization. let pq_storage = storage::PQStorage::new(pivots, data, None); - let table = pq_storage.load_pq_pivots_bin(pivots, pq_bytes, provider)?; + let table = pq_storage.load_pivots(provider)?; + if table.nchunks() != pq_bytes { + return Err(ANNError::message(format!( + "PQ pivot table mismatch: file has {} chunks but expected {} chunks.", + table.nchunks(), + pq_bytes + ))); + } + let table: FixedChunkPQTable = table.into(); Ok(Self::new(metric, num_points, table)) }) } diff --git a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs index 391873612..b5bab1873 100644 --- a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs @@ -162,7 +162,7 @@ impl MemoryQuantVectorProviderAsync { /// Load `self` from a pivots file and data file. /// - /// The pivots file follows the format in [`storage::PQStorage::load_pq_pivots_bin`] and + /// The pivots file follows the format in [`storage::PQStorage::load_pivots`] and /// the compressed code is saved in a canonical `.bin` format. /// /// See also: [`storage::bin::load_from_bin`]. @@ -178,7 +178,15 @@ impl MemoryQuantVectorProviderAsync { // We can use that information to load the pivots, then finish the rest // of initialization. let pq_storage = storage::PQStorage::new(pivots, data, None); - let table = pq_storage.load_pq_pivots_bin(pivots, pq_bytes, provider)?; + let table = pq_storage.load_pivots(provider)?; + if table.nchunks() != pq_bytes { + return Err(ANNError::message(format!( + "PQ pivot table mismatch: file has {} chunks but expected {} chunks.", + table.nchunks(), + pq_bytes + ))); + } + let table: FixedChunkPQTable = table.into(); Ok(Self::new(metric, num_points, table)) }) } diff --git a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs index e12e1bc0a..50a1b96e9 100644 --- a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs +++ b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs @@ -144,11 +144,6 @@ impl FixedChunkPQTable { Ok(Self { table }) } - /// Wrap an already-validated basic PQ table. - pub fn from_basic_table(table: BasicTable) -> Self { - Self { table } - } - /// Get chunk number. pub fn get_num_chunks(&self) -> usize { self.table.nchunks() @@ -435,6 +430,12 @@ impl FixedChunkPQTable { } } +impl From for FixedChunkPQTable { + fn from(table: BasicTable) -> Self { + Self { table } + } +} + // This goes against Rust's Orphan rule, so we cannot implement it directly. // However, we can use a wrapper type to implement the conversion. // This is a workaround to allow the conversion from `product::TableCompressionError` to @@ -698,8 +699,9 @@ mod fixed_chunk_pq_table_test { fn load_test_pivots() -> FixedChunkPQTable { let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); PQStorage::new(PQ_PIVOTS_PATH, "", None) - .load_pq_pivots_bin(PQ_PIVOTS_PATH, 1, &storage_provider) + .load_pivots(&storage_provider) .unwrap() + .into() } #[test] diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index a95e43fa3..a57e45a1d 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -347,17 +347,18 @@ where return Err(ANNError::message("ERROR: PQ k-means pivot file not found.")); } else { (_, full_dim) = pq_storage.read_existing_pivot_metadata(storage_provider)?; - table = pq_storage.load_pivots( - pq_storage.get_pivot_data_path(), - Some(num_pq_chunks), - storage_provider, - )?; + table = pq_storage.load_pivots(storage_provider)?; - if table.ncenters() != num_centers || table.dim() != full_dim { + if table.nchunks() != num_pq_chunks + || table.ncenters() != num_centers + || table.dim() != full_dim + { return Err(ANNError::message(format!( - "PQ pivot table mismatch: file has {} centers in {} dimensions but expected {} centers in {} dimensions.", + "PQ pivot table mismatch: file has {} chunks, {} centers in {} dimensions but expected {} chunks, {} centers in {} dimensions.", + table.nchunks(), table.ncenters(), table.dim(), + num_pq_chunks, num_centers, full_dim ))); @@ -929,14 +930,13 @@ mod pq_test { // use membuf function to generate pq // use pivot data generated by original function - let (full_pivot_data, centroid, offsets) = pq_storage - .load_existing_pivot_data( - &num_pq_chunks, - &NUM_PQ_CENTROIDS, - &train_dim, - &storage_provider, - ) - .unwrap(); + let table = pq_storage.load_pivots(&storage_provider).unwrap(); + assert_eq!(table.nchunks(), num_pq_chunks); + assert_eq!(table.ncenters(), NUM_PQ_CENTROIDS); + assert_eq!(table.dim(), train_dim); + let pivot_data_view = table.view_pivots(); + let full_pivot_data = pivot_data_view.as_slice(); + let offsets = table.view_offsets(); let mut membuf_pq_data: Vec = vec![0; num_pq_chunks * num_train]; @@ -950,9 +950,9 @@ mod pq_test { .for_each_in_pool(pool.as_ref(), |(i, membuf_slice)| { generate_pq_data_from_pivots_from_membuf( &full_data_vector[train_dim * i..train_dim * (i + 1)], - &full_pivot_data, + full_pivot_data, NUM_PQ_CENTROIDS, - &offsets, + offsets.as_slice(), membuf_slice, ) .unwrap(); @@ -984,7 +984,8 @@ mod pq_test { let full_data = MatrixView::try_from(full_data_vector.as_slice(), num_train, train_dim).unwrap(); let pivot_view = - MatrixView::try_from(full_pivot_data.as_slice(), NUM_PQ_CENTROIDS, train_dim).unwrap(); + MatrixView::try_from(full_pivot_data, NUM_PQ_CENTROIDS, train_dim).unwrap(); + let centroid = vec![0.0; train_dim]; // Due to difference in numerical rounding, the results between the two APIs can // vary slightly. diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 7bc06adf0..4df1f7cb9 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -19,17 +19,12 @@ use tracing::info; use crate::{ model::{ - FixedChunkPQTable, NUM_PQ_CENTROIDS, + NUM_PQ_CENTROIDS, pq::{METADATA_SIZE, accum_row_inplace}, }, utils::{gen_random_slice, read_bin_from, write_bin_from}, }; -// Create types to make return values easier to understand -type FullPivotDataType = Vec; -type CentroidType = Vec; -type ChunkOffsetsType = Vec; - #[derive(Debug, Clone)] pub struct PQStorage { /// Pivot table path @@ -160,69 +155,24 @@ impl PQStorage { Ok(Metadata::read(reader)?.into_dims()) } - /// Load the raw pivot data, centroid, and chunk offsets from a pivot file. - /// - /// Unlike [`Self::load_pq_pivots_bin`], this method returns the centroid - /// separately without folding it into the pivot data. Callers that need the - /// effective (centroid-adjusted) pivots must apply the centroid themselves, - /// e.g. via [`accum_row_inplace`](crate::model::pq::accum_row_inplace). - /// - /// For files written without legacy centering (`centroid = None` in - /// [`Self::write_pivot_data`]), the returned centroid will be all zeros and - /// can safely be accumulated as a no-op. - pub fn load_existing_pivot_data( - &self, - num_pq_chunks: &usize, - num_centers: &usize, - dim: &usize, - storage_provider: &Storage, - ) -> ANNResult<(FullPivotDataType, CentroidType, ChunkOffsetsType)> - where - Storage: StorageReadProvider, - { - let (pivots, centroid, chunk_offsets) = self.read_pivot_file( - &self.pivot_data_path, - Some(*num_pq_chunks), - Some(*num_centers), - Some(*dim), - storage_provider, - )?; - let table = - Self::pivot_data_into_basic_table(&self.pivot_data_path, pivots, chunk_offsets)?; - let (pivots, chunk_offsets) = table.into_parts(); - - Ok(( - pivots.into_inner().into_vec(), - centroid.into_inner().into_vec(), - chunk_offsets.into_inner().into_vec(), - )) - } - /// Load the effective PQ pivot table from a pivot file. /// - /// If `expected_num_pq_chunks` is `None`, the chunk count is inferred from the - /// file. The loader verifies the pivot file layout and folds any stored legacy - /// centroid into the pivots. `BasicTable::new` validates the resulting table - /// invariants, including chunk-offset bounds and monotonicity. + /// The loader validates internal consistency: file layout, + /// centroid/pivot dimensions, offset monotonicity and bounds, and + /// `BasicTable` invariants. Callers are responsible for validating that the + /// resulting table is compatible with their build or search configuration. pub fn load_pivots( &self, - pq_pivots: &str, - expected_num_pq_chunks: Option, storage_provider: &Storage, ) -> ANNResult { - let (mut pivots, centroid, chunk_offsets) = self.read_pivot_file( - pq_pivots, - expected_num_pq_chunks, - None, - None, - storage_provider, - )?; + let (mut pivots, centroid, chunk_offsets) = + self.read_pivot_file(&self.pivot_data_path, None, None, storage_provider)?; if centroid.as_slice().iter().any(|c| *c != 0.0) { accum_row_inplace(pivots.as_mut_view(), centroid.as_slice()) } - Self::pivot_data_into_basic_table(pq_pivots, pivots, chunk_offsets) + Self::pivot_data_into_basic_table(&self.pivot_data_path, pivots, chunk_offsets) } /// Load the compressed pq dataset from file. @@ -260,25 +210,9 @@ impl PQStorage { Ok(data) } - /// Load pre-trained pivot table - pub fn load_pq_pivots_bin( - &self, - pq_pivots: &str, - num_pq_chunks: usize, - storage_provider: &Storage, - ) -> ANNResult { - let table = self.load_pivots( - pq_pivots, - (num_pq_chunks != 0).then_some(num_pq_chunks), - storage_provider, - )?; - Ok(FixedChunkPQTable::from_basic_table(table)) - } - fn read_pivot_file( &self, pq_pivots: &str, - expected_num_pq_chunks: Option, expected_num_centers: Option, expected_dim: Option, storage_provider: &Storage, @@ -356,15 +290,6 @@ impl PQStorage { } let chunk_offsets_m = read_bin_from::(&mut reader, file_offset_data[(2, 0)])?; - if let Some(num_pq_chunks) = expected_num_pq_chunks - && chunk_offsets_m.nrows() != num_pq_chunks + 1 - { - return Err(ANNError::message(format!( - "Error reading pq_pivots file at chunk offsets; file has nr={}, but expecting nr={}.", - chunk_offsets_m.nrows(), - num_pq_chunks + 1 - ))); - } if chunk_offsets_m.ncols() != 1 { return Err(ANNError::message(format!( "Error reading pq_pivots file at chunk offsets; file has nc={}, but expecting nc=1.", @@ -429,10 +354,6 @@ impl PQStorage { pub fn get_compressed_data_path(&self) -> &str { &self.compressed_data_path } - - pub fn get_pivot_data_path(&self) -> &str { - &self.pivot_data_path - } } #[cfg(test)] @@ -530,18 +451,15 @@ mod pq_storage_tests { fn load_pivot_data_test() { let storage_provider = VirtualStorageProvider::new_overlay(test_data_root()); let result = PQStorage::new(PQ_PIVOT_PATH, PQ_COMPRESSED_PATH, Some(DATA_FILE)); - let (pq_pivot_data, centroids, chunk_offsets) = result - .load_existing_pivot_data(&1, &256, &128, &storage_provider) - .unwrap(); + let table = result.load_pivots(&storage_provider).unwrap(); - assert_eq!(pq_pivot_data.len(), 256 * 128); - assert_eq!(centroids.len(), 128); - assert_eq!(chunk_offsets.len(), 2); + assert_eq!(table.view_pivots().as_slice().len(), 256 * 128); + assert_eq!(table.dim(), 128); + assert_eq!(table.nchunks(), 1); } - /// Write pivot data with `centroid = None`, read it back via - /// `load_existing_pivot_data`, and verify the pivots are unchanged and the - /// centroid is all zeros. + /// Write pivot data with `centroid = None`, read it back, and verify the + /// effective pivots are unchanged. #[test] fn write_read_roundtrip_no_centroid() { let storage_provider = VirtualStorageProvider::new_memory(); @@ -565,26 +483,11 @@ mod pq_storage_tests { ) .unwrap(); - let (loaded_pivots, loaded_centroid, loaded_offsets) = pq_storage - .load_existing_pivot_data(&num_pq_chunks, &num_centers, &dim, &storage_provider) - .unwrap(); - - assert_eq!( - loaded_pivots, pivots, - "pivots should survive the round-trip unchanged" - ); - assert!( - loaded_centroid.iter().all(|&c| c == 0.0), - "centroid should be all zeros when written with None" - ); - assert_eq!(loaded_offsets, chunk_offsets); - - // Check that `load_pq_pivots_bin` correctly loads the pivots. - let table = pq_storage - .load_pq_pivots_bin(pivot_path, num_pq_chunks, &storage_provider) - .unwrap(); + let table = pq_storage.load_pivots(&storage_provider).unwrap(); - assert_eq!(loaded_pivots, table.view_pivots().as_slice()); + assert_eq!(table.view_pivots().as_slice(), pivots); + assert_eq!(table.view_offsets().as_slice(), chunk_offsets); + assert_eq!(table.nchunks(), num_pq_chunks); } #[test] @@ -609,9 +512,7 @@ mod pq_storage_tests { ) .unwrap(); - let table = pq_storage - .load_pivots(pivot_path, None, &storage_provider) - .unwrap(); + let table = pq_storage.load_pivots(&storage_provider).unwrap(); assert_eq!(table.view_pivots().as_slice(), pivots); } @@ -627,34 +528,14 @@ mod pq_storage_tests { .unwrap(); let table = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pq_pivots_bin(pivot_path, 0, &storage_provider) + .load_pivots(&storage_provider) .unwrap(); assert_eq!(table.view_pivots().as_slice(), pivots); } #[test] - fn load_pivot_data_rejects_mismatched_shape() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/mismatched_shape_pivots.bin"; - - write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 4]); - let pq_storage = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None); - - assert!( - pq_storage - .load_existing_pivot_data(&2, &4, &4, &storage_provider) - .is_err() - ); - assert!( - pq_storage - .load_existing_pivot_data(&2, &3, &5, &storage_provider) - .is_err() - ); - } - - #[test] - fn load_pivot_data_rejects_invalid_centroid_and_chunk_count() { + fn load_pq_pivots_rejects_invalid_centroid() { let storage_provider = VirtualStorageProvider::new_memory(); let wrong_centroid_path = "/wrong_centroid_pivots.bin"; @@ -668,15 +549,7 @@ mod pq_storage_tests { ); assert!( PQStorage::new(wrong_centroid_path, PQ_COMPRESSED_PATH, None) - .load_existing_pivot_data(&2, &3, &4, &storage_provider) - .is_err() - ); - - let wrong_count_path = "/wrong_chunk_count_pivots.bin"; - write_test_pivots(&storage_provider, wrong_count_path, 3, 4, None, &[0, 4]); - assert!( - PQStorage::new(wrong_count_path, PQ_COMPRESSED_PATH, None) - .load_existing_pivot_data(&2, &3, &4, &storage_provider) + .load_pivots(&storage_provider) .is_err() ); } @@ -689,7 +562,7 @@ mod pq_storage_tests { write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_existing_pivot_data(&1, &3, &4, &storage_provider) + .load_pivots(&storage_provider) .unwrap_err(); assert!(err.to_string().contains("offsets must begin at 0")); @@ -703,7 +576,7 @@ mod pq_storage_tests { write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 3]); let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_existing_pivot_data(&2, &3, &4, &storage_provider) + .load_pivots(&storage_provider) .unwrap_err(); assert!(err.to_string().contains("offsets expect 3")); @@ -717,7 +590,7 @@ mod pq_storage_tests { write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(pivot_path, None, &storage_provider) + .load_pivots(&storage_provider) .unwrap_err(); assert!(err.to_string().contains("offsets must begin at 0")); @@ -765,7 +638,7 @@ mod pq_storage_tests { } let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(pivot_path, None, &storage_provider) + .load_pivots(&storage_provider) .unwrap_err(); let message = err.to_string(); @@ -789,7 +662,7 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(pivot_path, None, &storage_provider) + .load_pivots(&storage_provider) .is_err() ); } @@ -807,19 +680,15 @@ mod pq_storage_tests { assert!( PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(pivot_path, None, &storage_provider) + .load_pivots(&storage_provider) .is_err() ); } /// Write pivot data with a non-zero centroid, read it back, and verify that - /// folding the centroid via `accum_row_inplace` produces the expected - /// adjusted pivots. + /// loading folds the centroid into the effective pivots. #[test] fn write_read_roundtrip_with_legacy_centroid() { - use crate::model::pq::accum_row_inplace; - use diskann_utils::views::MutMatrixView; - let storage_provider = VirtualStorageProvider::new_memory(); let pivot_path = "/roundtrip_legacy_centroid_pivots.bin"; @@ -842,42 +711,22 @@ mod pq_storage_tests { ) .unwrap(); - let (mut loaded_pivots, loaded_centroid, loaded_offsets) = pq_storage - .load_existing_pivot_data(&num_pq_chunks, &num_centers, &dim, &storage_provider) - .unwrap(); - - assert_eq!( - loaded_pivots, pivots, - "raw pivots should match what was written" - ); - assert_eq!( - loaded_centroid, centroid, - "centroid should round-trip exactly" - ); - assert_eq!(loaded_offsets, chunk_offsets); - - // Fold the centroid into the pivots — this is what production callers do. - let mut pivot_mat = - MutMatrixView::try_from(loaded_pivots.as_mut_slice(), num_centers, dim).unwrap(); - accum_row_inplace(pivot_mat.as_mut_view(), &loaded_centroid); + let table = pq_storage.load_pivots(&storage_provider).unwrap(); + let pivot_view = table.view_pivots(); + let loaded_pivots = pivot_view.as_slice(); // Each pivot row should have the centroid added element-wise. - for (idx, (pivot, &orig)) in loaded_pivots.iter().zip(pivots.iter()).enumerate() { + for (idx, (&pivot, &orig)) in loaded_pivots.iter().zip(pivots.iter()).enumerate() { let d = idx % dim; let expected = orig + centroid[d]; assert_eq!( - *pivot, expected, + pivot, expected, "pivot[{}]: expected {expected}, got {pivot}", idx ); } - - // Check that `load_pq_pivots_bin` correctly does the centroid folding. - let table = pq_storage - .load_pq_pivots_bin(pivot_path, num_pq_chunks, &storage_provider) - .unwrap(); - - assert_eq!(loaded_pivots, table.view_pivots().as_slice()); + assert_eq!(table.view_offsets().as_slice(), chunk_offsets); + assert_eq!(table.nchunks(), num_pq_chunks); } #[test] diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index 822940645..849a2ba07 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -5,7 +5,7 @@ use crate::traits::CompressInto; use crate::views::{ChunkOffsetsBase, ChunkOffsetsView}; -use diskann_utils::views::{DenseData, MatrixBase, MatrixView, MutDenseData, MutMatrixView}; +use diskann_utils::views::{DenseData, MatrixBase, MatrixView}; use diskann_vector::{PureDistanceFunction, distance::SquaredL2}; use thiserror::Error; @@ -86,14 +86,6 @@ where self.pivots.as_view() } - /// Return a mutable view over the pivot table. - pub fn view_pivots_mut(&mut self) -> MutMatrixView<'_, f32> - where - T: MutDenseData, - { - self.pivots.as_mut_view() - } - /// Return a view over the schema offsets. pub fn view_offsets(&self) -> ChunkOffsetsView<'_> { self.offsets.as_view() From 729c6f0faa5b014d88650886492609bcf0342f05 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Mon, 3 Aug 2026 22:09:06 +0800 Subject: [PATCH 11/13] Remove redundant PQ loading helpers --- .../fast_memory_quant_vector_provider.rs | 3 +-- .../async_/memory_quant_vector_provider.rs | 3 +-- diskann-providers/src/storage/pq_storage.rs | 26 ++----------------- .../src/product/tables/basic.rs | 5 ---- diskann-quantization/src/views.rs | 5 ---- 5 files changed, 4 insertions(+), 38 deletions(-) diff --git a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs index 2235f418c..81afa59a8 100644 --- a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs @@ -253,8 +253,7 @@ impl FastMemoryQuantVectorProviderAsync { pq_bytes ))); } - let table: FixedChunkPQTable = table.into(); - Ok(Self::new(metric, num_points, table)) + Ok(Self::new(metric, num_points, table.into())) }) } diff --git a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs index b5bab1873..79255bffe 100644 --- a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs @@ -186,8 +186,7 @@ impl MemoryQuantVectorProviderAsync { pq_bytes ))); } - let table: FixedChunkPQTable = table.into(); - Ok(Self::new(metric, num_points, table)) + Ok(Self::new(metric, num_points, table.into())) }) } diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 4df1f7cb9..8672343c6 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -166,7 +166,7 @@ impl PQStorage { storage_provider: &Storage, ) -> ANNResult { let (mut pivots, centroid, chunk_offsets) = - self.read_pivot_file(&self.pivot_data_path, None, None, storage_provider)?; + self.read_pivot_file(&self.pivot_data_path, storage_provider)?; if centroid.as_slice().iter().any(|c| *c != 0.0) { accum_row_inplace(pivots.as_mut_view(), centroid.as_slice()) @@ -213,8 +213,6 @@ impl PQStorage { fn read_pivot_file( &self, pq_pivots: &str, - expected_num_centers: Option, - expected_dim: Option, storage_provider: &Storage, ) -> ANNResult<(Matrix, Matrix, Matrix)> { if !storage_provider.exists(pq_pivots) { @@ -248,16 +246,7 @@ impl PQStorage { info!(" Offset data: {:?}", file_offset_data.as_slice()); let pivots = read_bin_from::(&mut reader, file_offset_data[(0, 0)])?; - if let Some(num_centers) = expected_num_centers { - if pivots.nrows() != num_centers { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", - pq_pivots, - pivots.nrows(), - num_centers - ))); - } - } else if pivots.nrows() > NUM_PQ_CENTROIDS { + if pivots.nrows() > NUM_PQ_CENTROIDS { return Err(ANNError::message(format!( "Error reading pq_pivots file {}. file_num_centers = {}, but expecting {} centers.", pq_pivots, @@ -266,17 +255,6 @@ impl PQStorage { ))); } - if let Some(dim) = expected_dim - && pivots.ncols() != dim - { - return Err(ANNError::message(format!( - "Error reading pq_pivots file {}. file_dim = {} but expecting {} dimensions.", - pq_pivots, - pivots.ncols(), - dim - ))); - } - let centroid = read_bin_from::(&mut reader, file_offset_data[(1, 0)])?; if centroid.nrows() != pivots.ncols() || centroid.ncols() != 1 { return Err(ANNError::message(format!( diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index 849a2ba07..9469fa58d 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -105,11 +105,6 @@ where pub fn dim(&self) -> usize { self.pivots.ncols() } - - /// Consume this table and return the underlying pivots and offsets. - pub fn into_parts(self) -> (MatrixBase, ChunkOffsetsBase) { - (self.pivots, self.offsets) - } } #[derive(Error, Debug)] diff --git a/diskann-quantization/src/views.rs b/diskann-quantization/src/views.rs index 03ed01223..04c4a0953 100644 --- a/diskann-quantization/src/views.rs +++ b/diskann-quantization/src/views.rs @@ -211,11 +211,6 @@ where pub fn as_slice(&self) -> &[usize] { self.offsets.as_slice() } - - /// Consume the offsets, returning the inner representation. - pub fn into_inner(self) -> T { - self.offsets - } } pub type ChunkOffsetsView<'a> = ChunkOffsetsBase<&'a [usize]>; From 3a162c7a96960cdb2aaa3f198d67c9f1c464c309 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Tue, 4 Aug 2026 10:14:21 +0800 Subject: [PATCH 12/13] Simplify PQ membuf compression tests --- diskann-providers/src/model/mod.rs | 4 +- diskann-providers/src/model/pq/mod.rs | 5 +- .../src/model/pq/pq_construction.rs | 131 +++--------------- diskann-providers/src/storage/pq_storage.rs | 86 ------------ 4 files changed, 27 insertions(+), 199 deletions(-) diff --git a/diskann-providers/src/model/mod.rs b/diskann-providers/src/model/mod.rs index 69631c2f9..eaf9c1766 100644 --- a/diskann-providers/src/model/mod.rs +++ b/diskann-providers/src/model/mod.rs @@ -12,8 +12,8 @@ pub mod pq; pub use pq::{ FixedChunkPQTable, GeneratePivotArguments, MAX_PQ_TRAINING_SET_SIZE, NUM_KMEANS_REPS_PQ, NUM_PQ_CENTROIDS, compute_pq_distance, compute_pq_distance_for_pq_coordinates, distance, - generate_pq_data_from_pivots_from_membuf, generate_pq_data_from_pivots_from_membuf_batch, - generate_pq_pivots, generate_pq_pivots_from_membuf, + generate_pq_data_from_pivots_from_membuf_batch, generate_pq_pivots, + generate_pq_pivots_from_membuf, }; pub mod statistics; diff --git a/diskann-providers/src/model/pq/mod.rs b/diskann-providers/src/model/pq/mod.rs index e3c6f21a0..5f4b2deba 100644 --- a/diskann-providers/src/model/pq/mod.rs +++ b/diskann-providers/src/model/pq/mod.rs @@ -11,9 +11,8 @@ pub use fixed_chunk_pq_table::{ mod pq_construction; pub use pq_construction::{ MAX_PQ_TRAINING_SET_SIZE, NUM_KMEANS_REPS_PQ, NUM_PQ_CENTROIDS, accum_row_inplace, - generate_pq_data_from_pivots, generate_pq_data_from_pivots_from_membuf, - generate_pq_data_from_pivots_from_membuf_batch, generate_pq_pivots, - generate_pq_pivots_from_membuf, move_train_data_by_centroid, + generate_pq_data_from_pivots, generate_pq_data_from_pivots_from_membuf_batch, + generate_pq_pivots, generate_pq_pivots_from_membuf, move_train_data_by_centroid, }; /// all metadata of individual sub-component files is written in first 4KB for unified files diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index a57e45a1d..aca742516 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -448,76 +448,12 @@ where Ok(()) } -/// Compute the PQ codes for a single vector argument. +/// Compute PQ codes for a batch of vectors using in-memory pivots. /// -/// Given training data in train_data of dimensions `dim` and -/// PQ pivots computed earlier, partition the co-ordinates into -/// `num_pq_chunks`, and find the closest pivots for each point in each chunk. -/// This API doesn't involve reading/writing to disk and is used for in-memory. -/// -/// If `centroid` is `Some(_)` subtract the centroid from each point before finding the -/// closest pivots. -/// -/// Output `pq_out` which must be pre-allocated and will be used to determine -/// `num_pq_chunks`. -/// -/// # Arguments -/// * `vector_data` - A single vector to be encoded. -/// * `pivot_data` - A logical 2-dimensional array containing the PQ pivots in row-major -/// order. -/// * `num_pivots` - The size of the first dimension of the `pivot_data` matrix. -/// * `centroid` - An optional centroid to use for zero centering `vector_data`. -/// -/// If `Some(_)`, then `vector_data` will be transformed by subtracting each component by -/// its corresponding entry in `centroid`. -/// -/// If `None`, then no centering will take place. -/// * `offsets` - A prefix-sum style encoding of the start and stop positions of each -/// chunk in `pivot_data`. -/// * `pq_out` - Output buffer for the PQ codes. -/// -/// # Returns -/// An `ANNResult<()>` indicating success or failure. -pub fn generate_pq_data_from_pivots_from_membuf>( - vector_data: &[T], - pivot_data: &[f32], - num_pivots: usize, - offsets: &[usize], - pq_out: &mut [u8], -) -> ANNResult<()> { - // Number of dimensions in the vector to encode. - let dim = vector_data.len(); - - // Create a `BasicTableView` of the pivots. - // - // This does not allocate memory, but does validate the following invariants: - // * `pivot_data.len() == num_pivots * dim`. - // * `offsets` begins at zero, ends at `dim`, and is monotonic. - let table = BasicTableView::new( - MatrixView::try_from(pivot_data, num_pivots, dim).bridge_err()?, - diskann_quantization::views::ChunkOffsetsView::new(offsets).bridge_err()?, - ) - .map_err(ANNError::new)?; - - let data = vector_data - .iter() - .map(|x| (*x).into()) - .collect::>(); - - table - .compress_into(data.as_slice(), pq_out) - .map_err(ANNError::new) -} - -/// Legacy compatibility function for providing an batch data generation. -/// -/// Compute the PQ codes for a single vector argument. -/// -/// Given training data in train_data of dimensions `dim` and -/// PQ pivots computed earlier, partition the co-ordinates into -/// `num_pq_chunks`, and find the closest pivots for each point in each chunk. -/// This API doesn't involve reading/writing to disk and is used for in-memory. -pub fn generate_pq_data_from_pivots_from_membuf_batch>( +/// Given vector data with dimensions from `parameters` and PQ pivots computed +/// earlier, partition the coordinates into `num_pq_chunks` and find the closest +/// pivot for each vector chunk. This API does not read or write storage. +pub fn generate_pq_data_from_pivots_from_membuf_batch( parameters: &GeneratePivotArguments, vector_data: &[T], pivot_data: &[f32], @@ -528,7 +464,7 @@ pub fn generate_pq_data_from_pivots_from_membuf_batch // Perform minimal error checking at this level, mainly on the sizes of `vector_data` // and `pq_out`. // - // More dimentionality checking is deferred to the inner function. + // More dimensionality checking is deferred to the table construction. let num_train = parameters.num_train(); let num_pq_chunks = parameters.num_pq_chunks(); let dim = parameters.dim(); @@ -552,10 +488,8 @@ pub fn generate_pq_data_from_pivots_from_membuf_batch .par_chunks_mut(num_pq_chunks) .zip(vector_data.par_chunks(dim)) .try_for_each_in_pool(pool, |(pq_slice, vector)| { - let data = vector.iter().map(|x| (*x).into()).collect::>(); - table - .compress_into(data.as_slice(), pq_slice) - .map_err(ANNError::new) + let data = T::as_f32(vector).map_err(ANNError::new)?; + table.compress_into(&*data, pq_slice).map_err(ANNError::new) }) } @@ -849,26 +783,16 @@ mod pq_test { ) .unwrap(); + let table = FixedChunkPQTable::new(dim, pivot_data.into(), offsets.into()).unwrap(); let mut pq: Vec = vec![0; num_pq_chunks]; for i in 0..num_train { - generate_pq_data_from_pivots_from_membuf( - &train_data[dim * i..dim * (i + 1)], - &pivot_data, - num_centers, - &offsets, - &mut pq, - ) - .unwrap(); + table + .compress_into(&train_data[dim * i..dim * (i + 1)], &mut pq) + .unwrap(); } - assert!( - !offsets.contains(&usize::MAX), - "offsets contains max value!" - ); - assert!( - !pivot_data.contains(&f32::MAX), - "pivot_data contains max value!" - ); + assert!(!table.get_chunk_offsets().contains(&usize::MAX)); + assert!(!table.get_pq_table().contains(&f32::MAX)); } #[rstest] @@ -934,10 +858,6 @@ mod pq_test { assert_eq!(table.nchunks(), num_pq_chunks); assert_eq!(table.ncenters(), NUM_PQ_CENTROIDS); assert_eq!(table.dim(), train_dim); - let pivot_data_view = table.view_pivots(); - let full_pivot_data = pivot_data_view.as_slice(); - let offsets = table.view_offsets(); - let mut membuf_pq_data: Vec = vec![0; num_pq_chunks * num_train]; // `from_membuf` switched to an implementation optimized for a single vector. @@ -948,14 +868,12 @@ mod pq_test { .par_chunks_mut(num_pq_chunks) .enumerate() .for_each_in_pool(pool.as_ref(), |(i, membuf_slice)| { - generate_pq_data_from_pivots_from_membuf( - &full_data_vector[train_dim * i..train_dim * (i + 1)], - full_pivot_data, - NUM_PQ_CENTROIDS, - offsets.as_slice(), - membuf_slice, - ) - .unwrap(); + table + .compress_into( + &full_data_vector[train_dim * i..train_dim * (i + 1)], + membuf_slice, + ) + .unwrap(); }); // use pq generated by original function as the gt @@ -983,8 +901,7 @@ mod pq_test { let offset_view = chunk_offsets.as_view(); let full_data = MatrixView::try_from(full_data_vector.as_slice(), num_train, train_dim).unwrap(); - let pivot_view = - MatrixView::try_from(full_pivot_data, NUM_PQ_CENTROIDS, train_dim).unwrap(); + let pivot_view = table.view_pivots(); let centroid = vec![0.0; train_dim]; // Due to difference in numerical rounding, the results between the two APIs can @@ -1106,13 +1023,11 @@ mod pq_test { ); assert!(result.is_ok()); + let table = FixedChunkPQTable::new(dim, full_pivot_data.into(), offsets.into()).unwrap(); let mut membuf_pq_data: Vec = vec![0; num_pq_chunks]; for i in 0..npts { - let result = generate_pq_data_from_pivots_from_membuf( + let result = table.compress_into( &full_data_vector[(dim * i)..(dim * (i + 1))], - &full_pivot_data, - NUM_PQ_CENTROIDS, - &offsets, &mut membuf_pq_data, ); assert!(result.is_ok()); diff --git a/diskann-providers/src/storage/pq_storage.rs b/diskann-providers/src/storage/pq_storage.rs index 8672343c6..55fb43c29 100644 --- a/diskann-providers/src/storage/pq_storage.rs +++ b/diskann-providers/src/storage/pq_storage.rs @@ -468,50 +468,6 @@ mod pq_storage_tests { assert_eq!(table.nchunks(), num_pq_chunks); } - #[test] - fn load_pq_pivots_infer_chunks_loads_without_expected_count() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/infer_chunk_count_pivots.bin"; - - let num_centers = 3; - let dim = 4; - let pivots: Vec = (0..num_centers * dim).map(|i| i as f32).collect(); - let chunk_offsets = vec![0, 2, dim]; - - let pq_storage = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None); - pq_storage - .write_pivot_data( - &pivots, - None, - &chunk_offsets, - num_centers, - dim, - &storage_provider, - ) - .unwrap(); - - let table = pq_storage.load_pivots(&storage_provider).unwrap(); - - assert_eq!(table.view_pivots().as_slice(), pivots); - } - - #[test] - fn load_pq_pivots_zero_chunk_count_infers_from_file() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/zero_chunk_count_pivots.bin"; - let pivots: Vec = (0..12).map(|i| i as f32).collect(); - - PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .write_pivot_data(&pivots, None, &[0, 2, 4], 3, 4, &storage_provider) - .unwrap(); - - let table = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(&storage_provider) - .unwrap(); - - assert_eq!(table.view_pivots().as_slice(), pivots); - } - #[test] fn load_pq_pivots_rejects_invalid_centroid() { let storage_provider = VirtualStorageProvider::new_memory(); @@ -532,48 +488,6 @@ mod pq_storage_tests { ); } - #[test] - fn load_pivot_data_rejects_invalid_chunk_bounds() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/legacy_chunk_bounds_pivots.bin"; - - write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); - - let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(&storage_provider) - .unwrap_err(); - - assert!(err.to_string().contains("offsets must begin at 0")); - } - - #[test] - fn load_pivot_data_rejects_chunk_offsets_dim_mismatch() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/chunk_offsets_dim_mismatch_pivots.bin"; - - write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[0, 2, 3]); - - let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(&storage_provider) - .unwrap_err(); - - assert!(err.to_string().contains("offsets expect 3")); - } - - #[test] - fn load_pq_pivots_infer_rejects_invalid_chunk_bounds() { - let storage_provider = VirtualStorageProvider::new_memory(); - let pivot_path = "/infer_wrong_chunk_bounds_pivots.bin"; - - write_test_pivots(&storage_provider, pivot_path, 3, 4, None, &[1, 4]); - - let err = PQStorage::new(pivot_path, PQ_COMPRESSED_PATH, None) - .load_pivots(&storage_provider) - .unwrap_err(); - - assert!(err.to_string().contains("offsets must begin at 0")); - } - #[test] fn load_pq_pivots_reports_chunk_offset_column_mismatch() { let storage_provider = VirtualStorageProvider::new_memory(); From ce57c16392a76c71676214217db7d14929b8c773 Mon Sep 17 00:00:00 2001 From: "Xinyu Wen (from Dev Box)" Date: Tue, 4 Aug 2026 11:03:02 +0800 Subject: [PATCH 13/13] Validate fixed PQ table center count --- diskann-disk/src/storage/disk_index_reader.rs | 3 +- .../fast_memory_quant_vector_provider.rs | 2 +- .../async_/memory_quant_vector_provider.rs | 2 +- .../src/model/pq/fixed_chunk_pq_table.rs | 36 ++++++++++++++++--- 4 files changed, 36 insertions(+), 7 deletions(-) diff --git a/diskann-disk/src/storage/disk_index_reader.rs b/diskann-disk/src/storage/disk_index_reader.rs index e776bf056..909954075 100644 --- a/diskann-disk/src/storage/disk_index_reader.rs +++ b/diskann-disk/src/storage/disk_index_reader.rs @@ -31,7 +31,8 @@ impl DiskIndexReader { storage_provider: &Storage, ) -> ANNResult { let pq_storage = PQStorage::new(&pq_pivot_path, &pq_compressed_data_path, None); - let pq_pivot_table: FixedChunkPQTable = pq_storage.load_pivots(storage_provider)?.into(); + let pq_pivot_table = + FixedChunkPQTable::try_from(pq_storage.load_pivots(storage_provider)?)?; // Auto-detect number of points from compressed PQ file metadata let metadata = load_metadata_from_file(storage_provider, &pq_compressed_data_path)?; diff --git a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs index 81afa59a8..e599855ac 100644 --- a/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/fast_memory_quant_vector_provider.rs @@ -253,7 +253,7 @@ impl FastMemoryQuantVectorProviderAsync { pq_bytes ))); } - Ok(Self::new(metric, num_points, table.into())) + Ok(Self::new(metric, num_points, table.try_into()?)) }) } diff --git a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs index 79255bffe..237fd33e3 100644 --- a/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs +++ b/diskann-providers/src/model/graph/provider/async_/memory_quant_vector_provider.rs @@ -186,7 +186,7 @@ impl MemoryQuantVectorProviderAsync { pq_bytes ))); } - Ok(Self::new(metric, num_points, table.into())) + Ok(Self::new(metric, num_points, table.try_into()?)) }) } diff --git a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs index 50a1b96e9..9ca23ac5b 100644 --- a/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs +++ b/diskann-providers/src/model/pq/fixed_chunk_pq_table.rs @@ -430,9 +430,19 @@ impl FixedChunkPQTable { } } -impl From for FixedChunkPQTable { - fn from(table: BasicTable) -> Self { - Self { table } +impl TryFrom for FixedChunkPQTable { + type Error = ANNError; + + fn try_from(table: BasicTable) -> Result { + if table.ncenters() > NUM_PQ_CENTROIDS { + return Err(ANNError::message(format!( + "PQ pivot table mismatch: file has {} centers but supports at most {} centers.", + table.ncenters(), + NUM_PQ_CENTROIDS + ))); + } + + Ok(Self { table }) } } @@ -701,7 +711,8 @@ mod fixed_chunk_pq_table_test { PQStorage::new(PQ_PIVOTS_PATH, "", None) .load_pivots(&storage_provider) .unwrap() - .into() + .try_into() + .unwrap() } #[test] @@ -762,6 +773,23 @@ mod fixed_chunk_pq_table_test { } } + #[test] + fn conversion_rejects_too_many_centers() { + let dim = 5; + let table = BasicTable::new( + MatrixBase::try_from( + vec![0.0; dim * (NUM_PQ_CENTROIDS + 1)].into_boxed_slice(), + NUM_PQ_CENTROIDS + 1, + dim, + ) + .unwrap(), + ChunkOffsetsBase::new(vec![0, 2, dim].into_boxed_slice()).unwrap(), + ) + .unwrap(); + + assert!(FixedChunkPQTable::try_from(table).is_err()); + } + #[test] fn test_compute_pq_distance() { let num_pq_chunks = 17;