From 429c1a503489974f9b6b724d4bf2a2f0b5f6bae6 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 6 Oct 2026 16:48:03 -0700 Subject: [PATCH 1/3] PQ training returns a BasicTable. --- diskann-benchmark/src/exhaustive/product.rs | 12 +- .../src/model/pq/pq_construction.rs | 19 ++- .../src/product/tables/basic.rs | 65 +++++++- diskann-quantization/src/product/train.rs | 146 ++++++------------ diskann-quantization/src/views.rs | 38 ++++- 5 files changed, 148 insertions(+), 132 deletions(-) diff --git a/diskann-benchmark/src/exhaustive/product.rs b/diskann-benchmark/src/exhaustive/product.rs index 200d76ed0c..3b206e28d6 100644 --- a/diskann-benchmark/src/exhaustive/product.rs +++ b/diskann-benchmark/src/exhaustive/product.rs @@ -99,7 +99,7 @@ mod imp { let offsets = diskann_quantization::views::ChunkOffsets::partition(dim, input.num_pq_chunks)?; - let base = { + let table = { let threadpool = rayon::ThreadPoolBuilder::new() .num_threads(input.compression_threads.get()) .build()?; @@ -114,16 +114,6 @@ mod imp { })? }; - // TODO: Training should return a `BasicTable` directly. - let table = tables::BasicTable::new( - rowmajor::Owned::try_from_data( - base.flatten().into(), - input.num_pq_centers.get(), - data.ncols(), - )?, - offsets, - )?; - let training_time: MicroSeconds = start.elapsed().into(); // Compressing diff --git a/diskann-providers/src/model/pq/pq_construction.rs b/diskann-providers/src/model/pq/pq_construction.rs index f6202810e0..74bcb0befd 100644 --- a/diskann-providers/src/model/pq/pq_construction.rs +++ b/diskann-providers/src/model/pq/pq_construction.rs @@ -116,7 +116,7 @@ where parameters.max_k_means_reps(), ); - let full_pivot_data = pool.install(|| -> Result, ANNError> { + let basic_table = pool.install(|| -> Result<_, ANNError> { let result = trainer .train( rowmajor::Ref::try_from_data(train_data, parameters.num_train(), parameters.dim()) @@ -126,13 +126,12 @@ where &random_provider, &diskann_quantization::cancel::DontCancel, ) - .map_err(ANNError::new)? - .flatten(); + .map_err(ANNError::new)?; Ok(result) })?; pq_storage.write_pivot_data( - &full_pivot_data, + basic_table.view_pivots().as_slice(), centroid.as_deref(), chunk_offsets.as_slice(), parameters.num_centers(), @@ -200,7 +199,7 @@ pub fn generate_pq_pivots_from_membuf>( ); let rng_builder = create_rnd_provider_from_seed(rand::distr::StandardUniform {}.sample(rng)); - let trained = pool.install(|| -> Result, ANNError> { + let basic_table = pool.install(|| -> Result<_, ANNError> { // SAFETY: The pointer for `cancellation_token` is valid for this local lifetime, // and we do not otherwise access `cancellation_token`. // @@ -216,7 +215,7 @@ pub fn generate_pq_pivots_from_membuf>( let atomic_bool: &AtomicBool = unsafe { AtomicBool::from_ptr(cancellation_token) }; let cancelation = diskann_quantization::cancel::AtomicCancelation::new(atomic_bool); - let result = trainer + let table = trainer .train( rowmajor::Ref::try_from_data( train_data.as_slice(), @@ -229,12 +228,12 @@ pub fn generate_pq_pivots_from_membuf>( &rng_builder, &cancelation, ) - .map_err(ANNError::new)? - .flatten(); - Ok(result) + .map_err(ANNError::new)?; + + Ok(table) })?; - full_pivot_data.copy_from_slice(&trained); + full_pivot_data.copy_from_slice(basic_table.view_pivots().as_slice()); Ok(()) } diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index f8db42bd84..bc0312f846 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -5,9 +5,12 @@ use crate::traits::CompressInto; use crate::views::{ChunkOffsetsBase, ChunkOffsetsView}; -use diskann_utils::views::{ - DenseData, - rowmajor::{self, Matrix}, +use diskann_utils::{ + strided::Strided, + views::{ + DenseData, + rowmajor::{self, Matrix}, + }, }; use diskann_vector::{PureDistanceFunction, distance::SquaredL2}; use thiserror::Error; @@ -113,6 +116,27 @@ where offsets: self.offsets.as_view(), } } + + /// Return a [`Strided`] for the raw pivots of the requested chunk. + /// + /// Returns `None` if `chunk >= self.nchunks`. + pub fn pivots_for(&self, chunk: usize) -> Option> { + let range = self.offsets.get(chunk)?; + + #[expect( + clippy::expect_used, + reason = "the BasicTable's invariants mean this panic should be unreachable" + )] + Some( + Strided::try_from_data( + &self.pivots.as_slice()[range.start..], + self.pivots.nrows(), + range.len(), + self.pivots.ncols(), + ) + .expect("BasicTable asserts that this layout is valid"), + ) + } } #[derive(Error, Debug)] @@ -304,6 +328,41 @@ mod tests { let (pivots, offsets) = create_pivot_tables(schema.to_owned(), num_centers); let table = BasicTable::new(pivots, offsets).unwrap(); + + // Check that `pivots_for` works as expected with repsect to the documented + // table configuration for `create_pivot_tables`. + for chunk in 0..schema.len() { + let strided = table.pivots_for(chunk).unwrap(); + assert_eq!(strided.nrows(), num_centers); + assert_eq!(strided.ncols(), schema.at(chunk).len()); + + for (center, row) in strided.rows().enumerate() { + let base = ((center + chunk) % num_centers) as f32; + row.iter().enumerate().for_each(|(dim, b)| { + let offset = if dim.is_multiple_of(2) { 0.25 } else { -0.25 }; + + if dim.is_multiple_of(2) { + assert_eq!( + *b, + base + offset, + "failed: chunk {} of {}, center {} of {}, dim {} of {}", + chunk, + schema.len(), + center, + num_centers, + dim, + row.len(), + ); + } + }) + } + } + + assert!( + table.pivots_for(schema.len()).is_none(), + "out of bounds access should return `None`", + ); + let (data, expected) = create_dataset(schema, num_centers, num_data, &mut rng); let mut output = vec![0; schema.len()]; diff --git a/diskann-quantization/src/product/train.rs b/diskann-quantization/src/product/train.rs index 037e2ebddc..364ed764f7 100644 --- a/diskann-quantization/src/product/train.rs +++ b/diskann-quantization/src/product/train.rs @@ -19,6 +19,7 @@ use crate::{ algorithms::kmeans::{self, common::square_norm}, cancel::Cancelation, multi_vector::BlockTransposed, + product::tables::BasicTable, random::{BoxedRngBuilder, RngBuilder}, }; @@ -39,45 +40,6 @@ impl LightPQTrainingParameters { } } -#[derive(Debug)] -pub struct SimplePivots { - dim: usize, - ncenters: usize, - pivots: Vec>, -} - -fn flatten( - pivots: &[rowmajor::Owned], - ncenters: usize, - dim: usize, -) -> rowmajor::Owned { - let mut flattened = rowmajor::Owned::from_element(ncenters, dim, T::default()); - let mut col_start = 0; - for matrix in pivots { - assert_eq!(matrix.nrows(), flattened.nrows()); - for (row_index, row) in matrix.rows().enumerate() { - let dst = &mut flattened.row_mut(row_index)[col_start..col_start + row.len()]; - dst.copy_from_slice(row); - } - col_start += matrix.ncols(); - } - flattened -} - -impl SimplePivots { - /// Return the selected pivots for each chunk. - pub fn pivots(&self) -> &[rowmajor::Owned] { - &self.pivots - } - - /// Concatenate the individual pivots into a dense representation. - pub fn flatten(&self) -> Vec { - flatten(self.pivots(), self.ncenters, self.dim) - .into_inner() - .into() - } -} - pub trait TrainQuantizer { type Quantizer; type Error: std::error::Error; @@ -96,7 +58,7 @@ pub trait TrainQuantizer { } impl TrainQuantizer for LightPQTrainingParameters { - type Quantizer = SimplePivots; + type Quantizer = BasicTable; type Error = PQTrainingError; /// Perform product quantization training on the provided training set and return a @@ -135,7 +97,7 @@ impl TrainQuantizer for LightPQTrainingParameters { parallelism: Parallelism, rng_builder: &(dyn BoxedRngBuilder + Sync), cancelation: &(dyn Cancelation + Sync), - ) -> Result { + ) -> Result { // Make sure we're provided sane values for our schema. assert_eq!(data.ncols(), schema.dim()); @@ -225,20 +187,32 @@ impl TrainQuantizer for LightPQTrainingParameters { Ok(centers) }; - let pivots: Result, _> = match parallelism { + let pivots: Result>, _> = match parallelism { Parallelism::Sequential => (0..schema.len()).map(thunk).collect(), #[cfg(feature = "rayon")] Parallelism::Rayon => (0..schema.len()).into_par_iter().map(thunk).collect(), }; - let dim = data.ncols(); - let ncenters = trainer.ncenters; - Ok(SimplePivots { - dim, - ncenters, - pivots: pivots?, - }) + let pivots = pivots?; + let mut packed = rowmajor::Owned::from_element(trainer.ncenters, schema.dim(), 0.0); + + for (i, trained) in pivots.into_iter().enumerate() { + let offsets = schema.at(i); + for (dst, src) in std::iter::zip(packed.rows_mut(), trained.rows()) { + dst[offsets.clone()].copy_from_slice(src); + } + } + + let table = + BasicTable::new(packed, schema.to_owned()).map_err(|err| PQTrainingError { + chunk: schema.len(), + of: schema.len(), + dim: data.nrows(), + kind: PQTrainingErrorKind::InternalError(Box::new(err)), + })?; + + Ok(table) } train(self, data, schema, parallelism, rng_builder, cancelation) @@ -294,45 +268,22 @@ mod tests { use super::*; use crate::{cancel::DontCancel, error::format, random::StdRngBuilder}; - // With this test - we create sub-matrices that when flattened, will yield the output - // sequence `0, 1, 2, 3, 4, ...`. - #[test] - fn test_flatten() { - // The number of rows in the final matrix. - let nrows = 5; - // The dimensions in each sub-matrix. - let sub_dims = [1, 2, 3, 4, 5]; - // The prefix sum of the sub dimensions. - let prefix_sum: Vec = sub_dims - .iter() - .scan(0, |state, i| { - let this = *state; - *state += *i; - Some(this) - }) - .collect(); - - let dim: usize = sub_dims.iter().sum(); - - // Create the sub matrices. - let matrices: Vec> = - std::iter::zip(sub_dims.iter(), prefix_sum.iter()) - .map(|(&this_dim, &offset)| { - let mut m = rowmajor::Owned::from_element(nrows, this_dim, 0); - for r in 0..nrows { - for c in 0..this_dim { - *m.element_mut(r, c) = dim * r + offset + c; - } - } - m - }) - .collect(); - - let flattened = flatten(&matrices, nrows, dim); - // Check that the output is correct. - for (i, v) in flattened.as_slice().iter().enumerate() { - assert_eq!(*v, i, "failed at index {i}"); + fn flatten( + pivots: &[rowmajor::Owned], + ncenters: usize, + dim: usize, + ) -> rowmajor::Owned { + let mut flattened = rowmajor::Owned::from_element(ncenters, dim, T::default()); + let mut col_start = 0; + for matrix in pivots { + assert_eq!(matrix.nrows(), flattened.nrows()); + for (row_index, row) in matrix.rows().enumerate() { + let dst = &mut flattened.row_mut(row_index)[col_start..col_start + row.len()]; + dst.copy_from_slice(row); + } + col_start += matrix.ncols(); } + flattened } struct DatasetBuilder { @@ -460,10 +411,11 @@ mod tests { // 1. We ensure that the quantizer's center actually aligns with a cluster (i.e., // training did not invent values out of thin air). // 2. Every clustering in the original dataset has a representative in the quantizer. - assert_eq!(quantizer.dim, schema.dim()); - assert_eq!(quantizer.ncenters, ncenters); - assert_eq!(quantizer.pivots.len(), schema.len()); - for (i, pivot) in quantizer.pivots.iter().enumerate() { + assert_eq!(quantizer.dim(), schema.dim()); + assert_eq!(quantizer.ncenters(), ncenters); + assert_eq!(quantizer.nchunks(), schema.len()); + for i in 0..quantizer.nchunks() { + let pivot = quantizer.pivots_for(i).unwrap(); // Make sure the pivot has the correct dimension. assert_eq!( pivot.ncols(), @@ -507,13 +459,6 @@ mod tests { // Make sure that all clusters were seen. assert!(seen.iter().all(|i| *i), "not all clusters were seen"); } - - // Check `flatten`. - let flattened = quantizer.flatten(); - assert_eq!( - &flattened, - flatten(&quantizer.pivots, quantizer.ncenters, quantizer.dim).as_slice() - ); } #[test] @@ -615,14 +560,13 @@ mod tests { // pivot actually selected. // // All the rest should be zero. - let flat = flatten(&quantizer.pivots, quantizer.ncenters, quantizer.dim); assert!( - flat.row(0).iter().all(|i| *i == 1.0), + quantizer.view_pivots().row(0).iter().all(|i| *i == 1.0), "expected pivot 0 to be the non-zero pivot" ); - for (i, row) in flat.rows().enumerate() { + for (i, row) in quantizer.view_pivots().rows().enumerate() { // skip the first row. if i == 0 { continue; diff --git a/diskann-quantization/src/views.rs b/diskann-quantization/src/views.rs index 6bbe1cbdb3..11cf9393e9 100644 --- a/diskann-quantization/src/views.rs +++ b/diskann-quantization/src/views.rs @@ -181,14 +181,27 @@ where /// # Panics /// /// Panics if `i >= self.len()`. + #[expect( + clippy::panic, + reason = "this is documented to panic for out-of-bounds `i`" + )] pub fn at(&self, i: usize) -> core::ops::Range { - assert!( - i < self.len(), - "index {i} must be less than len {}", - self.len() - ); - let slice = self.offsets.as_slice(); - slice[i]..slice[i + 1] + match self.get(i) { + Some(range) => range, + None => panic!("index {i} must be less than len {}", self.len()), + } + } + + /// Return a range containing the start and one-past-the-end indices for chunk `i`. + /// + /// Returns `None` is `i >= self.len()`. + pub fn get(&self, i: usize) -> Option> { + if i < self.len() { + let slice = self.offsets.as_slice(); + Some(slice[i]..slice[i + 1]) + } else { + None + } } /// Return `self` as a view. @@ -446,12 +459,22 @@ mod tests { assert!(!offsets.is_empty()); assert_eq!(offsets.at(0), 0..1); + assert_eq!(offsets.get(0).unwrap(), 0..1); assert_eq!(offsets.at(1), 1..3); + assert_eq!(offsets.get(1).unwrap(), 1..3); assert_eq!(offsets.at(2), 3..6); + assert_eq!(offsets.get(2).unwrap(), 3..6); assert_eq!(offsets.at(3), 6..10); + assert_eq!(offsets.get(3).unwrap(), 6..10); assert_eq!(offsets.at(4), 10..12); + assert_eq!(offsets.get(4).unwrap(), 10..12); assert_eq!(offsets.at(5), 12..13); + assert_eq!(offsets.get(5).unwrap(), 12..13); assert_eq!(offsets.at(6), 13..14); + assert_eq!(offsets.get(6).unwrap(), 13..14); + + assert!(offsets.get(7).is_none()); + assert!(offsets.get(8).is_none()); // Finally, make sure the type is copyable. assert!(is_copyable(offsets)); @@ -496,6 +519,7 @@ mod tests { #[should_panic(expected = "index 5 must be less than len 3")] fn chunk_offset_indexing_panic() { let offsets = ChunkOffsets::new(Box::new([0, 1, 2, 3])).unwrap(); + assert!(offsets.get(5).is_none()); // panics let _ = offsets.at(5); From 16bed6b0aa87199cbe260953e9cfd7446e2152c7 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 6 Oct 2026 16:56:46 -0700 Subject: [PATCH 2/3] Small cleanups. --- .../src/product/tables/basic.rs | 45 +++++++++---------- diskann-quantization/src/product/train.rs | 4 +- 2 files changed, 23 insertions(+), 26 deletions(-) diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index bc0312f846..7056216863 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -119,7 +119,7 @@ where /// Return a [`Strided`] for the raw pivots of the requested chunk. /// - /// Returns `None` if `chunk >= self.nchunks`. + /// Returns `None` if `chunk >= self.nchunks()`. pub fn pivots_for(&self, chunk: usize) -> Option> { let range = self.offsets.get(chunk)?; @@ -127,15 +127,15 @@ where clippy::expect_used, reason = "the BasicTable's invariants mean this panic should be unreachable" )] - Some( - Strided::try_from_data( - &self.pivots.as_slice()[range.start..], - self.pivots.nrows(), - range.len(), - self.pivots.ncols(), - ) - .expect("BasicTable asserts that this layout is valid"), + let strided = Strided::try_from_data( + &self.pivots.as_slice()[range.start..], + self.pivots.nrows(), + range.len(), + self.pivots.ncols(), ) + .expect("BasicTable asserts that this layout is valid"); + + Some(strided) } } @@ -329,7 +329,7 @@ mod tests { let (pivots, offsets) = create_pivot_tables(schema.to_owned(), num_centers); let table = BasicTable::new(pivots, offsets).unwrap(); - // Check that `pivots_for` works as expected with repsect to the documented + // Check that `pivots_for` works as expected with respect to the documented // table configuration for `create_pivot_tables`. for chunk in 0..schema.len() { let strided = table.pivots_for(chunk).unwrap(); @@ -340,20 +340,17 @@ mod tests { let base = ((center + chunk) % num_centers) as f32; row.iter().enumerate().for_each(|(dim, b)| { let offset = if dim.is_multiple_of(2) { 0.25 } else { -0.25 }; - - if dim.is_multiple_of(2) { - assert_eq!( - *b, - base + offset, - "failed: chunk {} of {}, center {} of {}, dim {} of {}", - chunk, - schema.len(), - center, - num_centers, - dim, - row.len(), - ); - } + assert_eq!( + *b, + base + offset, + "failed: chunk {} of {}, center {} of {}, dim {} of {}", + chunk, + schema.len(), + center, + num_centers, + dim, + row.len(), + ); }) } } diff --git a/diskann-quantization/src/product/train.rs b/diskann-quantization/src/product/train.rs index 364ed764f7..9b64cf777b 100644 --- a/diskann-quantization/src/product/train.rs +++ b/diskann-quantization/src/product/train.rs @@ -62,7 +62,7 @@ impl TrainQuantizer for LightPQTrainingParameters { type Error = PQTrainingError; /// Perform product quantization training on the provided training set and return a - /// `SimplePivots` containing the result of kmeans clustering on each partition. + /// [`BasicTable`] containing the result of kmeans clustering on each partition. /// /// # Panics /// @@ -208,7 +208,7 @@ impl TrainQuantizer for LightPQTrainingParameters { BasicTable::new(packed, schema.to_owned()).map_err(|err| PQTrainingError { chunk: schema.len(), of: schema.len(), - dim: data.nrows(), + dim: data.ncols(), kind: PQTrainingErrorKind::InternalError(Box::new(err)), })?; From 9a47e2051d499b2e14a8a3fffead364fa5350faf Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Wed, 7 Oct 2026 06:57:07 -0700 Subject: [PATCH 3/3] Use `pivots_for`. --- .../src/product/tables/padded.rs | 21 +++---------------- 1 file changed, 3 insertions(+), 18 deletions(-) diff --git a/diskann-quantization/src/product/tables/padded.rs b/diskann-quantization/src/product/tables/padded.rs index 9e4601b0ce..5a9ddeecc4 100644 --- a/diskann-quantization/src/product/tables/padded.rs +++ b/diskann-quantization/src/product/tables/padded.rs @@ -45,10 +45,7 @@ use std::{marker::PhantomData, num::NonZeroUsize}; -use diskann_utils::{ - strided, - views::rowmajor::{self, Matrix, MatrixMut}, -}; +use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut}; use diskann_vector::distance::Metric as VectorMetric; use diskann_wide::{ SIMDFloat, SIMDSumTree, SIMDVector, @@ -131,20 +128,8 @@ impl PaddedTable { // Copy the pivots. (0..offsets.len()).for_each(|i| { - let range = offsets.at(i); - - #[expect( - clippy::expect_used, - reason = "the layout should be pre-validated by `BasicTable`" - )] - let view = strided::Strided::try_from_data( - &(pivots.as_slice()[range.start..]), - pivots.nrows(), - range.len(), - offsets.dim(), - ) - .expect("the check on `pivot_dim` and `offsets_dim` should cause this to never error"); - + #[expect(clippy::expect_used, reason = "`i` should be in-bounds")] + let view = basic.pivots_for(i).expect("`i` should be in-bounds"); view.rows().for_each(|src| { copy_from_slice_subset(padded.row_mut(row), src); row += 1;