diff --git a/diskann-benchmark/example/product-exhaustive.json b/diskann-benchmark/example/product-exhaustive.json index 6e12cb32a7..d384204aaa 100644 --- a/diskann-benchmark/example/product-exhaustive.json +++ b/diskann-benchmark/example/product-exhaustive.json @@ -27,7 +27,8 @@ "compression_threads": 1, "seed": 7831252621480178695, "num_pq_chunks": 16, - "num_pq_centers": 16 + "num_pq_centers": 16, + "table_style": "transposed" } } ] diff --git a/diskann-benchmark/src/exhaustive/product.rs b/diskann-benchmark/src/exhaustive/product.rs index ca52a0a5dc..200d76ed0c 100644 --- a/diskann-benchmark/src/exhaustive/product.rs +++ b/diskann-benchmark/src/exhaustive/product.rs @@ -34,8 +34,13 @@ mod imp { utils::{percentiles, MicroSeconds}, Benchmark, Output, }; - use diskann_quantization::{product::train::TrainQuantizer, CompressInto}; + use diskann_providers::model::pq::FixedChunkPQTable; + use diskann_quantization::{ + product::{tables, train::TrainQuantizer}, + CompressInto, + }; use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut}; + use diskann_vector::distance::Metric; use indicatif::{ProgressBar, ProgressStyle}; use rayon::iter::{IndexedParallelIterator, ParallelIterator}; use serde::Serialize; @@ -109,10 +114,14 @@ mod imp { })? }; - let quantizer = diskann_providers::model::pq::FixedChunkPQTable::new( - data.ncols(), - base.flatten().into(), - offsets.as_slice().into(), + // 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(); @@ -126,8 +135,14 @@ mod imp { let compression_progress = make_progress_bar("compressing", data.nrows(), output.draw_target())?; - let store = threadpool - .install(|| Store::new(data.as_view(), quantizer, &compression_progress))?; + let store = threadpool.install(|| { + Store::new( + data.as_view(), + table, + input.table_style, + &compression_progress, + ) + })?; compression_progress.finish(); store }; @@ -321,32 +336,246 @@ mod imp { } } + //------------// + // Compressor // + //------------// + + #[derive(Debug)] + enum Compressor { + FixedChunk(FixedChunkPQTable), + Transposed(tables::TransposedTable), + } + + impl Compressor { + fn new( + table: tables::BasicTable, + style: inputs::exhaustive::PQTableStyle, + ) -> anyhow::Result { + use inputs::exhaustive::PQTableStyle; + match style { + PQTableStyle::FixedChunk => Ok(Self::FixedChunk(table.try_into()?)), + PQTableStyle::Padded | PQTableStyle::Transposed => { + let table = tables::TransposedTable::from_parts( + table.view_pivots(), + table.view_offsets().to_owned(), + )?; + + Ok(Self::Transposed(table)) + } + } + } + + fn compress(&self, storage: &mut [u8], data: &[f32]) -> anyhow::Result<()> { + match self { + Self::FixedChunk(table) => table.compress_into(data, storage)?, + Self::Transposed(table) => table.compress_into(data, storage)?, + } + Ok(()) + } + } + + //-----------// + // Distances // + //-----------// + + trait ComputerImpl: std::fmt::Debug { + fn evaluate(&self, x: &[u8]) -> anyhow::Result; + } + + #[derive(Debug)] + struct Computer<'a>(Box); + + impl<'a> Computer<'a> { + fn new(inner: C) -> Self + where + C: ComputerImpl + 'a, + { + Self(Box::new(inner)) + } + } + + impl diskann_vector::PreprocessedDistanceFunction<&[u8], f32> for Computer<'_> { + fn evaluate_similarity(&self, x: &[u8]) -> f32 { + match self.0.evaluate(x) { + Ok(v) => v, + Err(err) => panic!("distance failed with {:#}", err), + } + } + } + + //-----------// + // Computers // + //-----------// + + impl ComputerImpl for diskann_providers::model::pq::distance::QueryComputer<'_> { + fn evaluate(&self, x: &[u8]) -> anyhow::Result { + Ok(>::evaluate_similarity(self, x)) + } + } + + #[derive(Debug)] + struct PaddedComputer<'a> { + table: &'a tables::PaddedTable, + vtable: tables::padded::VTable, + query: Vec, + } + + impl ComputerImpl for PaddedComputer<'_> { + fn evaluate(&self, x: &[u8]) -> anyhow::Result { + Ok(self.vtable.distance(self.table, &self.query, x)?) + } + } + + #[derive(Debug)] + struct LookupTable { + lookup: rowmajor::Owned, + } + + impl ComputerImpl for LookupTable { + fn evaluate(&self, x: &[u8]) -> anyhow::Result { + Ok(tables::lookup::lookup_single( + tables::lookup::Sum, + self.lookup.as_view(), + x, + )?) + } + } + + #[derive(Debug)] + struct CosineLookupTable { + lookup: rowmajor::Owned, + query_norm: f32, + } + + impl ComputerImpl for CosineLookupTable { + fn evaluate(&self, x: &[u8]) -> anyhow::Result { + let partial = + tables::lookup::lookup_single(tables::lookup::Sum, self.lookup.as_view(), x)?; + Ok(partial.finish_cosine(self.query_norm).into_inner()) + } + } + + #[derive(Debug)] + enum Distance { + FixedChunk(FixedChunkPQTable), + Padded(tables::PaddedTable), + Transposed(tables::TransposedTable), + } + + impl Distance { + fn new( + basic: tables::BasicTable, + style: inputs::exhaustive::PQTableStyle, + ) -> anyhow::Result { + use inputs::exhaustive::PQTableStyle; + match style { + PQTableStyle::FixedChunk => Ok(Self::FixedChunk(basic.try_into()?)), + PQTableStyle::Padded => Ok(Self::Padded(tables::PaddedTable::from_basic( + basic.as_view(), + ))), + PQTableStyle::Transposed => { + Ok(Self::Transposed(tables::TransposedTable::from_parts( + basic.view_pivots(), + basic.view_offsets().to_owned(), + )?)) + } + } + } + + fn computer(&self, query: &[f32], metric: Metric) -> anyhow::Result> { + match self { + Self::FixedChunk(table) => { + let inner = diskann_providers::model::pq::distance::QueryComputer::new( + table.into(), + metric, + query, + None, + )?; + Ok(Computer::new(inner)) + } + Self::Padded(table) => { + let inner = PaddedComputer { + table, + vtable: table.vtable(metric.into()), + query: query.into(), + }; + + Ok(Computer::new(inner)) + } + Self::Transposed(table) => match metric { + Metric::L2 => { + let mut lookup = + rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0); + table.process_into::( + query, + lookup.as_view_mut(), + ); + Ok(Computer::new(LookupTable { lookup })) + } + Metric::InnerProduct => { + let mut lookup = + rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0); + table.process_into::( + query, + lookup.as_view_mut(), + ); + Ok(Computer::new(LookupTable { lookup })) + } + Metric::Cosine | Metric::CosineNormalized => { + let mut lookup = rowmajor::Owned::from_element( + table.nchunks(), + table.ncenters(), + tables::lookup::DotAndNorm::default(), + ); + + let query_norm = <_ as diskann_vector::Norm<&[f32]>>::evaluate( + &diskann_vector::norm::FastL2Norm, + query, + ); + + table.process_into::( + query, + lookup.as_view_mut(), + ); + Ok(Computer::new(CosineLookupTable { lookup, query_norm })) + } + }, + } + } + } + /// A store for quantized data. pub(super) struct Store { data: rowmajor::Owned, - quantizer: diskann_providers::model::pq::FixedChunkPQTable, + distance: Distance, } impl Store { fn new( input: rowmajor::Ref, - quantizer: diskann_providers::model::pq::FixedChunkPQTable, + table: tables::BasicTable, + style: inputs::exhaustive::PQTableStyle, progress: &ProgressBar, ) -> anyhow::Result { - let mut data = - rowmajor::Owned::try_from_element(input.nrows(), quantizer.get_num_chunks(), 0)?; + let mut data = rowmajor::Owned::try_from_element(input.nrows(), table.nchunks(), 0)?; + + let compressor = Compressor::new(table.clone(), style)?; // Compress the data. #[expect(clippy::disallowed_methods)] data.par_rows_mut().zip(input.par_rows()).try_for_each( |(d, i)| -> anyhow::Result<()> { - quantizer.compress_into(i, d)?; + compressor.compress(d, i)?; progress.inc(1); Ok(()) }, )?; - Ok(Self { data, quantizer }) + let distance = Distance::new(table, style)?; + Ok(Self { data, distance }) } } @@ -366,19 +595,14 @@ mod imp { } impl algos::CreateQuantComputer for Plan { - type Computer<'a> = diskann_providers::model::pq::distance::QueryComputer<'a>; + type Computer<'a> = Computer<'a>; fn create_quant_computer<'a>( &self, store: &'a Store, query: &[f32], ) -> anyhow::Result> { - Ok(diskann_providers::model::pq::distance::QueryComputer::new( - (&store.quantizer).into(), - self.measure.into(), - query, - None, - )?) + store.distance.computer(query, self.measure.into()) } } } diff --git a/diskann-benchmark/src/inputs/exhaustive.rs b/diskann-benchmark/src/inputs/exhaustive.rs index 20583de85c..5331e71d68 100644 --- a/diskann-benchmark/src/inputs/exhaustive.rs +++ b/diskann-benchmark/src/inputs/exhaustive.rs @@ -200,6 +200,24 @@ impl From<&TransformKind> for diskann_quantization::algorithms::transforms::Tran // Product Quantization Methods // ////////////////////////////////// +#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub(crate) enum PQTableStyle { + FixedChunk, + Padded, + Transposed, +} + +impl PQTableStyle { + fn as_str(&self) -> &'static str { + match self { + PQTableStyle::FixedChunk => "fixed-chunk", + PQTableStyle::Padded => "padded", + PQTableStyle::Transposed => "transposed", + } + } +} + #[derive(Debug, Serialize, Deserialize)] pub(crate) struct Product { pub(crate) data: InputFile, @@ -210,6 +228,7 @@ pub(crate) struct Product { pub(crate) seed: u64, pub(crate) num_pq_chunks: NonZeroUsize, pub(crate) num_pq_centers: NonZeroUsize, + pub(crate) table_style: PQTableStyle, } impl Product { @@ -251,6 +270,7 @@ impl Example for Product { seed: 0x6cae32c479ac3407, num_pq_chunks: NUM_PQ_CHUNKS, num_pq_centers: NUM_PQ_CENTERS, + table_style: PQTableStyle::FixedChunk, } } } @@ -264,6 +284,7 @@ impl std::fmt::Display for Product { write_field!(f, "seed", self.seed)?; write_field!(f, "PQ Chunks", self.num_pq_chunks.get())?; write_field!(f, "PQ Centers", self.num_pq_centers.get())?; + write_field!(f, "Table Style", self.table_style.as_str())?; Ok(()) } } diff --git a/diskann-disk/src/search/pq/quantizer_preprocess.rs b/diskann-disk/src/search/pq/quantizer_preprocess.rs index 1bc1e676f9..6a3227571e 100644 --- a/diskann-disk/src/search/pq/quantizer_preprocess.rs +++ b/diskann-disk/src/search/pq/quantizer_preprocess.rs @@ -38,13 +38,13 @@ impl PQScratch { // We're keeping that behavior here - treating `Cosine` and `CosineNormalized` // as L2 until a more thorough evaluation can be made. Metric::L2 | Metric::Cosine | Metric::CosineNormalized => { - table.process_into::( + table.process_into::( &self.query_scratch, dst, ); } Metric::InnerProduct => { - table.process_into::( + table.process_into::( &self.query_scratch, dst, ); diff --git a/diskann-disk/src/storage/quant/generator.rs b/diskann-disk/src/storage/quant/generator.rs index 51748561c9..7209dd100b 100644 --- a/diskann-disk/src/storage/quant/generator.rs +++ b/diskann-disk/src/storage/quant/generator.rs @@ -136,7 +136,7 @@ where // process `BATCH_SIZE` many dataset vectors at a time. const BATCH_SIZE: usize = 128; - // Wrap the data in `rowmajor::Mut` so we do not need to manually construct view + // Wrap the data in `rowmajor::Mut` so we do not need to manually construct a view // in the compression loop. let mut compressed_block = views::rowmajor::Mut::try_from_data( block_compressed_base, diff --git a/diskann-quantization/src/distances.rs b/diskann-quantization/src/distances.rs index 5e93355974..81975bd884 100644 --- a/diskann-quantization/src/distances.rs +++ b/diskann-quantization/src/distances.rs @@ -82,6 +82,10 @@ pub struct SquaredL2; #[derive(Debug, Clone, Copy)] pub struct InnerProduct; +/// Compute the cosine similarity between vector-like types. +#[derive(Debug, Clone, Copy)] +pub struct Cosine; + /// Compute the hamming distance between bit-vectors. #[derive(Debug, Clone, Copy)] pub struct Hamming; diff --git a/diskann-quantization/src/product/mod.rs b/diskann-quantization/src/product/mod.rs index 06cc1330a8..9f7f09f999 100644 --- a/diskann-quantization/src/product/mod.rs +++ b/diskann-quantization/src/product/mod.rs @@ -5,10 +5,9 @@ //! Product quantization training and compression. +pub mod tables; pub mod train; -mod tables; - ///////////// // Exports // ///////////// diff --git a/diskann-quantization/src/product/tables/basic.rs b/diskann-quantization/src/product/tables/basic.rs index 4816916316..f8db42bd84 100644 --- a/diskann-quantization/src/product/tables/basic.rs +++ b/diskann-quantization/src/product/tables/basic.rs @@ -105,6 +105,14 @@ where pub fn dim(&self) -> usize { self.pivots.ncols() } + + /// Return `self` as a [`BasicTableView`]. + pub fn as_view(&self) -> BasicTableView<'_> { + BasicTableView { + pivots: self.pivots.as_view(), + offsets: self.offsets.as_view(), + } + } } #[derive(Error, Debug)] diff --git a/diskann-quantization/src/product/tables/lookup.rs b/diskann-quantization/src/product/tables/lookup.rs new file mode 100644 index 0000000000..dffc788a67 --- /dev/null +++ b/diskann-quantization/src/product/tables/lookup.rs @@ -0,0 +1,343 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_utils::views::rowmajor::{self, Matrix}; +use thiserror::Error; + +/// Policy for processing entries in [`lookup_single`]. +/// +/// The lookup operation may be performed by using multiple independent accumulators, each +/// processing a subset of the total lookup operation. +pub trait Lookup { + /// The type of the accumulator. + /// + /// Multiple such accumulators may be instantiated. + type Accumulator; + + /// The type of the final result. + type Output; + + /// Create a default accumulator. + fn default(&self) -> Self::Accumulator; + + /// Accumulate the element `v` with the current accumulator. + fn accumulate(&self, v: &T, acc: Self::Accumulator) -> Self::Accumulator; + + /// Combine the results of two independent accumulators. + fn reduce(&self, a: Self::Accumulator, b: Self::Accumulator) -> Self::Accumulator; + + /// Process the final accumulator, returning the result. + fn finish(self, acc: Self::Accumulator) -> Self::Output; +} + +/// Use `indices` to retrieve one item from each row of `data` and accumulate the results +/// across all rows using `policy`. +/// ``` +/// use diskann_quantization::product::tables::lookup; +/// use diskann_utils::views::rowmajor::{Owned, Matrix}; +/// +/// // Make the following matrix: +/// // 0 1 +/// // 2 3 +/// // 4 5 +/// let data = Owned::from_fn(3, 2, |rc| 2 * rc.row + rc.col); +/// let sum = lookup::lookup_single( +/// lookup::Sum, +/// data.as_view(), +/// &[0, 1, 0] +/// ).unwrap(); +/// +/// assert_eq!(sum, 0 + 3 + 4); +/// ``` +/// +/// # Errors +/// +/// Returns an error if `indices.len() != data.nrows()` (there must be a one-to-one +/// correspondence between indices and rows) or if any index in `indices` exceeds `data.ncols()`. +pub fn lookup_single( + policy: P, + data: rowmajor::Ref<'_, T>, + indices: &[u8], +) -> Result +where + P: Lookup, +{ + // Check 1. + if indices.len() != data.nrows() { + return Err(LookupError::InvalidLength); + } + + // Check 2. + // + // Conversion fails if `data.ncols()` is 256 or greater. + // + // In this case, all indices will be in bounds anyway. + if let Ok(ncols) = u8::try_from(data.ncols()) + && let Some(max) = indices.iter().max() + && *max >= ncols + { + return Err(LookupError::OutOfBounds); + } + + const UNROLL: usize = 4; + + let mut i = 0; + let mut a = if data.nrows() >= UNROLL { + let mut a0 = policy.default(); + let mut a1 = policy.default(); + let mut a2 = policy.default(); + let mut a3 = policy.default(); + + while i + 4 <= data.nrows() { + // SAFETY: Since `i + 4 <= data.nrows()`, the row index `i` is in-bounds. Check 1 + // ensures that `i` is also a valid index for `indices`. + // + // Check 2 guarantees that `indices[i] < data.ncols()`, so the matrix access is + // valid. + let v0 = unsafe { data.element_unchecked(i, (*indices.get_unchecked(i)).into()) }; + a0 = policy.accumulate(v0, a0); + + // SAFETY: Same as above. + let v1 = + unsafe { data.element_unchecked(i + 1, (*indices.get_unchecked(i + 1)).into()) }; + a1 = policy.accumulate(v1, a1); + + // SAFETY: Same as above. + let v2 = + unsafe { data.element_unchecked(i + 2, (*indices.get_unchecked(i + 2)).into()) }; + a2 = policy.accumulate(v2, a2); + + // SAFETY: Same as above. + let v3 = + unsafe { data.element_unchecked(i + 3, (*indices.get_unchecked(i + 3)).into()) }; + a3 = policy.accumulate(v3, a3); + + i += UNROLL; + } + + policy.reduce(policy.reduce(a0, a1), policy.reduce(a2, a3)) + } else { + policy.default() + }; + + if let Some(remainder) = data.nrows().checked_sub(i) { + // Hint to LLVM that the loop below is bounded. + let remainder = remainder.min(UNROLL - 1); + for j in 0..remainder { + let k = i + j; + + // SAFETY: See justification in the main unrolled loop. + let v = unsafe { data.element_unchecked(k, (*indices.get_unchecked(k)).into()) }; + a = policy.accumulate(v, a); + } + } + + Ok(policy.finish(a)) +} + +/// Errors from [`lookup_single`]. +#[derive(Debug, Error, Clone, Copy)] +#[non_exhaustive] +pub enum LookupError { + #[error("number of lookup indices does not match the number of data rows")] + InvalidLength, + #[error("at least one of the lookup indices exceeds the number of data columns")] + OutOfBounds, +} + +/// A simple [`Lookup`] that uses `std::ops::Add` to accumulate results. +#[derive(Debug, Clone, Copy)] +pub struct Sum; + +impl Lookup for Sum +where + T: Default + std::ops::Add + Copy, +{ + type Accumulator = T; + type Output = T; + + fn default(&self) -> T { + T::default() + } + + fn accumulate(&self, v: &T, acc: T) -> T { + *v + acc + } + + fn reduce(&self, a: T, b: T) -> T { + a + b + } + + fn finish(self, acc: T) -> T { + acc + } +} + +/// An element for [`lookup_single`] that is used for computing cosine similarity. +/// +/// Each [`DotAndNorm`] consists of a partial dot-product (e.g. the dot-product between a +/// query chunk and a PQ center) as well as the PQ center's squared norm. +/// +/// After the lookup operation, the final [`DotAndNorm`] consists of the dot-product between +/// the query and the effective data vector as well as the total squared norm of the effective +/// data vector. +#[derive(Debug, Clone, Copy, Default)] +#[repr(C)] +pub struct DotAndNorm { + dot: f32, + square_norm: f32, +} + +impl DotAndNorm { + /// Construct a new [`DotAndNorm`]. + pub const fn new(dot: f32, square_norm: f32) -> Self { + Self { dot, square_norm } + } + + /// Return the current value of the dot-product. + pub fn dot(&self) -> f32 { + self.dot + } + + /// Return the current value of the squared norm. + pub fn square_norm(&self) -> f32 { + self.square_norm + } + + /// Finish a cosine computation using `query_norm`. This computes: + /// ```math + /// 1.0 - (self.dot) / (self.square_norm.sqrt() * query_norm) + /// ``` + /// taking care to avoid division by zero. + /// + /// Note that this returns a [`diskann_vector::SimilarityScore`] for use in similarity + /// reranking. + pub fn finish_cosine(&self, query_norm: f32) -> diskann_vector::SimilarityScore { + use diskann_vector::SimilarityScore; + + if self.square_norm < f32::MIN_POSITIVE || query_norm < f32::MIN_POSITIVE { + SimilarityScore::new(1.0) + } else { + let v = self.dot / (self.square_norm.sqrt() * query_norm); + SimilarityScore::new(1.0 - (-1.0f32).max(1.0f32.min(v))) + } + } +} + +impl std::ops::Add for DotAndNorm { + type Output = Self; + fn add(self, rhs: Self) -> Self { + Self { + dot: self.dot + rhs.dot, + square_norm: self.square_norm + rhs.square_norm, + } + } +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use diskann_utils::assert_contains; + use rand::{ + SeedableRng, + distr::{Distribution, Uniform}, + rngs::StdRng, + }; + + fn expected_sum(x: &[u8]) -> f32 { + let mut sum = 0.0; + for (i, v) in x.iter().enumerate() { + sum += (i as f32) + (*v as f32) + } + sum + } + + #[test] + fn test_lookup_sum() { + let ntrials = if cfg!(miri) { 1 } else { 10 }; + + let mut rng = StdRng::seed_from_u64(0xd0cc501bde4c9ddd); + for nrows in 0..12 { + let mut codes = vec![0u8; nrows]; + for ncols in [1, 2, 255, 256] { + let dist = if ncols == 0 { + Uniform::new(0, 1).unwrap() + } else { + Uniform::new(0, ncols).unwrap() + }; + + let table = rowmajor::Owned::from_fn(nrows, ncols, |rc| (rc.row + rc.col) as f32); + + // If it's possible for a value to be out-of-bounds, make sure we return an + // error if anything *is* out-of-bounds. + if ncols < 256 { + let nc = u8::try_from(ncols).unwrap(); + codes.fill(0); + for r in 0..nrows { + codes[r] = nc; + let err = lookup_single(Sum, table.as_view(), &codes).unwrap_err(); + assert_contains!( + err.to_string(), + "at least one of the lookup indices exceeds the number of data columns" + ); + codes[r] = 0; + } + } + + // Check that too long and too short codes are detected. + if nrows > 0 { + let too_short = vec![0; nrows - 1]; + let err = lookup_single(Sum, table.as_view(), &too_short).unwrap_err(); + assert_contains!(err.to_string(), "number of lookup indices does not match"); + } + + { + let too_long = vec![0; nrows + 1]; + let err = lookup_single(Sum, table.as_view(), &too_long).unwrap_err(); + assert_contains!(err.to_string(), "number of lookup indices does not match"); + } + + // Test all zeros + codes.iter_mut().for_each(|c| *c = 0); + assert_eq!( + lookup_single(Sum, table.as_view(), &codes).unwrap(), + expected_sum(&codes), + "all zeros - nrows = {nrows}, ncols = {ncols}", + ); + + // Test all max + codes + .iter_mut() + .for_each(|c| *c = (ncols - 1).try_into().unwrap()); + assert_eq!( + lookup_single(Sum, table.as_view(), &codes).unwrap(), + expected_sum(&codes), + "all max - nrows = {nrows}, ncols = {ncols}", + ); + + for trial in 0..ntrials { + codes + .iter_mut() + .for_each(|c| *c = u8::try_from(dist.sample(&mut rng)).unwrap()); + + println!("codes = {:?}, ncols = {}", codes, ncols); + + assert_eq!( + lookup_single(Sum, table.as_view(), &codes).unwrap(), + expected_sum(&codes), + "nrows = {nrows}, ncols = {ncols}, trial = {} of {}", + trial + 1, + ntrials, + ); + } + } + } + } +} diff --git a/diskann-quantization/src/product/tables/mod.rs b/diskann-quantization/src/product/tables/mod.rs index 96218db82c..e7f5975503 100644 --- a/diskann-quantization/src/product/tables/mod.rs +++ b/diskann-quantization/src/product/tables/mod.rs @@ -4,8 +4,11 @@ */ mod basic; +pub mod padded; mod transposed; +pub mod lookup; + #[cfg(test)] pub(super) mod test; @@ -13,6 +16,6 @@ pub(super) mod test; // Exports // ///////////// -// Error types pub use basic::{BasicTable, BasicTableBase, BasicTableView, TableCompressionError}; +pub use padded::PaddedTable; pub use transposed::TransposedTable; diff --git a/diskann-quantization/src/product/tables/padded.rs b/diskann-quantization/src/product/tables/padded.rs new file mode 100644 index 0000000000..9e4601b0ce --- /dev/null +++ b/diskann-quantization/src/product/tables/padded.rs @@ -0,0 +1,1383 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! A PQ table optimized for computing distances between quantized vectors. +//! +//! During similarity search index construction, it is not uncommon to compute distances +//! among quantized elements within a table. The [`PaddedTable`] is designed to facilitate +//! such distances. +//! +//! ``` +//! use diskann_quantization::{views::ChunkOffsets, product::tables}; +//! use diskann_utils::views::rowmajor::{Owned, Matrix, MatrixMut}; +//! +//! // We're creating the following pivot table. +//! // +//! // | chunk 0 | chunk 1 | +//! // | 0 0 | 1 1 | pivot 0 +//! // | 1 1 | 2 2 | pivot 1 +//! // | 2 2 | 3 3 | pivot 2 +//! +//! let mut pivots = Owned::from_element(3, 4, 0.0f32); +//! pivots.row_mut(0).copy_from_slice(&[0.0, 0.0, 1.0, 1.0]); +//! pivots.row_mut(1).copy_from_slice(&[1.0, 1.0, 2.0, 2.0]); +//! pivots.row_mut(2).copy_from_slice(&[2.0, 2.0, 3.0, 3.0]); +//! +//! let offsets = ChunkOffsets::new(Box::new([0, 2, 4])).unwrap(); +//! +//! let basic = tables::BasicTable::new(pivots, offsets).unwrap(); +//! let padded = tables::PaddedTable::from_basic(basic.as_view()); +//! +//! // Distances are provided through a v-table. +//! let vtable = padded.vtable(tables::padded::Metric::SquaredL2); +//! +//! // Compute the distance between [1, 1, 1, 1] and the chunk defined by [2, 1], which +//! // should translate to the compressed vector [2, 2, 2, 2]. +//! let distance = vtable.distance(&padded, &[1.0, 1.0, 1.0, 1.0], &[2, 1]).unwrap(); +//! assert_eq!(distance, 4.0); +//! +//! // Compute the distance between the two compressed vectors encoded by [0, 2] and [2, 0]. +//! let distance = vtable.self_distance(&padded, &[0, 2], &[2, 0]).unwrap(); +//! assert_eq!(distance, 16.0); +//! ``` + +use std::{marker::PhantomData, num::NonZeroUsize}; + +use diskann_utils::{ + strided, + views::rowmajor::{self, Matrix, MatrixMut}, +}; +use diskann_vector::distance::Metric as VectorMetric; +use diskann_wide::{ + SIMDFloat, SIMDSumTree, SIMDVector, + arch::{Architecture, Dispatched3, FTarget3, Scalar, Target, dispatch_no_features}, + lifetime::Ref, +}; +use thiserror::Error; + +#[cfg(target_arch = "x86_64")] +use diskann_wide::arch::x86_64::{V3, V4}; + +#[cfg(target_arch = "aarch64")] +use diskann_wide::arch::aarch64::Neon; + +use crate::{product::tables::BasicTableView, views::ChunkOffsets}; + +/// A PQ table that stores pivots grouped by chunk in the following dense, row-major form: +/// ```text +/// | -- pivot 0 -- | -- pivot 1 -- | .... | -- pivot K-1 -- | +/// +------------------+------------------+------+--------------------+ +/// chunk 0 | c000 c001 ... 0X | c010 c011 ... 0X | .... | c0K0 c0K1 ... 0X | +/// chunk 1 | c100 c101 ... 0X | c110 c111 ... 0X | .... | c1K0 c1K1 ... 0X | +/// ... | ... | ... | .... | ... | +/// chunk N-1 | cN00 cN01 ... 0X | cN10 cN11 ... 0X | .... | cNK0 cNK1 ... 0X | +/// ``` +/// where `cCPD` is dimension `D` of pivot `P` in chunk `C`, and trailing `0X`s denote +/// potential zero-padding to the SIMD-aligned pivot width. +/// +/// The member `offsets` describes the number of *unpadded* dimensions of each chunk. +/// +/// Importantly, though, the storage for each pivot is rounded up to a multiple of the +/// runtime system's preferred SIMD width and all pivots are padded to the same length. +/// This makes distance computations between pivots very fast when computing distances +/// between two product-quantized vectors. +#[derive(Debug, Clone)] +pub struct PaddedTable { + /// Invariants: + /// * `pivots.ncols()` is at least as large as the largest chunk in [`Self::offsets`] + /// and is always a multiple of `simd_width`. + /// * `pivots.nrows() == offsets.len() * pivots_per_chunk`. + pivots: rowmajor::Owned, + offsets: ChunkOffsets, + pivots_per_chunk: usize, + simd_width: SIMDWidth, + arch: RuntimeArch, +} + +impl PaddedTable { + /// Construct a [`PaddedTable`] with the same contents as `basic`. + pub fn from_basic(basic: BasicTableView<'_>) -> Self { + let arch = RuntimeArch::new(); + Self::from_basic_with(basic, arch) + } + + fn from_basic_with(basic: BasicTableView<'_>, arch: RuntimeArch) -> Self { + let pivots = basic.view_pivots(); + let offsets = basic.view_offsets(); + + let pivots_per_chunk = pivots.nrows(); + + // Compute the padded dimension of the pivots. + // + // There exists a corner case where `pivots` barely fits within the `isize` limit + // of an allocation and padding will put us beyond that threshold, but that is + // exceedingly unlikely for typical data. + let max_chunk_dim = offsets.max_chunk_dim(); + let simd_width = arch.select_simd_width(max_chunk_dim); + let padded_dim = max_chunk_dim + .get() + .next_multiple_of(simd_width.as_nonzero().get()); + + let rows = pivots_per_chunk * offsets.len(); + let mut padded = rowmajor::Owned::from_element(rows, padded_dim, 0.0); + let mut row = 0; + + // Since we padded, we use a custom "copy_from_slice" that allows `dst` to shrink. + fn copy_from_slice_subset(dst: &mut [f32], src: &[f32]) { + dst[..src.len()].copy_from_slice(src) + } + + // 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"); + + view.rows().for_each(|src| { + copy_from_slice_subset(padded.row_mut(row), src); + row += 1; + }); + }); + + Self { + pivots: padded, + offsets: offsets.to_owned(), + pivots_per_chunk, + simd_width, + arch, + } + } + + /// Return the distance [`VTable`] for the requested [`Metric`]. + /// + /// Note that [`VTable`]s are generally specific to the [`PaddedTable`] that generated + /// them and cannot be reliably shared with different tables. + /// + /// While this is not a safety issue, incorrect [`VTable`]s may yield errors due to + /// mismatched SIMD widths. + pub fn vtable(&self, metric: Metric) -> VTable { + self.arch.dispatch(self.simd_width, metric) + } + + /// Return the number of PQ centers per chunk. + pub fn ncenters(&self) -> usize { + self.pivots_per_chunk + } + + /// Return the number of PQ chunks. + pub fn nchunks(&self) -> usize { + self.offsets.len() + } + + /// Return the full-precision dimension expected by this table. + pub fn dim(&self) -> usize { + self.offsets.dim() + } +} + +/// Distance metrics used by [`PaddedTable`]. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Metric { + SquaredL2, + InnerProduct, + Cosine, +} + +impl From for Metric { + fn from(metric: VectorMetric) -> Self { + match metric { + VectorMetric::L2 => Self::SquaredL2, + VectorMetric::InnerProduct => Self::InnerProduct, + VectorMetric::Cosine => Self::Cosine, + VectorMetric::CosineNormalized => Self::Cosine, + } + } +} + +type Distance = Dispatched3, Ref, Ref<[f32]>, Ref<[u8]>>; +type SelfDistance = + Dispatched3, Ref, Ref<[u8]>, Ref<[u8]>>; + +/// A distance [`VTable`] for a [`PaddedTable`]. +/// +/// See: [`PaddedTable::vtable`]. +#[derive(Debug, Clone, Copy)] +pub struct VTable { + distance: Distance, + self_distance: SelfDistance, +} + +impl VTable { + fn new(arch: O::Arch) -> Self + where + O: Op, + { + Self { + distance: arch.dispatch3::< + OpTarget, + Result, + Ref, + Ref<[f32]>, + Ref<[u8]>, + >(), + self_distance: arch.dispatch3::< + OpTarget, + Result, + Ref, + Ref<[u8]>, + Ref<[u8]>, + >(), + } + } + + /// Compute the distance between `vector` and the PQ vector encoded by `codes`. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * `vector.len() != padded.dim()` + /// * `codes.len() != padded.nchunks()` + /// * Any element in `codes` is equal to or greater than `padded.ncenters()`. + /// + /// In addition, an error may be returned if `self` was created for a different + /// [`PaddedTable`]. Mixing [`VTable`]s in this way is not a safety issue, but is also + /// not guaranteed to work. + #[inline] + pub fn distance( + &self, + padded: &PaddedTable, + vector: &[f32], + codes: &[u8], + ) -> Result { + (self.distance).call(padded, vector, codes) + } + + /// Compute the distance between two compressed vectors. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * `a.len() != padded.nchunks()` + /// * `b.len() != padded.nchunks()` + /// * Any element in `a` or `b` is equal to or greater than `padded.ncenters()`. + /// + /// In addition, an error may be returned if `self` was created for a different + /// [`PaddedTable`]. Mixing [`VTable`]s in this way is not a safety issue, but is also + /// not guaranteed to work. + #[inline] + pub fn self_distance( + &self, + padded: &PaddedTable, + a: &[u8], + b: &[u8], + ) -> Result { + (self.self_distance).call(padded, a, b) + } +} + +//-----------------// +// Width Selection // +//-----------------// + +const FOUR: NonZeroUsize = NonZeroUsize::new(4).unwrap(); + +#[cfg(target_arch = "x86_64")] +const EIGHT: NonZeroUsize = NonZeroUsize::new(8).unwrap(); + +#[derive(Debug, Clone, Copy)] +enum SIMDWidth { + Four, + #[cfg(target_arch = "x86_64")] + Eight, +} + +impl SIMDWidth { + fn as_nonzero(self) -> NonZeroUsize { + match self { + Self::Four => FOUR, + #[cfg(target_arch = "x86_64")] + Self::Eight => EIGHT, + } + } +} + +/// The runtime architecture. +#[derive(Debug, Clone, Copy)] +enum RuntimeArch { + Scalar(Scalar), + #[cfg(target_arch = "x86_64")] + V3(V3), + #[cfg(target_arch = "aarch64")] + Neon(Neon), +} + +impl RuntimeArch { + fn new() -> Self { + dispatch_no_features(GetArch) + } + + #[cfg_attr( + not(target_arch = "x86_64"), + expect( + unused_variables, + reason = "the same result is returned regardless of the actual chunk size" + ) + )] + fn select_simd_width(&self, max_chunk_size: NonZeroUsize) -> SIMDWidth { + match self { + Self::Scalar(_) => SIMDWidth::Four, + // Use a smaller width if available. + #[cfg(target_arch = "x86_64")] + Self::V3(_) => match max_chunk_size.get() { + 0..=4 => SIMDWidth::Four, + _ => SIMDWidth::Eight, + }, + #[cfg(target_arch = "aarch64")] + Self::Neon(_) => SIMDWidth::Four, + } + } + + fn dispatch(&self, simd_width: SIMDWidth, metric: Metric) -> VTable { + diskann_wide::alias!(f32x4 = f32x4); + diskann_wide::alias!(f32x8 = f32x8); + + match (*self, simd_width, metric) { + (Self::Scalar(a), _, Metric::SquaredL2) => VTable::new::>>(a), + (Self::Scalar(a), _, Metric::InnerProduct) => { + VTable::new::>>(a) + } + (Self::Scalar(a), _, Metric::Cosine) => VTable::new::>>(a), + + // V3 // + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Four, Metric::SquaredL2) => { + VTable::new::>>(a) + } + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Four, Metric::InnerProduct) => { + VTable::new::>>(a) + } + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Four, Metric::Cosine) => VTable::new::>>(a), + + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Eight, Metric::SquaredL2) => { + VTable::new::>>(a) + } + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Eight, Metric::InnerProduct) => { + VTable::new::>>(a) + } + #[cfg(target_arch = "x86_64")] + (Self::V3(a), SIMDWidth::Eight, Metric::Cosine) => VTable::new::>>(a), + + // Neon // + #[cfg(target_arch = "aarch64")] + (Self::Neon(a), _, Metric::SquaredL2) => VTable::new::>>(a), + #[cfg(target_arch = "aarch64")] + (Self::Neon(a), _, Metric::InnerProduct) => VTable::new::>>(a), + #[cfg(target_arch = "aarch64")] + (Self::Neon(a), _, Metric::Cosine) => VTable::new::>>(a), + } + } +} + +#[derive(Debug, Clone, Copy)] +struct GetArch; + +impl Target for GetArch { + #[inline(always)] + fn run(self, arch: Scalar) -> RuntimeArch { + RuntimeArch::Scalar(arch) + } +} + +#[cfg(target_arch = "x86_64")] +impl Target for GetArch { + #[inline(always)] + fn run(self, arch: V3) -> RuntimeArch { + RuntimeArch::V3(arch) + } +} + +#[cfg(target_arch = "x86_64")] +impl Target for GetArch { + #[inline(always)] + fn run(self, arch: V4) -> RuntimeArch { + RuntimeArch::V3(arch.retarget()) + } +} + +#[cfg(target_arch = "aarch64")] +impl Target for GetArch { + #[inline(always)] + fn run(self, arch: Neon) -> RuntimeArch { + RuntimeArch::Neon(arch) + } +} + +//----------// +// SIMD Ops // +//----------// + +trait Op { + type Arch: Architecture; + type Accum; + type Vector: SIMDVector; + + fn init(arch: Self::Arch) -> Self::Accum; + fn accum(acc: Self::Accum, x: Self::Vector, y: Self::Vector) -> Self::Accum; + fn reduce_pair(a: Self::Accum, b: Self::Accum) -> f32; +} + +#[derive(Debug)] +struct SquaredL2(PhantomData); + +impl Op for SquaredL2 +where + V: SIMDFloat + SIMDSumTree, + V::Arch: Architecture, +{ + type Arch = V::Arch; + type Accum = V; + type Vector = V; + + fn init(arch: Self::Arch) -> V { + V::default(arch) + } + + fn accum(acc: V, x: V, y: V) -> V { + let d = x - y; + d.mul_add_simd(d, acc) + } + + fn reduce_pair(a: V, b: V) -> f32 { + (a + b).sum_tree() + } +} + +#[derive(Debug)] +struct InnerProduct(PhantomData); + +impl Op for InnerProduct +where + V: SIMDFloat + SIMDSumTree, + V::Arch: Architecture, +{ + type Arch = V::Arch; + type Accum = V; + type Vector = V; + + fn init(arch: Self::Arch) -> V { + V::default(arch) + } + + fn accum(acc: V, x: V, y: V) -> V { + x.mul_add_simd(y, acc) + } + + fn reduce_pair(a: V, b: V) -> f32 { + -(a + b).sum_tree() + } +} + +#[derive(Debug)] +struct Cosine(PhantomData); + +#[derive(Debug)] +struct CosineAccumulator { + xy: V, + xnorm: V, + ynorm: V, +} + +fn finish_cosine(xy: f32, xnorm: f32, ynorm: f32) -> f32 { + if xnorm < f32::MIN_POSITIVE || ynorm < f32::MIN_POSITIVE { + 1.0 + } else { + let v = xy / (xnorm.sqrt() * ynorm.sqrt()); + 1.0 - (-1.0f32).max(1.0f32.min(v)) + } +} + +impl Op for Cosine +where + V: SIMDFloat + SIMDSumTree, + V::Arch: Architecture, +{ + type Arch = V::Arch; + type Accum = CosineAccumulator; + type Vector = V; + + fn init(arch: Self::Arch) -> Self::Accum { + CosineAccumulator { + xy: V::default(arch), + xnorm: V::default(arch), + ynorm: V::default(arch), + } + } + + fn accum(acc: Self::Accum, x: V, y: V) -> Self::Accum { + CosineAccumulator { + xy: x.mul_add_simd(y, acc.xy), + xnorm: x.mul_add_simd(x, acc.xnorm), + ynorm: y.mul_add_simd(y, acc.ynorm), + } + } + + fn reduce_pair(a: Self::Accum, b: Self::Accum) -> f32 { + let xy = (a.xy + b.xy).sum_tree(); + let xnorm = (a.xnorm + b.xnorm).sum_tree(); + let ynorm = (a.ynorm + b.ynorm).sum_tree(); + finish_cosine(xy, xnorm, ynorm) + } +} + +#[derive(Debug, Clone, Copy)] +struct OpTarget(PhantomData); + +impl FTarget3, &PaddedTable, &[f32], &[u8]> for OpTarget +where + O: Op, +{ + #[inline(always)] + fn run(arch: O::Arch, table: &PaddedTable, a: &[f32], b: &[u8]) -> Result { + distance::(arch, table, a, b) + } +} + +impl FTarget3, &PaddedTable, &[u8], &[u8]> + for OpTarget +where + O: Op, +{ + #[inline(always)] + fn run( + arch: O::Arch, + table: &PaddedTable, + a: &[u8], + b: &[u8], + ) -> Result { + self_distance::(arch, table, a, b) + } +} + +//--------------------------------// +// Full Precision-Quant Distances // +//--------------------------------// + +/// Errors for [`VTable::distance`]. +#[derive(Debug, Error)] +#[non_exhaustive] +pub enum DistanceError { + #[error("table has dimension {} but full vector has length {}", dim, alen)] + ALen { dim: usize, alen: usize }, + #[error("table has {} chunks but codes slice has length {}", chunks, blen)] + BLen { chunks: usize, blen: usize }, + #[error("codes have a value that exceeds the number of pivots {}", ncenters)] + OutOfBounds { ncenters: usize }, + #[error("An invalid SIMD width was chosen - ensure the correct vtable is used")] + InvalidSimd, +} + +#[inline(always)] +fn distance( + arch: O::Arch, + table: &PaddedTable, + a: &[f32], + b: &[u8], +) -> Result +where + O: Op, +{ + // Check 1 + if !table + .simd_width + .as_nonzero() + .get() + .is_multiple_of(O::Vector::LANES) + { + return Err(DistanceError::InvalidSimd); + } + + // Check 2 + if a.len() != table.dim() { + return Err(DistanceError::ALen { + dim: table.dim(), + alen: a.len(), + }); + } + + // Check 3 + if b.len() != table.nchunks() { + return Err(DistanceError::BLen { + chunks: table.nchunks(), + blen: b.len(), + }); + } + + // Check 4 + if let Ok(ncenters) = u8::try_from(table.ncenters()) + && let Some(max) = b.iter().max() + && *max >= ncenters + { + return Err(DistanceError::OutOfBounds { + ncenters: table.ncenters(), + }); + } + + // All checks passed - we're good to go! + + let pivots = &table.pivots; + + let pivot_stride = pivots.ncols(); + let chunk_stride = table.pivots_per_chunk * pivot_stride; + + let nchunks = b.len(); + let mut d0 = O::init(arch); + let mut d1 = O::init(arch); + + let mut i = 0; + let mut p = pivots.as_ptr(); + + let offsets = table.offsets.as_slice(); + let aptr = a.as_ptr(); + + let lanes = O::Vector::LANES; + + // Here is our conundrum. Unrolling is nice, but there is an issue if two adjacent chunks + // have wildly different lengths. + // + // So here, we process two chunks at a time. We first tackle the common full-width prefix. + // Then we process the tails independently. + // + // SAFETY INVARIANTS (SI) + // + // 1. `offsets.len() == nchunks + 1`. Offsets are strictly increasing with + // `offsets[0] == 0` and `offsets[nchunks] == a.len()` by the table invariant and + // Check 2. + // + // 2. At the start of each outer iteration, `p` points to the beginning of pivot storage + // for chunk `i`. + // + // 3. Pivot chunk `i` begins at `pivots.as_ptr().add(i * chunk_stride)`. Within a chunk, + // pivot `j` begins at `chunk_base.add(j * pivot_stride)`. Each chunk and pivot row + // contains `chunk_stride` and `pivot_stride` initialized elements, respectively. + // + // 4. Each code selects an existing pivot. Check 4 establishes this when the number of + // pivots fits in `u8`; otherwise every possible `u8` is less than + // `pivots_per_chunk`. + // + // 5. `pivot_stride` is a multiple of the table's SIMD width, and Check 1 establishes + // that width is a multiple of `lanes`. Each pivot row is at least as long as every + // unpadded chunk rounded up to `lanes`. + while i + 2 <= nchunks { + // SAFETY: `i + 2 <= nchunks`, so offsets `i` through `i + 2` exist by SI-1. + let (o0, o1, o2) = unsafe { + ( + *offsets.get_unchecked(i), + *offsets.get_unchecked(i + 1), + *offsets.get_unchecked(i + 2), + ) + }; + + // SAFETY: `o0` and `o1` are in-bounds chunk starts in `a` by SI-1. + let (a0, a1) = unsafe { (aptr.add(o0), aptr.add(o1)) }; + + // SAFETY: By SI-(2,3,4), `p` begins chunk `i` and the code selects an in-bounds + // pivot row in that chunk. + let b0 = unsafe { p.add(pivot_stride * (*b.get_unchecked(i) as usize)) }; + + // SAFETY: `i + 1 < nchunks`; by SI-(2,3,4), this selects an in-bounds pivot row in + // the following chunk. + let b1 = unsafe { p.add(chunk_stride + pivot_stride * (*b.get_unchecked(i + 1) as usize)) }; + + let full0 = (o1 - o0) / lanes; + let full1 = (o2 - o1) / lanes; + + let common = full0.min(full1); + for j in 0..common { + // SAFETY: `j < full0` and `j < full1`, so both query loads end at or before + // their respective chunk boundaries. By SI-(3,4,5), the corresponding pivot + // loads remain within their initialized padded rows. + unsafe { + let va = O::Vector::load_simd(arch, a0.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b0.add(lanes * j)); + d0 = O::accum(d0, va, vb); + + let va = O::Vector::load_simd(arch, a1.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b1.add(lanes * j)); + d1 = O::accum(d1, va, vb); + } + } + + // Handle whatever happens to be left of `a0`. + for j in common..full0 { + // SAFETY: `j < full0`, so the query load remains within chunk 0. By SI-(3,4,5), + // the corresponding pivot load remains within its padded row. + unsafe { + let va = O::Vector::load_simd(arch, a0.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b0.add(lanes * j)); + d0 = O::accum(d0, va, vb); + } + } + + let a0_remaining = (o1 - o0) - full0 * lanes; + if a0_remaining != 0 { + // SAFETY: `full0 * lanes + a0_remaining == o1 - o0` and `a0_remaining < lanes`, + // so this reads exactly the initialized remainder. + let va = + unsafe { O::Vector::load_simd_first(arch, a0.add(lanes * full0), a0_remaining) }; + // SAFETY: `(full0 + 1) * lanes` is the chunk length rounded up to `lanes`, + // which is at most `pivot_stride` by SI-5. + let vb = unsafe { O::Vector::load_simd(arch, b0.add(lanes * full0)) }; + d0 = O::accum(d0, va, vb); + } + + // Handle whatever happens to be left of `a1`. + for j in common..full1 { + // SAFETY: Same proof as the remaining full-vector loop for chunk 0. + unsafe { + let va = O::Vector::load_simd(arch, a1.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b1.add(lanes * j)); + d1 = O::accum(d1, va, vb); + } + } + + let a1_remaining = (o2 - o1) - full1 * lanes; + if a1_remaining != 0 { + // SAFETY: Same remainder proof as chunk 0. + let va = + unsafe { O::Vector::load_simd_first(arch, a1.add(lanes * full1), a1_remaining) }; + // SAFETY: Same padded-row proof as chunk 0. + let vb = unsafe { O::Vector::load_simd(arch, b1.add(lanes * full1)) }; + d1 = O::accum(d1, va, vb); + } + + i += 2; + + // SAFETY: Before incrementing, `p` began chunk `i - 2` by SI-2. The loop bound + // established `i <= nchunks`, so advancing two chunk strides reaches chunk `i` or + // one-past the final chunk, preserving SI-2. + p = unsafe { p.add(2 * chunk_stride) }; + } + + if i < nchunks { + debug_assert!(i + 1 == nchunks); + + // SAFETY: `i < nchunks`, so offsets `i` and `i + 1` exist by SI-1. + let (o0, o1) = unsafe { (*offsets.get_unchecked(i), *offsets.get_unchecked(i + 1)) }; + + // SAFETY: `o0` is an in-bounds chunk start in `a` by SI-1. + let a0 = unsafe { aptr.add(o0) }; + + // SAFETY: By SI-(2,3,4), this selects an in-bounds pivot row in chunk `i`. + let b0 = unsafe { p.add(pivot_stride * (*b.get_unchecked(i) as usize)) }; + + // The number of unprocessed elements. + let full = (o1 - o0) / lanes; + for j in 0..full { + // SAFETY: `j < full`, so the query load remains within the chunk. By SI-(3,4,5), + // the corresponding pivot load remains within its padded row. + unsafe { + let va = O::Vector::load_simd(arch, a0.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b0.add(lanes * j)); + d0 = O::accum(d0, va, vb); + } + } + + let remaining = (o1 - o0) - full * lanes; + if remaining != 0 { + // SAFETY: `full * lanes + remaining == o1 - o0` and `remaining < lanes`, so + // this reads exactly the initialized remainder. + let va = unsafe { O::Vector::load_simd_first(arch, a0.add(lanes * full), remaining) }; + // SAFETY: `(full + 1) * lanes` is at most `pivot_stride` by SI-5. + let vb = unsafe { O::Vector::load_simd(arch, b0.add(lanes * full)) }; + d0 = O::accum(d0, va, vb); + } + } + + Ok(O::reduce_pair(d0, d1)) +} + +//-----------------------// +// Quant-Quant Distances // +//-----------------------// + +/// Errors for [`VTable::self_distance`]. +#[derive(Debug, Error)] +#[non_exhaustive] +pub enum SelfDistanceError { + #[error("table has {} chunks but codes slice has length {}", chunks, len)] + Len { chunks: usize, len: usize }, + #[error("codes have a value that exceeds the number of pivots {}", ncenters)] + OutOfBounds { ncenters: usize }, + #[error("An invalid SIMD width was chosen - ensure the correct vtable is used")] + InvalidSimd, +} + +#[inline(always)] +fn self_distance( + arch: O::Arch, + table: &PaddedTable, + a: &[u8], + b: &[u8], +) -> Result +where + O: Op, +{ + // Check 1 + if !table + .simd_width + .as_nonzero() + .get() + .is_multiple_of(O::Vector::LANES) + { + return Err(SelfDistanceError::InvalidSimd); + } + + let chunks = table.nchunks(); + // Check 2 + if a.len() != chunks { + return Err(SelfDistanceError::Len { + chunks, + len: a.len(), + }); + } + + // Check 3 + if b.len() != chunks { + return Err(SelfDistanceError::Len { + chunks, + len: b.len(), + }); + } + + // Check 4 + if let Ok(ncenters) = u8::try_from(table.ncenters()) { + if let Some(max) = a.iter().max() + && *max >= ncenters + { + return Err(SelfDistanceError::OutOfBounds { + ncenters: ncenters.into(), + }); + } + + if let Some(max) = b.iter().max() + && *max >= ncenters + { + return Err(SelfDistanceError::OutOfBounds { + ncenters: ncenters.into(), + }); + } + } + + // All checks passed - we're good to go! + + let pivots = &table.pivots; + + // The number of SIMD steps to process for each pivot. + let steps = pivots.ncols() / O::Vector::LANES; + + let pivot_stride = pivots.ncols(); + let chunk_stride = table.pivots_per_chunk * pivot_stride; + + let len = a.len(); + let mut d0 = O::init(arch); + let mut d1 = O::init(arch); + + let mut p = pivots.as_ptr(); + + let lanes = O::Vector::LANES; + let mut i = 0; + + // SAFETY INVARIANTS (SI) + // + // 1. Both code slices contain exactly `len == table.nchunks()` entries by Checks 2 + // and 3. + // + // 2. At the start of each outer iteration, `p` points to the beginning of pivot storage + // for chunk `i`. + // + // 3. Pivot chunk `i` begins at `pivots.as_ptr().add(i * chunk_stride)`. Within a chunk, + // pivot `j` begins at `chunk_base.add(j * pivot_stride)`. Each chunk and pivot row + // contains `chunk_stride` and `pivot_stride` initialized elements, respectively. + // + // 4. Every code selects an existing pivot. Check 4 establishes this when the number of + // pivots fits in `u8`; otherwise every possible `u8` is in bounds. + // + // 5. `pivot_stride` is a multiple of the table's SIMD width, and Check 1 establishes + // that width is a multiple of `lanes`. Therefore + // `steps * lanes == pivot_stride`. + while i + 2 <= len { + // SAFETY: `i + 2 <= len`; by SI-(1,2,3,4), `p` begins chunk `i` and this code + // selects an in-bounds pivot row in that chunk. + let a0 = unsafe { p.add(pivot_stride * (*a.get_unchecked(i) as usize)) }; + + // SAFETY: By the same invariants, this selects an in-bounds pivot row in chunk + // `i + 1`. + let a1 = unsafe { p.add(chunk_stride + pivot_stride * (*a.get_unchecked(i + 1) as usize)) }; + + // SAFETY: Same proof as `a0`; SI-1 establishes that `b[i]` exists. + let b0 = unsafe { p.add(pivot_stride * (*b.get_unchecked(i) as usize)) }; + // SAFETY: Same proof as `a1`. + let b1 = unsafe { p.add(chunk_stride + pivot_stride * (*b.get_unchecked(i + 1) as usize)) }; + + for j in 0..steps { + // SAFETY: `j < steps` and `steps * lanes == pivot_stride` by SI-5, so every + // full-vector load remains within its selected initialized pivot row. + unsafe { + // Unroll 0 + let va = O::Vector::load_simd(arch, a0.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b0.add(lanes * j)); + d0 = O::accum(d0, va, vb); + + // Unroll 1 + let va = O::Vector::load_simd(arch, a1.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b1.add(lanes * j)); + d1 = O::accum(d1, va, vb); + } + } + + i += 2; + + // SAFETY: Before incrementing, `p` began chunk `i - 2` by SI-2. The loop bound + // established `i <= len`, so advancing two chunk strides reaches chunk `i` or + // one-past the final chunk, preserving SI-2. + p = unsafe { p.add(2 * chunk_stride) }; + } + + if i < len { + debug_assert!(i + 1 == len); + + // SAFETY: By SI-(1,2,3,4), `p` begins chunk `i` and this selects an in-bounds + // pivot row. + let a0 = unsafe { p.add(pivot_stride * (*a.get_unchecked(i) as usize)) }; + // SAFETY: Same proof for `b[i]`. + let b0 = unsafe { p.add(pivot_stride * (*b.get_unchecked(i) as usize)) }; + + for j in 0..steps { + // SAFETY: `j < steps`, so both loads remain within their initialized pivot rows + // by SI-5. + unsafe { + let va = O::Vector::load_simd(arch, a0.add(lanes * j)); + let vb = O::Vector::load_simd(arch, b0.add(lanes * j)); + d0 = O::accum(d0, va, vb); + } + } + } + + Ok(O::reduce_pair(d0, d1)) +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use diskann_utils::assert_contains; + + use crate::{ + product::tables::test::{self as table_test, DistanceTestTable, QueryLike, SelfLike}, + test_util::Check, + }; + + fn cases() -> &'static [Case] { + const CASES: &[Case] = &[ + // chunk dimensions, pivots, start + Case::new(&[1], 1, 0.0, Check::exact()), + Case::new(&[1, 1, 1], 2, -1.0, Check::exact()), + Case::new(&[3, 2, 2], 15, -7.0, Check::exact()), + Case::new(&[2, 2, 2, 2], 16, -8.0, Check::exact()), + Case::new(&[3, 3, 3, 2, 2], 17, -11.0, Check::exact()), + Case::new(&[8, 8, 8, 8], 32, -16.0, Check::exact()), + Case::new(&[6, 6, 5, 5, 5, 5, 5], 33, -20.0, Check::exact()), + Case::new(&[13, 12, 12], 33, -20.0, Check::exact()), + Case::new(&[4, 4, 3, 3, 3], 256, -128.0, Check::exact()), + // Exercise sharply imbalanced chunks in either position of an unrolled pair, + // followed by an odd chunk that must use the remainder loop. + Case::new(&[1, 23, 23], 17, -9.0, Check::exact()), + Case::new(&[23, 1, 23], 17, -9.0, Check::exact()), + Case::new(&[1, 17, 2, 9, 3], 17, -9.0, Check::exact()), + ]; + CASES + } + + #[derive(Debug, Clone, Copy)] + struct Case { + chunk_dims: &'static [usize], + pivots: usize, + start: f32, + check: Check, + } + + impl Case { + const fn new( + chunk_dims: &'static [usize], + pivots: usize, + start: f32, + check: Check, + ) -> Self { + Self { + chunk_dims, + pivots, + start, + check, + } + } + } + + fn run_test( + cases: &[Case], + metric: Metric, + arch: Option, + reference: &dyn Fn(&[f32], &[f32]) -> f32, + ctx: &dyn std::fmt::Display, + ) { + let (num_queries, num_trials) = if cfg!(miri) { + // The driver will run some directed tests even if there are no regular random + // trials. + (1, 0) + } else { + (10, 10) + }; + + #[derive(Debug)] + struct Dut<'a> { + table: &'a PaddedTable, + query: Vec, + vtable: VTable, + } + + impl QueryLike for Dut<'_> { + fn preprocess(&mut self, query: &[f32]) { + self.query.clear(); + self.query.extend_from_slice(query); + } + + fn evaluate(&mut self, code: &[u8]) -> f32 { + self.vtable.distance(self.table, &self.query, code).unwrap() + } + } + + impl SelfLike for Dut<'_> { + fn evaluate(&mut self, a: &[u8], b: &[u8]) -> f32 { + self.vtable.self_distance(self.table, a, b).unwrap() + } + } + + for Case { + chunk_dims, + pivots, + start, + check, + } in cases.iter().copied() + { + let driver = DistanceTestTable::from_chunk_dims(chunk_dims, pivots, start); + let dim = driver.dim(); + let chunks = driver.chunks(); + let basic = driver.basic_table(); + + let table = match arch { + Some(arch) => PaddedTable::from_basic_with(basic.as_view(), arch), + None => PaddedTable::from_basic(basic.as_view()), + }; + + let vtable = table.vtable(metric); + + let mut dut = Dut { + table: &table, + query: Vec::new(), + vtable, + }; + + driver.drive_query_like( + num_queries, + num_trials, + &mut driver.rng(0xc0ffee), + check, + reference, + &mut dut, + format_args!( + "[{}] padded table - dim = {}, chunks = {}, pivots = {}", + ctx, dim, chunks, pivots + ), + ); + + driver.drive_self_like( + num_trials, + &mut driver.rng(0xc0ffee), + check, + reference, + &mut dut, + format_args!( + "[{}] padded table - dim = {}, chunks = {}, pivots = {}", + ctx, dim, chunks, pivots + ), + ); + } + } + + // L2 - query-like + #[test] + fn test_l2_query_like() { + run_test( + cases(), + Metric::SquaredL2, + None, + &table_test::squared_l2, + &"squared l2 full x quant - auto-detect", + ); + + run_test( + cases(), + Metric::SquaredL2, + Some(RuntimeArch::Scalar(Scalar::new())), + &table_test::squared_l2, + &"squared l2 full x quant - scalar", + ); + + #[cfg(target_arch = "x86_64")] + if let Some(arch) = V3::new_checked() { + run_test( + cases(), + Metric::SquaredL2, + Some(RuntimeArch::V3(arch)), + &table_test::squared_l2, + &"squared l2 full x quant - V3", + ); + } + + #[cfg(target_arch = "aarch64")] + if let Some(arch) = Neon::new_checked() { + run_test( + cases(), + Metric::SquaredL2, + Some(RuntimeArch::Neon(arch)), + &table_test::squared_l2, + &"squared l2 full x quant - V3", + ); + } + } + + // Inner Product - query-like + #[test] + fn test_inner_product_query_like() { + run_test( + cases(), + Metric::InnerProduct, + None, + &table_test::inner_product, + &"inner-product full x quant - auto-detect", + ); + + run_test( + cases(), + Metric::InnerProduct, + Some(RuntimeArch::Scalar(Scalar::new())), + &table_test::inner_product, + &"inner-product full x quant - scalar", + ); + + #[cfg(target_arch = "x86_64")] + if let Some(arch) = V3::new_checked() { + run_test( + cases(), + Metric::InnerProduct, + Some(RuntimeArch::V3(arch)), + &table_test::inner_product, + &"inner-product full x quant - V3", + ); + } + + #[cfg(target_arch = "aarch64")] + if let Some(arch) = Neon::new_checked() { + run_test( + cases(), + Metric::InnerProduct, + Some(RuntimeArch::Neon(arch)), + &table_test::inner_product, + &"inner-product full x quant - V3", + ); + } + } + + // Cosine - query-like + #[test] + fn test_cosine_query_like() { + run_test( + cases(), + Metric::Cosine, + None, + &table_test::cosine, + &"cosine full x quant - auto-detect", + ); + + run_test( + cases(), + Metric::Cosine, + Some(RuntimeArch::Scalar(Scalar::new())), + &table_test::cosine, + &"cosine full x quant - scalar", + ); + + #[cfg(target_arch = "x86_64")] + if let Some(arch) = V3::new_checked() { + run_test( + cases(), + Metric::Cosine, + Some(RuntimeArch::V3(arch)), + &table_test::cosine, + &"cosine full x quant - V3", + ); + } + + #[cfg(target_arch = "aarch64")] + if let Some(arch) = Neon::new_checked() { + run_test( + cases(), + Metric::Cosine, + Some(RuntimeArch::Neon(arch)), + &table_test::cosine, + &"cosine full x quant - V3", + ); + } + } + + //////////// + // Errors // + //////////// + + /// A table with dimension 7, 3 chunks, and 3 pivots per chunk. + fn error_table() -> (PaddedTable, VTable) { + let driver = DistanceTestTable::new(7, 3, 3, 0.0); + let basic = driver.basic_table(); + let table = + PaddedTable::from_basic_with(basic.as_view(), RuntimeArch::Scalar(Scalar::new())); + let vtable = table.vtable(Metric::SquaredL2); + (table, vtable) + } + + #[test] + fn test_distance_errors() { + let (table, vtable) = error_table(); + let vector = vec![0.0; table.dim()]; + let codes = vec![0; table.nchunks()]; + + for len in [table.dim() - 1, table.dim() + 1] { + let invalid = vec![0.0; len]; + let err = vtable.distance(&table, &invalid, &codes).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "table has dimension {} but full vector has length {len}", + table.dim() + ) + ); + } + + for len in [table.nchunks() - 1, table.nchunks() + 1] { + let invalid = vec![0; len]; + let err = vtable.distance(&table, &vector, &invalid).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "table has {} chunks but codes slice has length {len}", + table.nchunks() + ) + ); + } + + let mut invalid = codes; + invalid[1] = u8::try_from(table.ncenters()).unwrap(); + let err = vtable.distance(&table, &vector, &invalid).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "codes have a value that exceeds the number of pivots {}", + table.ncenters() + ) + ); + } + + #[test] + fn test_self_distance_errors() { + let (table, vtable) = error_table(); + let codes = vec![0; table.nchunks()]; + + for len in [table.nchunks() - 1, table.nchunks() + 1] { + let invalid = vec![0; len]; + + let err = vtable.self_distance(&table, &invalid, &codes).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "table has {} chunks but codes slice has length {len}", + table.nchunks() + ) + ); + + let err = vtable.self_distance(&table, &codes, &invalid).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "table has {} chunks but codes slice has length {len}", + table.nchunks() + ) + ); + } + + for invalid_operand in 0..2 { + let mut a = codes.clone(); + let mut b = codes.clone(); + let invalid = if invalid_operand == 0 { &mut a } else { &mut b }; + invalid[1] = u8::try_from(table.ncenters()).unwrap(); + + let err = vtable.self_distance(&table, &a, &b).unwrap_err(); + assert_contains!( + err.to_string(), + format!( + "codes have a value that exceeds the number of pivots {}", + table.ncenters() + ) + ); + } + } + + #[test] + fn test_mismatched_simd_width_errors() { + diskann_wide::alias!(f32x8 = f32x8); + + let (table, _) = error_table(); + let invalid_vtable = VTable::new::>>(Scalar::new()); + let vector = vec![0.0; table.dim()]; + let codes = vec![0; table.nchunks()]; + + let err = invalid_vtable + .distance(&table, &vector, &codes) + .unwrap_err(); + assert_contains!(err.to_string(), "An invalid SIMD width was chosen"); + + let err = invalid_vtable + .self_distance(&table, &codes, &codes) + .unwrap_err(); + assert_contains!(err.to_string(), "An invalid SIMD width was chosen"); + } +} diff --git a/diskann-quantization/src/product/tables/test.rs b/diskann-quantization/src/product/tables/test.rs index 2b7cfeb51e..791cc2319c 100644 --- a/diskann-quantization/src/product/tables/test.rs +++ b/diskann-quantization/src/product/tables/test.rs @@ -4,16 +4,346 @@ */ // A collection of test helpers to ensure uniformity across tables. + +use std::num::NonZeroUsize; + use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut}; +use diskann_vector::{PureDistanceFunction, distance}; #[cfg(not(miri))] use rand::seq::IndexedRandom; use rand::{ Rng, SeedableRng, distr::{Distribution, Uniform}, + rngs::StdRng, +}; + +use crate::{ + product::tables::BasicTable, + test_util::Check, + traits::CompressInto, + views::{self, ChunkOffsets, ChunkOffsetsView}, }; -use crate::traits::CompressInto; -use crate::views::{self, ChunkOffsets, ChunkOffsetsView}; +////////////////////// +// Distance Helpers // +////////////////////// + +/// To test the implementation of distances, we need a way to seed the source pivot table +/// with known contents. +/// +/// The layout of the pivot table will look like this: +/// +/// chunk 0 chunk 1 ... chunk K +/// +/// | S S ... | S+1 S+1 ... | ... | S+K S+K ... | pivot 0 +/// | S+1 S+1 ... | S+2 S+2 ... | ... | S+K+1 S+K+1 ... | pivot 1 +/// | S+2 S+2 ... | S+3 S+3 ... | ... | S+K+2 S+K+2 ... | pivot 2 +/// | ... | ... | ... | ... | ... +/// | S+N S+N ... | S+N+1 S+N+1 ... | ... | S+K+N S+K+N ... | pivot N +/// +/// where +/// +/// * S: The configured start value for chunk 0, pivot 0 (i.e., [`Self::start`]) +/// * K + 1: The number of PQ chunks ([`Self::chunks`]). +/// * N + 1: The number of PQ pivots ([`Self::pivots`]). +#[derive(Debug, Clone)] +pub(super) struct DistanceTestTable { + /// The chunking schema. + pub(super) offsets: ChunkOffsets, + /// The number of pivots per chunk. + pub(super) pivots: usize, + /// The starting value for chunk 0, pivot 0. + pub(super) start: f32, +} + +/// The position within the chunking scheme. +#[derive(Debug, Clone, Copy)] +struct Location { + /// The chunk number. + chunk: usize, + /// The pivot. + pivot: usize, +} + +#[derive(Debug, Clone)] +pub(super) struct UniformFloat(Uniform); + +impl UniformFloat { + pub(super) fn new(low: usize, high: usize) -> Result { + Uniform::new(low, high).map(Self) + } +} + +impl Distribution for UniformFloat { + fn sample(&self, rng: &mut R) -> f32 { + self.0.sample(rng) as f32 + } +} + +type DriveFn<'a> = &'a mut (dyn FnMut(&[u8], &[f32], std::fmt::Arguments<'_>) + 'a); + +impl DistanceTestTable { + pub(super) fn new(dim: usize, chunks: usize, pivots: usize, start: f32) -> Self { + Self { + offsets: ChunkOffsets::partition( + NonZeroUsize::new(dim).unwrap(), + NonZeroUsize::new(chunks).unwrap(), + ) + .unwrap(), + pivots, + start, + } + } + + pub(super) fn from_chunk_dims(chunk_dims: &[usize], pivots: usize, start: f32) -> Self { + let mut offsets = Vec::with_capacity(chunk_dims.len() + 1); + let mut offset = 0usize; + offsets.push(offset); + for &dim in chunk_dims { + offset = offset.checked_add(dim).unwrap(); + offsets.push(offset); + } + + Self { + offsets: ChunkOffsets::new(offsets.into_boxed_slice()).unwrap(), + pivots, + start, + } + } + + /// This is mainly a convenience so we don't always have to import `StdRng` and + /// `SeedableRng` and all that jazz. + pub(super) fn rng(&self, seed: u64) -> StdRng { + StdRng::seed_from_u64(seed) + } + + pub(super) fn offsets(&self) -> ChunkOffsetsView<'_> { + self.offsets.as_view() + } + + pub(super) fn chunks(&self) -> usize { + self.offsets.len() + } + + pub(super) fn dim(&self) -> usize { + self.offsets.dim() + } + + pub(super) fn pivots(&self) -> usize { + self.pivots + } + + fn value(&self, loc: Location) -> f32 { + (loc.chunk + loc.pivot) as f32 + self.start + } + + pub(super) fn basic_table(&self) -> BasicTable { + // This creates a base vector like + // | chunk 0 | chunk 1 | ... | chunk K | + // | 0 0 ... 0 | 1 1 ... 1 | ... | K K ... K | + let mut base = Vec::::new(); + for i in 0..self.offsets.len() { + let v = (i as f32) + self.start; + for _ in self.offsets.at(i) { + base.push(v); + } + } + + // Use our base vector to build the rest of the pivot matrix. + let pivots = rowmajor::Owned::from_fn(self.pivots(), self.dim(), |rc| { + (rc.row as f32) + base[rc.col] + }); + + BasicTable::new(pivots, self.offsets().to_owned()).unwrap() + } + + pub(super) fn expected_vector_into(&self, v: &mut [f32], codes: &[u8]) { + assert_eq!(v.len(), self.dim()); + assert_eq!(codes.len(), self.chunks()); + + let mut i = 0; + for (chunk, pivot) in codes.iter().copied().enumerate() { + let pivot = usize::from(pivot); + assert!(pivot < self.pivots()); + let loc = Location { chunk, pivot }; + + for _ in self.offsets.at(chunk) { + v[i] = self.value(loc); + i += 1; + } + } + } + + pub(super) fn drive( + &self, + num_trials: usize, + rng: &mut StdRng, + f: DriveFn<'_>, + ctx: std::fmt::Arguments<'_>, + ) { + // Run two fixed trials - one with all zeros and one with the max setting. + // + // Then we perform random trials. + let mut codes = vec![0u8; self.chunks()]; + let mut vector = vec![0.0; self.dim()]; + self.expected_vector_into(&mut vector, &codes); + f(&codes, &mut vector, format_args!("{ctx}, all zeros")); + + let max = u8::try_from(self.pivots() - 1).unwrap(); + codes.fill(max); + self.expected_vector_into(&mut vector, &codes); + f(&codes, &mut vector, format_args!("{ctx}, all {max}")); + + // Begin random trials. + let dist = Uniform::new(0, self.pivots()).unwrap(); + for trial in 0..num_trials { + codes + .iter_mut() + .for_each(|c| *c = u8::try_from(dist.sample(rng)).unwrap()); + self.expected_vector_into(&mut vector, &codes); + f( + &codes, + &mut vector, + format_args!("{ctx}, trial {} of {}", trial + 1, num_trials), + ); + } + } + + pub(super) fn drive_unary( + &self, + num_trials: usize, + rng: &mut StdRng, + check: Check, + reference: &dyn Fn(&[f32]) -> f32, + dut: &mut dyn FnMut(&[u8]) -> f32, + ctx: std::fmt::Arguments<'_>, + ) { + let mut f = |codes: &[u8], vector: &[f32], ctx: std::fmt::Arguments<'_>| { + let expected = reference(vector); + let got = dut(codes); + + if let Err(reason) = check.check(got, expected) { + panic!("Check failed: {} -- {}", reason, ctx); + } + }; + + self.drive(num_trials, rng, &mut f, ctx) + } + + #[expect(clippy::too_many_arguments, reason = "this is a test function")] + pub(super) fn drive_query_like( + &self, + num_queries: usize, + num_trials: usize, + rng: &mut StdRng, + check: Check, + f: &dyn Fn(&[f32], &[f32]) -> f32, + dut: &mut dyn QueryLike, + ctx: std::fmt::Arguments<'_>, + ) { + let dist = UniformFloat::new(0, self.chunks() + self.pivots()).unwrap(); + let mut query = vec![0.0f32; self.dim()]; + for trial in 0..num_queries { + query.iter_mut().for_each(|q| *q = dist.sample(rng)); + + dut.preprocess(&query); + self.drive_unary( + num_trials, + rng, + check, + &|vector: &[f32]| f(&query, vector), + &mut |code| dut.evaluate(code), + format_args!("{ctx}, query {} of {}", trial + 1, num_queries), + ) + } + } + + pub(super) fn drive_self_like( + &self, + num_trials: usize, + rng: &mut StdRng, + check: Check, + f: &dyn Fn(&[f32], &[f32]) -> f32, + dut: &mut dyn SelfLike, + ctx: std::fmt::Arguments<'_>, + ) { + let mut run_check = |lhs_code: &[u8], + lhs_vector: &[f32], + rng: &mut StdRng, + ctx: std::fmt::Arguments<'_>| { + self.drive( + num_trials, + rng, + &mut |rhs_code: &[u8], rhs_vector: &[f32], ctx: std::fmt::Arguments<'_>| { + let expected = f(lhs_vector, rhs_vector); + let got = dut.evaluate(lhs_code, rhs_code); + + if let Err(reason) = check.check(got, expected) { + panic!("Check failed: {} -- {}", reason, ctx); + } + }, + ctx, + ); + }; + + let mut a_code = vec![0u8; self.chunks()]; + let mut a_vector = vec![0.0; self.dim()]; + self.expected_vector_into(&mut a_vector, &a_code); + + // Test with all zeros. + run_check( + &a_code, + &a_vector, + rng, + format_args!("{ctx}, all-zeros lhs"), + ); + + // Test with all max. + let max = u8::try_from(self.pivots() - 1).unwrap(); + a_code.fill(max); + self.expected_vector_into(&mut a_vector, &a_code); + run_check( + &a_code, + &a_vector, + rng, + format_args!("{ctx}, all-{max} lhs"), + ); + + // Begin random trials. + let dist = Uniform::new(0, self.pivots()).unwrap(); + for _ in 0..num_trials { + a_code + .iter_mut() + .for_each(|c| *c = u8::try_from(dist.sample(rng)).unwrap()); + self.expected_vector_into(&mut a_vector, &a_code); + + run_check(&a_code, &a_vector, rng, ctx); + } + } +} + +pub(super) fn squared_l2(x: &[f32], y: &[f32]) -> f32 { + distance::SquaredL2::evaluate(x, y) +} + +pub(super) fn inner_product(x: &[f32], y: &[f32]) -> f32 { + distance::InnerProduct::evaluate(x, y) +} + +pub(super) fn cosine(x: &[f32], y: &[f32]) -> f32 { + distance::Cosine::evaluate(x, y) +} + +/// A trait modeling query-like distances with split pre-processing and evaluation. +pub(super) trait QueryLike { + fn preprocess(&mut self, query: &[f32]); + fn evaluate(&mut self, code: &[u8]) -> f32; +} + +/// A trait modeling self-like distances with split pre-processing and evaluation. +pub(super) trait SelfLike { + fn evaluate(&mut self, a: &[u8], b: &[u8]) -> f32; +} ///////////////////////// // Compression Helpers // diff --git a/diskann-quantization/src/product/tables/transposed/mod.rs b/diskann-quantization/src/product/tables/transposed/mod.rs index ffadb6f24d..582adb7dba 100644 --- a/diskann-quantization/src/product/tables/transposed/mod.rs +++ b/diskann-quantization/src/product/tables/transposed/mod.rs @@ -7,3 +7,215 @@ mod pivots; mod table; pub use table::TransposedTable; + +/////////// +// Tests // +/////////// + +/// These tests check the distance formulation as a result of pre-processing in the transposed +/// table. They ensure we have an end-to-end working example of full distance calculations. +/// +/// The tests are broken into metric-specific tests, mainly so they can run more efficiently +/// in parallel as there are a decent number of cases that must be covered. +#[cfg(test)] +mod tests { + use super::*; + + use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut}; + use diskann_vector::{Norm, norm::FastL2Norm}; + + use crate::{ + distances, + product::tables::{ + lookup::{self, DotAndNorm}, + test::{self as table_test, DistanceTestTable, QueryLike}, + }, + test_util::Check, + }; + + /// Common cases between all distances. + fn cases() -> &'static [Case] { + const CASES: &[Case] = &[ + // dim, chunks, pivots, start + Case::new(1, 1, 1, 0.0, Check::exact()), + Case::new(3, 3, 2, -1.0, Check::exact()), + Case::new(7, 3, 15, -7.0, Check::exact()), + Case::new(8, 4, 16, -8.0, Check::exact()), + Case::new(13, 5, 17, -11.0, Check::exact()), + Case::new(32, 4, 32, -16.0, Check::exact()), + Case::new(37, 7, 33, -20.0, Check::exact()), + Case::new(17, 5, 256, -128.0, Check::exact()), + ]; + CASES + } + + #[derive(Debug, Clone, Copy)] + struct Case { + dim: usize, + chunks: usize, + pivots: usize, + start: f32, + check: Check, + } + + impl Case { + const fn new(dim: usize, chunks: usize, pivots: usize, start: f32, check: Check) -> Self { + Self { + dim, + chunks, + pivots, + start, + check, + } + } + } + + fn run_test( + cases: &[Case], + create: &dyn Fn(&TransposedTable) -> Box, + reference: &dyn Fn(&[f32], &[f32]) -> f32, + ctx: &dyn std::fmt::Display, + ) { + let (num_queries, num_trials) = if cfg!(miri) { + // The driver will run some directed tests even if there are no regular random + // trials. + (1, 0) + } else { + (10, 10) + }; + + for Case { + dim, + chunks, + pivots, + start, + check, + } in cases.iter().copied() + { + let driver = DistanceTestTable::new(dim, chunks, pivots, start); + let basic = driver.basic_table(); + let transposed = + TransposedTable::from_parts(basic.view_pivots(), basic.view_offsets().to_owned()) + .unwrap(); + + let mut dut = create(&transposed); + + driver.drive_query_like( + num_queries, + num_trials, + &mut driver.rng(0xc0ffee), + check, + reference, + &mut *dut, + format_args!( + "[{}] transposed table - dim = {}, chunks = {}, pivots = {}", + ctx, dim, chunks, pivots + ), + ); + } + } + + // L2 + #[test] + fn test_l2() { + #[derive(Debug)] + struct Dut<'a> { + table: &'a TransposedTable, + lut: rowmajor::Owned, + } + + impl QueryLike for Dut<'_> { + fn preprocess(&mut self, query: &[f32]) { + self.table + .process_into::(query, self.lut.as_view_mut()) + } + + fn evaluate(&mut self, code: &[u8]) -> f32 { + lookup::lookup_single(lookup::Sum, self.lut.as_view(), code).unwrap() + } + } + + run_test( + cases(), + &|table: &TransposedTable| { + let lut = rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0f32); + Box::new(Dut { table, lut }) + }, + &table_test::squared_l2, + &"squared l2", + ); + } + + // IP + #[test] + fn test_ip() { + #[derive(Debug)] + struct Dut<'a> { + table: &'a TransposedTable, + lut: rowmajor::Owned, + } + + impl QueryLike for Dut<'_> { + fn preprocess(&mut self, query: &[f32]) { + self.table + .process_into::(query, self.lut.as_view_mut()) + } + + fn evaluate(&mut self, code: &[u8]) -> f32 { + lookup::lookup_single(lookup::Sum, self.lut.as_view(), code).unwrap() + } + } + + run_test( + cases(), + &|table: &TransposedTable| { + let lut = rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0f32); + Box::new(Dut { table, lut }) + }, + &table_test::inner_product, + &"inner product", + ); + } + + // Cosine + #[test] + fn test_cosine() { + #[derive(Debug)] + struct Dut<'a> { + table: &'a TransposedTable, + lut: rowmajor::Owned, + query_norm: f32, + } + + impl QueryLike for Dut<'_> { + fn preprocess(&mut self, query: &[f32]) { + self.query_norm = (FastL2Norm).evaluate(query); + self.table + .process_into::(query, self.lut.as_view_mut()) + } + + fn evaluate(&mut self, code: &[u8]) -> f32 { + let partial = lookup::lookup_single(lookup::Sum, self.lut.as_view(), code).unwrap(); + partial.finish_cosine(self.query_norm).into_inner() + } + } + + run_test( + cases(), + &|table: &TransposedTable| { + let lut = rowmajor::Owned::from_element( + table.nchunks(), + table.ncenters(), + DotAndNorm::default(), + ); + Box::new(Dut { + table, + lut, + query_norm: 0.0f32, + }) + }, + &table_test::cosine, + &"cosine", + ); + } +} diff --git a/diskann-quantization/src/product/tables/transposed/pivots.rs b/diskann-quantization/src/product/tables/transposed/pivots.rs index 8909a73f7c..e9b0c4fa20 100644 --- a/diskann-quantization/src/product/tables/transposed/pivots.rs +++ b/diskann-quantization/src/product/tables/transposed/pivots.rs @@ -6,16 +6,18 @@ use std::fmt; use diskann_utils::strided::Strided; -use diskann_wide::{SIMDMask, SIMDMulAdd, SIMDPartialOrd, SIMDSelect, SIMDVector}; +use diskann_wide::{LoHi, SIMDMask, SIMDMulAdd, SIMDPartialOrd, SIMDSelect, SIMDVector}; use crate::{ algorithms::kmeans, - distances::{InnerProduct, SquaredL2}, + distances::{Cosine, InnerProduct, SquaredL2}, multi_vector::BlockTransposed, + product::tables::lookup::DotAndNorm, }; // The `Wide` type used as the group granularity for `Chunk`. diskann_wide::alias!(f32s = f32x8); +diskann_wide::alias!(f32x16 = f32x16); diskann_wide::alias!(u32s = u32x8); /// Error types returned by Chunk construction. @@ -187,6 +189,11 @@ impl Chunk { self.data.remainder() } + /// Return an iterator over the square norms. + pub(super) fn square_norm_chunks(&self) -> (&[[f32; 16]], &[f32]) { + self.square_norms.as_chunks::<16>() + } + /// Retrieve the value originally stored in `(row, col)` of the input matrix. /// /// # Panics @@ -982,7 +989,7 @@ impl ComputeKernel for InnerProductMathematical { /// * [`SquaredL2`]: Compute the squared l2 distance between `from` and all pivots. /// * [`InnerProduct`]: Compute the inner product (as a [`diskann_vector::SimilarityScore`]) /// between `from` and all pivots. -pub trait ProcessInto { +pub trait ProcessInto { /// Do the specified operation. /// /// # Panics @@ -992,10 +999,10 @@ pub trait ProcessInto { /// of dimensions as the pivots stored in `chunk`. /// * `into.len() != chunk.num_centers()`: This routine will produce one result per /// pivot and `into` must be sized accordingly. - fn process_into(chunk: &Chunk, from: &[f32], into: &mut [f32]); + fn process_into(chunk: &Chunk, from: &[f32], into: &mut [T]); } -impl ProcessInto for T +impl ProcessInto for T where T: ComputeKernel, { @@ -1055,6 +1062,60 @@ where } } +impl ProcessInto for Cosine { + fn process_into(chunk: &Chunk, from: &[f32], into: &mut [DotAndNorm]) { + assert_eq!(from.len(), chunk.dimension(), "incorrect input vector dim"); + assert_eq!( + into.len(), + chunk.num_centers(), + "incorrect output vector dim" + ); + + let (norm_chunks, norm_remainder) = chunk.square_norm_chunks(); + let (into_chunks, into_remainder) = into.as_chunks_mut::<16>(); + + debug_assert_eq!( + norm_chunks.len(), + into_chunks.len(), + "Check 1 already proves this" + ); + + debug_assert_eq!( + norm_remainder.len(), + into_remainder.len(), + "Check 1 already proves this" + ); + + // NOTE: This code generated for constructing `DotAndNorm` from the computed + // dot products and norms is not particularly efficient. + // + // It's not terrible, but LLVM refuses to implement a shuffle on its own. + + for (block, (into, norms)) in + std::iter::zip(into_chunks.iter_mut(), norm_chunks.iter()).enumerate() + { + let (lo, hi) = chunk.compute_in_block::(from, block); + + let dot: f32x16 = LoHi { lo, hi }.join(); + + std::iter::zip(into.iter_mut(), norms.iter()) + .zip(dot.to_array()) + .for_each(|((i, n), dot)| *i = DotAndNorm::new(dot, *n)); + } + + // Process the remainder. + if !into_remainder.is_empty() { + let (lo, hi) = + chunk.compute_in_block::(from, chunk.full_blocks()); + let dot: f32x16 = LoHi { lo, hi }.join(); + + std::iter::zip(into_remainder.iter_mut(), norm_remainder.iter()) + .zip(dot.to_array()) + .for_each(|((i, n), dot)| *i = DotAndNorm::new(dot, *n)); + } + } +} + /////////// // Tests // /////////// @@ -1065,7 +1126,7 @@ mod tests { lazy_format, views::{self, rowmajor::Matrix}, }; - use diskann_vector::{PureDistanceFunction, distance}; + use diskann_vector::{Norm, PureDistanceFunction, SimilarityScore, distance, norm::FastL2Norm}; use rand::{ SeedableRng, distr::{Distribution, Uniform}, @@ -1542,30 +1603,44 @@ mod tests { let chunk = Chunk::new(base.as_view().into()).unwrap(); let mut input = vec![0.0; dim]; - let mut output = vec![0.0; total]; + + let mut output_f32 = vec![0.0; total]; + let mut output_dot = vec![DotAndNorm::default(); total]; for _ in 0..PROCESS_INTO_TRIALS { input .iter_mut() .for_each(|i| *i = distribution.sample(rng) as f32); + let input_norm = (FastL2Norm).evaluate(&*input); + // Inner Product - InnerProduct::process_into(&chunk, &input, &mut output); + InnerProduct::process_into(&chunk, &input, &mut output_f32); // Check outputs - std::iter::zip(base.rows(), output.iter()).for_each(|(row, got)| { + std::iter::zip(base.rows(), output_f32.iter()).for_each(|(row, got)| { let expected: f32 = distance::InnerProduct::evaluate(row, input.as_slice()); assert_eq!(*got, expected); }); // Squared L2 - SquaredL2::process_into(&chunk, &input, &mut output); + SquaredL2::process_into(&chunk, &input, &mut output_f32); // Check outputs - std::iter::zip(base.rows(), output.iter()).for_each(|(row, got)| { + std::iter::zip(base.rows(), output_f32.iter()).for_each(|(row, got)| { let expected: f32 = distance::SquaredL2::evaluate(row, input.as_slice()); assert_eq!(*got, expected); }); + + // Cosine + Cosine::process_into(&chunk, &input, &mut output_dot); + + // Check outputs + std::iter::zip(base.rows(), output_dot.iter()).for_each(|(row, got)| { + let expected: f32 = distance::Cosine::evaluate(row, input.as_slice()); + let got: SimilarityScore = got.finish_cosine(input_norm); + assert_eq!(got.into_inner(), expected); + }); } } @@ -1598,6 +1673,20 @@ mod tests { InnerProduct::process_into(&chunk, query.as_slice(), dst.as_mut_slice()); } + #[test] + #[should_panic = "incorrect input vector dim"] + fn test_process_into_cosine_panics_on_from() { + let data = views::rowmajor::Owned::::from_element(5, 10, 0.0); + let chunk = Chunk::new(data.as_view().into()).unwrap(); + assert_eq!(chunk.dimension(), 10); + assert_eq!(chunk.num_centers(), 5); + + // Query is too large. + let query: Vec = vec![0.0; chunk.dimension() + 1]; + let mut dst = vec![DotAndNorm::default(); chunk.num_centers()]; + Cosine::process_into(&chunk, query.as_slice(), dst.as_mut_slice()); + } + #[test] #[should_panic] fn test_process_into_panics_on_into() { @@ -1611,4 +1700,18 @@ mod tests { let mut dst = vec![0.0; chunk.num_centers() + 1]; InnerProduct::process_into(&chunk, query.as_slice(), dst.as_mut_slice()); } + + #[test] + #[should_panic = "incorrect output vector dim"] + fn test_process_into_cosine_panics_on_into() { + let data = views::rowmajor::Owned::::from_element(5, 10, 0.0); + let chunk = Chunk::new(data.as_view().into()).unwrap(); + assert_eq!(chunk.dimension(), 10); + assert_eq!(chunk.num_centers(), 5); + + let query: Vec = vec![0.0; chunk.dimension()]; + // Dst is too big. + let mut dst = vec![DotAndNorm::default(); chunk.num_centers() + 1]; + Cosine::process_into(&chunk, query.as_slice(), dst.as_mut_slice()); + } } diff --git a/diskann-quantization/src/product/tables/transposed/table.rs b/diskann-quantization/src/product/tables/transposed/table.rs index 04d67a7d6c..e5909e5def 100644 --- a/diskann-quantization/src/product/tables/transposed/table.rs +++ b/diskann-quantization/src/product/tables/transposed/table.rs @@ -285,9 +285,9 @@ impl TransposedTable { /// * `query.len() != self.dim()`. /// * `partisl.nrows() != self.nchunks()`. /// * `partisl.ncols() != self.ncenters()`. - pub fn process_into(&self, query: &[f32], mut partials: rowmajor::Mut<'_, f32>) + pub fn process_into(&self, query: &[f32], mut partials: rowmajor::Mut<'_, U>) where - T: pivots::ProcessInto, + T: pivots::ProcessInto, { // Check Requirements assert_eq!( @@ -506,7 +506,9 @@ where mod test_compression { use std::collections::HashSet; - use diskann_vector::{PureDistanceFunction, distance}; + use diskann_vector::{ + MathematicalValue, Norm, PureDistanceFunction, distance, norm::FastL2NormSquared, + }; use rand::{ Rng, SeedableRng, distr::{Distribution, StandardUniform, Uniform}, @@ -519,9 +521,12 @@ mod test_compression { check_pqtable_batch_compression_errors, check_pqtable_single_compression_errors, }; use crate::{ - distances::{InnerProduct, SquaredL2}, + distances::{Cosine, InnerProduct, SquaredL2}, error::format, - product::tables::test::{create_dataset, create_pivot_tables}, + product::tables::{ + lookup::DotAndNorm, + test::{create_dataset, create_pivot_tables}, + }, }; use diskann_utils::lazy_format; @@ -866,12 +871,13 @@ mod test_compression { let mut output = views::rowmajor::Owned::::from_element(num_chunks, num_centers, 0.0); + let query: Vec<_> = (0..dim) .map(|_| value_distribution.sample(rng) as f32) .collect(); // Inner Product - table.process_into::(&query, output.as_view_mut()); + table.process_into::(&query, output.as_view_mut()); for chunk in 0..num_chunks { let range = offsets.at(chunk); @@ -892,7 +898,7 @@ mod test_compression { } // Squared L2 - table.process_into::(&query, output.as_view_mut()); + table.process_into::(&query, output.as_view_mut()); for chunk in 0..num_chunks { let range = offsets.at(chunk); @@ -911,6 +917,46 @@ mod test_compression { ); } } + + // Cosine + let mut output = views::rowmajor::Owned::::from_element( + num_chunks, + num_centers, + DotAndNorm::default(), + ); + + table.process_into::(&query, output.as_view_mut()); + for chunk in 0..num_chunks { + let range = offsets.at(chunk); + let query_chunk = &query[range.clone()]; + + for center in 0..num_centers { + let data_chunk = &pivots.row(center)[range.clone()]; + let expected_dot: MathematicalValue = + distance::InnerProduct::evaluate(query_chunk, data_chunk); + + assert_eq!( + output.element(chunk, center).dot(), + expected_dot.into_inner(), + "failed on (chunk, center) = ({}, {}) - offsets = {:?} - trial = {}", + chunk, + center, + offsets, + trial, + ); + + let expected_norm: f32 = (FastL2NormSquared).evaluate(data_chunk); + assert_eq!( + output.element(chunk, center).square_norm(), + expected_norm, + "failed on (chunk, center) = ({}, {}) - offsets = {:?} - trial = {}", + chunk, + center, + offsets, + trial, + ); + } + } } } @@ -945,7 +991,7 @@ mod test_compression { let query = vec![0.0; table.dim() - 1]; let mut partials = views::rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0); - table.process_into::(&query, partials.as_view_mut()); + table.process_into::(&query, partials.as_view_mut()); } #[test] @@ -960,7 +1006,7 @@ mod test_compression { // partials has the wrong numbers of rows. let mut partials = views::rowmajor::Owned::from_element(table.nchunks() - 1, table.ncenters(), 0.0); - table.process_into::(&query, partials.as_view_mut()); + table.process_into::(&query, partials.as_view_mut()); } #[test] @@ -975,6 +1021,6 @@ mod test_compression { // partials has the wrong numbers of rows. let mut partials = views::rowmajor::Owned::from_element(table.nchunks(), table.ncenters() - 1, 0.0); - table.process_into::(&query, partials.as_view_mut()); + table.process_into::(&query, partials.as_view_mut()); } } diff --git a/diskann-quantization/src/test_util.rs b/diskann-quantization/src/test_util.rs index 20bf6fac6e..d7f8f95cbb 100644 --- a/diskann-quantization/src/test_util.rs +++ b/diskann-quantization/src/test_util.rs @@ -212,6 +212,9 @@ pub(crate) enum Check { /// ``` AbsRel { abs: f32, rel: f32 }, + /// Two values must be exact. + Exact, + /// Skip the check entirely. #[cfg(not(miri))] Skip, @@ -226,6 +229,10 @@ impl Check { Self::AbsRel { abs, rel } } + pub(crate) const fn exact() -> Self { + Self::Exact + } + #[cfg(not(miri))] pub(crate) const fn skip() -> Self { Self::Skip @@ -274,6 +281,13 @@ impl Check { }) } } + Self::Exact => { + if got == expected { + Ok(()) + } else { + Err(CheckFailed::Exact { got, expected }) + } + } #[cfg(not(miri))] Self::Skip => Ok(()), } @@ -284,6 +298,8 @@ impl Check { pub(crate) enum CheckFailed { #[error("not within {ulp} ulp - got {got}, expected {expected}")] Ulp { ulp: usize, got: f32, expected: f32 }, + #[error("not exact got {got}, expected {expected}")] + Exact { got: f32, expected: f32 }, #[error( "not within {abs_limit}/{rel_limit} - errors {abs_got}/{rel_got} - \ got {got}, expected {expected}" diff --git a/diskann-quantization/src/views.rs b/diskann-quantization/src/views.rs index 04c4a09539..6bbe1cbdb3 100644 --- a/diskann-quantization/src/views.rs +++ b/diskann-quantization/src/views.rs @@ -211,6 +211,29 @@ where pub fn as_slice(&self) -> &[usize] { self.offsets.as_slice() } + + /// Return the maximum chunk dimension. + pub fn max_chunk_dim(&self) -> NonZeroUsize { + let mut max = NonZeroUsize::MIN; + let mut itr = self.offsets.as_slice().iter(); + + let Some(mut previous) = itr.next() else { + // NOTE: this is unreachable since we maintain the invariant that the number + // of chunks is at least 1. + return max; + }; + + for next in itr { + // This cannot underflow because offsets are verified to be strictly monotonic. + if let Some(dim) = NonZeroUsize::new(next - previous) { + max = max.max(dim); + } + + previous = next; + } + + max + } } pub type ChunkOffsetsView<'a> = ChunkOffsetsBase<&'a [usize]>; @@ -452,6 +475,21 @@ mod tests { offsets_view.as_slice().as_ptr(), offsets_owned.as_slice().as_ptr() ); + + // `max_chunk_dim` + assert_eq!(offsets.max_chunk_dim().get(), 4); + } + + #[test] + fn chunk_offset_max_dim() { + let offsets = ChunkOffsetsView::new(&[0, 10, 11, 14]).unwrap(); + assert_eq!(offsets.max_chunk_dim().get(), 10); + + let offsets = ChunkOffsetsView::new(&[0, 1, 2, 3]).unwrap(); + assert_eq!(offsets.max_chunk_dim().get(), 1); + + let offsets = ChunkOffsetsView::new(&[0, 1, 2, 5, 20, 21]).unwrap(); + assert_eq!(offsets.max_chunk_dim().get(), 15); } #[test]