Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 1 addition & 11 deletions diskann-benchmark/src/exhaustive/product.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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()?;
Expand All @@ -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
Expand Down
19 changes: 9 additions & 10 deletions diskann-providers/src/model/pq/pq_construction.rs
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ where
parameters.max_k_means_reps(),
);

let full_pivot_data = pool.install(|| -> Result<Vec<f32>, ANNError> {
let basic_table = pool.install(|| -> Result<_, ANNError> {
let result = trainer
.train(
rowmajor::Ref::try_from_data(train_data, parameters.num_train(), parameters.dim())
Expand All @@ -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(),
Expand Down Expand Up @@ -200,7 +199,7 @@ pub fn generate_pq_pivots_from_membuf<T: Copy + Into<f32>>(
);

let rng_builder = create_rnd_provider_from_seed(rand::distr::StandardUniform {}.sample(rng));
let trained = pool.install(|| -> Result<Vec<f32>, 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`.
//
Expand All @@ -216,7 +215,7 @@ pub fn generate_pq_pivots_from_membuf<T: Copy + Into<f32>>(
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(),
Expand All @@ -229,12 +228,12 @@ pub fn generate_pq_pivots_from_membuf<T: Copy + Into<f32>>(
&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(())
}

Expand Down
62 changes: 59 additions & 3 deletions diskann-quantization/src/product/tables/basic.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<Strided<'_, f32>> {
let range = self.offsets.get(chunk)?;

#[expect(
clippy::expect_used,
reason = "the BasicTable's invariants mean this panic should be unreachable"
)]
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)
}
}

#[derive(Error, Debug)]
Expand Down Expand Up @@ -304,6 +328,38 @@ 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 respect 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 };
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()];
Expand Down
21 changes: 3 additions & 18 deletions diskann-quantization/src/product/tables/padded.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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;
Expand Down
Loading
Loading