diff --git a/diskann-benchmark/src/index/inmem2.rs b/diskann-benchmark/src/index/inmem2.rs index 0d574b16d8..45494d22ba 100644 --- a/diskann-benchmark/src/index/inmem2.rs +++ b/diskann-benchmark/src/index/inmem2.rs @@ -36,6 +36,7 @@ use diskann_inmem::{ }; use diskann_quantization::{ alloc::{GlobalAllocator, Poly}, + product::tables::BasicTable, spherical::iface, }; use diskann_utils::views::rowmajor::{self, Matrix}; @@ -56,7 +57,9 @@ pub(crate) fn register_benchmarks(registry: &mut Registry) -> anyhow::Result<()> registry.register("inmem2-f32", Build::::new())?; registry.register("inmem2-f16", Build::::new())?; registry.register("inmem2-u8", Build::::new())?; + registry.register("inmem2-spherical", SphericalBuild)?; + registry.register("inmem2-pq", ProductBuild)?; registry.register("inmem2-f32-stream", StreamingBenchmark::::new())?; Ok(()) @@ -109,6 +112,7 @@ mod dto { #[serde(rename_all = "kebab-case")] pub(super) enum Quantization { Spherical(Spherical), + Product(Product), } #[derive(Debug, Serialize, Deserialize)] @@ -124,6 +128,13 @@ mod dto { pub(super) bits: SphericalBits, } + #[derive(Debug, Serialize, Deserialize)] + pub(super) struct Product { + pub(super) chunks: NonZeroUsize, + pub(super) centers: NonZeroUsize, + pub(super) seed: u64, + } + //-----------// // Streaming // //-----------// @@ -357,6 +368,7 @@ impl Display for BuildParams { enum Quantization { None, Spherical(Spherical), + Product(Product), } impl Quantization { @@ -366,6 +378,7 @@ impl Quantization { dto::Quantization::Spherical(spherical) => { Self::Spherical(Spherical::from_raw(spherical)) } + dto::Quantization::Product(product) => Self::Product(Product::from_raw(product)), } } else { Self::None @@ -378,6 +391,13 @@ impl Quantization { _ => None, } } + + fn as_product(&self) -> Option<&Product> { + match self { + Self::Product(product) => Some(product), + _ => None, + } + } } impl std::fmt::Display for Quantization { @@ -389,6 +409,11 @@ impl std::fmt::Display for Quantization { kv.push("spherical", spherical); write!(f, "{}", kv) } + Self::Product(product) => { + let mut kv = KeyValue::new(); + kv.push("product", product); + write!(f, "{}", kv) + } } } } @@ -478,6 +503,78 @@ impl std::fmt::Display for Spherical { } } +#[derive(Debug)] +struct Product { + chunks: NonZeroUsize, + centers: NonZeroUsize, + seed: u64, +} + +impl Product { + fn from_raw(raw: dto::Product) -> Self { + let dto::Product { + chunks, + centers, + seed, + } = raw; + + Self { + chunks, + centers, + seed, + } + } + + fn train( + &self, + data: rowmajor::Ref<'_, f32>, + num_threads: NonZeroUsize, + ) -> anyhow::Result { + use diskann_quantization::{ + cancel::DontCancel, + product::{self, train::TrainQuantizer}, + random, + views::ChunkOffsets, + Parallelism, + }; + + let trainer = product::train::LightPQTrainingParameters::new(self.centers.get(), 5); + let threadpool = rayon::ThreadPoolBuilder::new() + .num_threads(num_threads.get()) + .build()?; + + threadpool.install(|| -> anyhow::Result<_> { + let table = trainer.train( + data, + ChunkOffsets::partition(NonZeroUsize::new(data.ncols()).unwrap(), self.chunks)? + .as_view(), + Parallelism::Rayon, + &random::StdRngBuilder::new(self.seed), + &DontCancel, + )?; + + Ok(table) + }) + } +} + +impl std::fmt::Display for Product { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let Self { + chunks, + centers, + seed, + } = self; + + let mut kv = KeyValue::new(); + kv.push("chunks", chunks); + kv.push("centers", centers); + kv.push("seed", seed); + + write!(f, "{}", kv) + } +} + //-------------// // StaticBuild // //-------------// @@ -912,6 +1009,149 @@ impl Benchmark for SphericalBuild { } } +//----------------------// +// Product Quantization // +//----------------------// + +#[derive(Debug)] +struct ProductBuild; + +impl Benchmark for ProductBuild { + type Input = StaticBuild; + type Output = (); + + fn try_match(&self, input: &StaticBuild, context: &MatchContext) -> Score { + let mut score = context.success(0); + + let DispatchParams { + data_type, + quantization, + distance, + } = input.dispatch_params(); + + if !matches!(quantization, Quantization::Product(_)) { + score.fail(2000, &"needed product-quantization"); + } + + if !f32::is_match(data_type) { + score.fail( + 1000, + &format_args!( + "expected data-type {}, instead got {}", + Quote(f32::DATA_TYPE), + Quote(data_type) + ), + ) + } + + accept_all(&distance); + score + } + + fn description(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "product-quantized index build-and-search",)?; + + Ok(()) + } + + fn run( + &self, + input: &StaticBuild, + checkpoint: Checkpoint<'_>, + mut output: &mut dyn Output, + ) -> anyhow::Result<()> { + writeln!(output, "{input}\n")?; + + let product = input.quantization.as_product().unwrap(); + + // Load data. + let data: Arc> = Arc::new(datafiles::load_dataset( + datafiles::BinFile(&input.data.data), + )?); + + let dim = data.ncols(); + let num_points = data.nrows(); + writeln!(output, "Loaded {num_points} points, dim={dim}")?; + + let table = product.train(data.as_view(), input.build.num_threads)?; + + // Compute the medoid of the dataset as the single start point. + let start = StartPointStrategy::Medoid.compute(data.as_view())?; + let config = repr::product::Product::config( + table, + input.data.distance.into(), + Capacity::new(num_points), + MaxDegree::new(input.build.config.max_degree().get()), + start, + repr::product::Rerank::F16, + )? + .thread_hint(Some(input.build.num_threads)); + + let provider = Provider::<_, u32>::new(config)?; + let index = Arc::new(DiskANNIndex::new( + input.build.config.clone(), + provider, + None, + )); + + // Build via SingleInsert. + let rt = benchmark_core::tokio::runtime(input.build.num_threads.get())?; + let builder = build_core::graph::SingleInsert::new( + index.clone(), + data, + Strategy, + build_core::ids::Identity::::new(), + ); + + let build_results = build_core::build_tracked( + builder, + build_core::Parallelism::dynamic(diskann::utils::ONE, input.build.num_threads), + &rt, + Some(&ProgressMeter::new(output)), + )?; + + let total_build_time = build_results.end_to_end_latency(); + writeln!( + output, + "\nBuild complete in {:.2}s", + total_build_time.as_seconds() + )?; + checkpoint.checkpoint(&total_build_time)?; + + // Search. + let queries: Arc> = Arc::new(datafiles::load_dataset( + datafiles::BinFile(&input.search.queries), + )?); + let max_k = input.search.maximum_recall_k(); + let groundtruth = datafiles::load_groundtruth( + datafiles::BinFile(&input.search.groundtruth), + Some(max_k), + )?; + + writeln!(output, "Loaded {} queries\n", queries.nrows())?; + + let knn = benchmark_core::search::graph::KNN::new( + index, + queries, + benchmark_core::search::graph::Strategy::broadcast(Strategy), + )?; + + let results = _knn( + &knn, + &groundtruth, + input.search.reps, + &input.search.num_threads, + &input.search.runs, + )?; + + let results = AggregatedSearchResults::Topk(results); + + writeln!(output, "{}", results)?; + + Ok(()) + } +} + fn _knn( runner: &dyn crate::index::search::knn::Knn, groundtruth: &dyn benchmark_core::recall::Rows, diff --git a/diskann-inmem/integration/index/object.rs b/diskann-inmem/integration/index/object.rs index db25c4259b..3fe4bfa7d9 100644 --- a/diskann-inmem/integration/index/object.rs +++ b/diskann-inmem/integration/index/object.rs @@ -237,3 +237,4 @@ macro_rules! index { index!({ T } repr::Full where T: repr::FullPrecision + FromSlice + AsDataType); index!(repr::Spherical); +index!(repr::Product); diff --git a/diskann-inmem/integration/index/runner.rs b/diskann-inmem/integration/index/runner.rs index 0e5d1e95a8..e6e3cfb5a2 100644 --- a/diskann-inmem/integration/index/runner.rs +++ b/diskann-inmem/integration/index/runner.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use std::{io::Write, sync::Arc}; +use std::{io::Write, num::NonZeroUsize, sync::Arc}; use anyhow::Context; use diskann::graph::{DiskANNIndex, search::Knn}; @@ -112,6 +112,17 @@ mod dto { // Quantization Parameters // //-------------------------// + pub(super) mod quantization { + use super::*; + + #[derive(Debug, Serialize, Deserialize)] + #[serde(rename_all = "kebab-case")] + pub(in crate::index::runner) enum Rerank { + None, + F16, + } + } + pub(super) mod spherical { use super::*; @@ -122,13 +133,6 @@ mod dto { Two, Four, } - - #[derive(Debug, Serialize, Deserialize)] - #[serde(rename_all = "kebab-case")] - pub(in crate::index::runner) enum Rerank { - None, - F16, - } } #[derive(Debug, Serialize, Deserialize)] @@ -139,7 +143,11 @@ mod dto { }, Spherical { bits: spherical::Bits, - rerank: spherical::Rerank, + rerank: quantization::Rerank, + }, + Product { + chunks: NonZeroUsize, + rerank: quantization::Rerank, }, } @@ -259,6 +267,50 @@ struct Bundle { groundtruth: rowmajor::Owned, } +mod quantization { + use super::*; + + #[derive(Debug, Clone, Copy)] + pub(super) enum Rerank { + None, + F16, + } + + impl Rerank { + pub(super) fn from_raw(raw: dto::quantization::Rerank) -> Self { + match raw { + dto::quantization::Rerank::None => Self::None, + dto::quantization::Rerank::F16 => Self::F16, + } + } + + pub(super) fn as_raw(&self) -> dto::quantization::Rerank { + match self { + Self::None => dto::quantization::Rerank::None, + Self::F16 => dto::quantization::Rerank::F16, + } + } + } + + impl From for inmem::repr::spherical::Rerank { + fn from(rerank: Rerank) -> Self { + match rerank { + Rerank::None => inmem::repr::spherical::Rerank::None, + Rerank::F16 => inmem::repr::spherical::Rerank::F16, + } + } + } + + impl From for inmem::repr::product::Rerank { + fn from(rerank: Rerank) -> Self { + match rerank { + Rerank::None => inmem::repr::product::Rerank::None, + Rerank::F16 => inmem::repr::product::Rerank::F16, + } + } + } +} + mod spherical { use super::*; @@ -286,37 +338,6 @@ mod spherical { } } } - - #[derive(Debug, Clone, Copy)] - pub(super) enum Rerank { - None, - F16, - } - - impl Rerank { - pub(super) fn from_raw(raw: dto::spherical::Rerank) -> Self { - match raw { - dto::spherical::Rerank::None => Self::None, - dto::spherical::Rerank::F16 => Self::F16, - } - } - - pub(super) fn as_raw(&self) -> dto::spherical::Rerank { - match self { - Self::None => dto::spherical::Rerank::None, - Self::F16 => dto::spherical::Rerank::F16, - } - } - } - - impl From for inmem::repr::spherical::Rerank { - fn from(rerank: Rerank) -> Self { - match rerank { - Rerank::None => inmem::repr::spherical::Rerank::None, - Rerank::F16 => inmem::repr::spherical::Rerank::F16, - } - } - } } #[derive(Debug)] @@ -326,7 +347,11 @@ enum Representation { }, Spherical { bits: spherical::Bits, - rerank: spherical::Rerank, + rerank: quantization::Rerank, + }, + Product { + chunks: NonZeroUsize, + rerank: quantization::Rerank, }, } @@ -336,7 +361,11 @@ impl Representation { dto::Representation::FullPrecision { data_type } => Self::FullPrecision { data_type }, dto::Representation::Spherical { bits, rerank } => Self::Spherical { bits: spherical::Bits::from_raw(bits), - rerank: spherical::Rerank::from_raw(rerank), + rerank: quantization::Rerank::from_raw(rerank), + }, + dto::Representation::Product { chunks, rerank } => Self::Product { + chunks, + rerank: quantization::Rerank::from_raw(rerank), }, } } @@ -350,6 +379,10 @@ impl Representation { bits: bits.as_raw(), rerank: rerank.as_raw(), }, + Self::Product { chunks, rerank } => dto::Representation::Product { + chunks: *chunks, + rerank: rerank.as_raw(), + }, } } } @@ -520,6 +553,9 @@ impl Test { Representation::Spherical { bits, rerank } => { self.create_spherical(data, *bits, *rerank) } + Representation::Product { chunks, rerank } => { + self.create_product(data, *chunks, *rerank) + } } } @@ -527,7 +563,7 @@ impl Test { &self, data: DatasetView<'_>, bits: spherical::Bits, - rerank: spherical::Rerank, + rerank: quantization::Rerank, ) -> anyhow::Result> { use diskann_quantization::{ algorithms::transforms, @@ -581,6 +617,56 @@ impl Test { Ok(finish(Provider::new(config)?, index_config)) } + + fn create_product( + &self, + data: DatasetView<'_>, + chunks: NonZeroUsize, + rerank: quantization::Rerank, + ) -> anyhow::Result> { + use diskann_quantization::{ + Parallelism, + cancel::DontCancel, + product::{self, train::TrainQuantizer}, + random, + views::ChunkOffsets, + }; + + let DatasetView::F32(data) = data else { + anyhow::bail!("spherical quantization only supports f32 data"); + }; + + // Step 1: Train a generic quantizer. + let trainer = product::train::LightPQTrainingParameters::new(256, 2); + let Some(dim) = NonZeroUsize::new(data.ncols()) else { + anyhow::bail!("cannot compress a zero dimensional dataset"); + }; + + let table = trainer.train( + data, + ChunkOffsets::partition(dim, chunks)?.as_view(), + Parallelism::Sequential, + &random::StdRngBuilder::new(0xc0ff33), + &DontCancel, + )?; + + let start_point = rowmajor::Owned::row_vector(Box::from( + ::compute_medoid(data), + )); + + // Step 2: Create the config. + let config = diskann_inmem::repr::Product::config( + table, + self.data.metric.into(), + Capacity::new(data.nrows()), + MaxDegree::new(self.build.config.max_degree().get()), + start_point, + rerank.into(), + )?; + + let index_config = self.build.config.clone(); + Ok(finish(Provider::new(config)?, index_config)) + } } fn finish(provider: DP, config: diskann::graph::Config) -> Arc @@ -671,7 +757,7 @@ impl diskann_benchmark_runner::Benchmark for FullPrecision { Representation::FullPrecision { .. } => { // We match all valid data-types } - Representation::Spherical { .. } => { + Representation::Spherical { .. } | Representation::Product { .. } => { let data_type = input.data.data_type; // Ensure that the data type if `f32`. if data_type != DataType::F32 { @@ -702,6 +788,7 @@ impl diskann_benchmark_runner::Benchmark for FullPrecision { let data_type = match input.representation { Representation::FullPrecision { data_type } => data_type, Representation::Spherical { .. } => input.data.data_type, + Representation::Product { .. } => input.data.data_type, }; // Load the data and perform any necessary data conversions. diff --git a/diskann-inmem/integration/jsons/graph/product/cosine-baseline.json b/diskann-inmem/integration/jsons/graph/product/cosine-baseline.json new file mode 100644 index 0000000000..681808b2ae --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/cosine-baseline.json @@ -0,0 +1,124 @@ +[ + { + "input": { + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "/yfcc/yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "/yfcc/groundtruth_cosine.bin", + "metric": "cosine", + "preprocess": [], + "queries": "/yfcc/yfcc_query_100.fbin" + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": 1, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + }, + "type": "integration-test" + }, + "results": { + "build": { + "append_neighbors": 39095, + "distance": 549876, + "get_neighbors": 298634, + "get_vector": 2008258, + "query_distance": 1724191, + "set_neighbors": 11168, + "set_vector": 10000 + }, + "knn": [ + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5625, + "get_vector": 34053, + "query_distance": 34053, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 34053, + "hops": 5625 + }, + "recall": { + "average": 0.523, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5977, + "get_vector": 37610, + "query_distance": 37610, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 37610, + "hops": 5977 + }, + "recall": { + "average": 0.528, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 10794, + "get_vector": 56255, + "query_distance": 56255, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 56255, + "hops": 10794 + }, + "recall": { + "average": 0.531, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + } + ] + } + } +] \ No newline at end of file diff --git a/diskann-inmem/integration/jsons/graph/product/cosine.json b/diskann-inmem/integration/jsons/graph/product/cosine.json new file mode 100644 index 0000000000..1702ea7996 --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/cosine.json @@ -0,0 +1,52 @@ +{ + "search_directories": [ + "yfcc" + ], + "output_directory": null, + "jobs": [ + { + "type": "integration-test", + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "groundtruth_cosine.bin", + "metric": "cosine", + "queries": "yfcc_query_100.fbin", + "preprocess": [] + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": null, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + } + } + ] +} diff --git a/diskann-inmem/integration/jsons/graph/product/ip-baseline.json b/diskann-inmem/integration/jsons/graph/product/ip-baseline.json new file mode 100644 index 0000000000..bbeacb85d3 --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/ip-baseline.json @@ -0,0 +1,124 @@ +[ + { + "input": { + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "/yfcc/yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "/yfcc/groundtruth_ip.bin", + "metric": "inner-product", + "preprocess": [], + "queries": "/yfcc/yfcc_query_100.fbin" + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": 1, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + }, + "type": "integration-test" + }, + "results": { + "build": { + "append_neighbors": 128225, + "distance": 5596404, + "get_neighbors": 384814, + "get_vector": 2467136, + "query_distance": 1545792, + "set_neighbors": 41655, + "set_vector": 10000 + }, + "knn": [ + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5254, + "get_vector": 31162, + "query_distance": 31162, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 31162, + "hops": 5254 + }, + "recall": { + "average": 0.243, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5420, + "get_vector": 32277, + "query_distance": 32277, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 32277, + "hops": 5420 + }, + "recall": { + "average": 0.244, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 10306, + "get_vector": 46065, + "query_distance": 46065, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 46065, + "hops": 10306 + }, + "recall": { + "average": 0.244, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + } + ] + } + } +] \ No newline at end of file diff --git a/diskann-inmem/integration/jsons/graph/product/ip.json b/diskann-inmem/integration/jsons/graph/product/ip.json new file mode 100644 index 0000000000..6c56c71c84 --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/ip.json @@ -0,0 +1,52 @@ +{ + "search_directories": [ + "yfcc" + ], + "output_directory": null, + "jobs": [ + { + "type": "integration-test", + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "groundtruth_ip.bin", + "metric": "inner-product", + "queries": "yfcc_query_100.fbin", + "preprocess": [] + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": null, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + } + } + ] +} diff --git a/diskann-inmem/integration/jsons/graph/product/l2-baseline.json b/diskann-inmem/integration/jsons/graph/product/l2-baseline.json new file mode 100644 index 0000000000..58eae3a798 --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/l2-baseline.json @@ -0,0 +1,246 @@ +[ + { + "input": { + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "/yfcc/yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "/yfcc/groundtruth.bin", + "metric": "l2", + "preprocess": [], + "queries": "/yfcc/yfcc_query_100.fbin" + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": 1, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + }, + "type": "integration-test" + }, + "results": { + "build": { + "append_neighbors": 37939, + "distance": 533446, + "get_neighbors": 297940, + "get_vector": 1985381, + "query_distance": 1702322, + "set_neighbors": 11098, + "set_vector": 10000 + }, + "knn": [ + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5568, + "get_vector": 33330, + "query_distance": 33330, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 33330, + "hops": 5568 + }, + "recall": { + "average": 0.527, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5985, + "get_vector": 37289, + "query_distance": 37289, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 37289, + "hops": 5985 + }, + "recall": { + "average": 0.529, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 10798, + "get_vector": 55660, + "query_distance": 55660, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 55660, + "hops": 10798 + }, + "recall": { + "average": 0.537, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + } + ] + } + }, + { + "input": { + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "/yfcc/yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "/yfcc/groundtruth.bin", + "metric": "l2", + "preprocess": [], + "queries": "/yfcc/yfcc_query_100.fbin" + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "f16" + } + }, + "search": { + "knn": [ + { + "beam_width": 1, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + }, + "type": "integration-test" + }, + "results": { + "build": { + "append_neighbors": 37939, + "distance": 533446, + "get_neighbors": 297940, + "get_vector": 1985381, + "query_distance": 1702322, + "set_neighbors": 11098, + "set_vector": 10000 + }, + "knn": [ + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5568, + "get_vector": 38430, + "query_distance": 38430, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 33330, + "hops": 5568 + }, + "recall": { + "average": 0.891, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 5985, + "get_vector": 42389, + "query_distance": 42389, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 37289, + "hops": 5985 + }, + "recall": { + "average": 0.891, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + }, + { + "counters": { + "append_neighbors": 0, + "distance": 0, + "get_neighbors": 10798, + "get_vector": 65760, + "query_distance": 65760, + "set_neighbors": 0, + "set_vector": 0 + }, + "misc": { + "cmps": 55660, + "hops": 10798 + }, + "recall": { + "average": 0.953, + "num_queries": 100, + "recall_k": 10, + "recall_n": 10 + } + } + ] + } + } +] \ No newline at end of file diff --git a/diskann-inmem/integration/jsons/graph/product/l2.json b/diskann-inmem/integration/jsons/graph/product/l2.json new file mode 100644 index 0000000000..ae2cda0b89 --- /dev/null +++ b/diskann-inmem/integration/jsons/graph/product/l2.json @@ -0,0 +1,96 @@ +{ + "search_directories": [ + "yfcc" + ], + "output_directory": null, + "jobs": [ + { + "type": "integration-test", + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "groundtruth.bin", + "metric": "l2", + "queries": "yfcc_query_100.fbin", + "preprocess": [] + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "none" + } + }, + "search": { + "knn": [ + { + "beam_width": null, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + } + }, + { + "type": "integration-test", + "content": { + "build": { + "alpha": 1.2000000476837158, + "l_build": 20, + "max_degree": 20, + "pruned_degree": 16 + }, + "data": { + "data": "yfcc_10k.fbin", + "data_type": "f32", + "groundtruth": "groundtruth.bin", + "metric": "l2", + "queries": "yfcc_query_100.fbin", + "preprocess": [] + }, + "representation": { + "product": { + "chunks": 16, + "rerank": "f16" + } + }, + "search": { + "knn": [ + { + "beam_width": null, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 50 + }, + { + "beam_width": 3, + "knn": 10, + "search_l": 100 + } + ] + } + } + } + ] +} diff --git a/diskann-inmem/integration/main.rs b/diskann-inmem/integration/main.rs index 0a8eff0aac..6ba14b9564 100644 --- a/diskann-inmem/integration/main.rs +++ b/diskann-inmem/integration/main.rs @@ -343,4 +343,38 @@ mod tests { "graph/spherical/four-bit-cosine-baseline.json", ); } + + //----------------------// + // Product Quantization // + //----------------------// + + #[test] + #[cfg(not(any(miri, coverage)))] + fn graph_product_l2() { + run_regression_example( + "graph/product/l2.json", + "checks.json", + "graph/product/l2-baseline.json", + ); + } + + #[test] + #[cfg(not(any(miri, coverage)))] + fn graph_product_ip() { + run_regression_example( + "graph/product/ip.json", + "checks.json", + "graph/product/ip-baseline.json", + ); + } + + #[test] + #[cfg(not(miri))] + fn graph_product_cosine() { + run_regression_example( + "graph/product/cosine.json", + "checks.json", + "graph/product/cosine-baseline.json", + ); + } } diff --git a/diskann-inmem/src/repr/internal/mod.rs b/diskann-inmem/src/repr/internal/mod.rs index 6d7770047b..aa9c6aa1fe 100644 --- a/diskann-inmem/src/repr/internal/mod.rs +++ b/diskann-inmem/src/repr/internal/mod.rs @@ -10,6 +10,9 @@ pub(super) mod intrusive; #[cfg(any(feature = "quantization", test))] pub(super) mod simple; +#[cfg(feature = "quantization")] +pub(super) mod quantization; + pub(super) mod macros; ////////// diff --git a/diskann-inmem/src/repr/internal/quantization/mod.rs b/diskann-inmem/src/repr/internal/quantization/mod.rs new file mode 100644 index 0000000000..250cb22c2f --- /dev/null +++ b/diskann-inmem/src/repr/internal/quantization/mod.rs @@ -0,0 +1,82 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +pub(in crate::repr) mod rerank; +pub(in crate::repr) use rerank::{Rerank, Reranker}; + +/// Distance metric used for quantization. +#[derive(Debug, Clone, Copy)] +pub(in crate::repr) enum Metric { + SquaredL2, + InnerProduct, + Cosine, +} + +impl Metric { + #[cfg(test)] + fn all() -> [Self; 3] { + [Self::SquaredL2, Self::InnerProduct, Self::Cosine] + } +} + +impl Metric { + pub(in crate::repr) fn as_vector_metric(&self) -> diskann_vector::distance::Metric { + use diskann_vector::distance::Metric as VMetric; + + match self { + Self::SquaredL2 => VMetric::L2, + Self::InnerProduct => VMetric::InnerProduct, + Self::Cosine => VMetric::Cosine, + } + } +} + +impl From for Metric { + #[inline] + fn from(m: diskann_vector::distance::Metric) -> Metric { + use diskann_vector::distance::Metric as VMetric; + + match m { + VMetric::L2 => Self::SquaredL2, + VMetric::InnerProduct => Self::InnerProduct, + VMetric::Cosine | VMetric::CosineNormalized => Self::Cosine, + } + } +} + +impl From for diskann_quantization::spherical::SupportedMetric { + #[inline] + fn from(m: Metric) -> diskann_quantization::spherical::SupportedMetric { + match m { + Metric::SquaredL2 => Self::SquaredL2, + Metric::InnerProduct => Self::InnerProduct, + Metric::Cosine => Self::Cosine, + } + } +} + +impl From for Metric { + #[inline] + fn from(m: diskann_quantization::spherical::SupportedMetric) -> Metric { + use diskann_quantization::spherical::SupportedMetric; + + match m { + SupportedMetric::SquaredL2 => Self::SquaredL2, + SupportedMetric::InnerProduct => Self::InnerProduct, + SupportedMetric::Cosine => Self::Cosine, + } + } +} + +impl From for diskann_quantization::product::tables::padded::Metric { + #[inline] + fn from(m: Metric) -> diskann_quantization::product::tables::padded::Metric { + match m { + Metric::SquaredL2 => Self::SquaredL2, + Metric::InnerProduct => Self::InnerProduct, + Metric::Cosine => Self::Cosine, + } + } +} diff --git a/diskann-inmem/src/repr/internal/quantization/rerank.rs b/diskann-inmem/src/repr/internal/quantization/rerank.rs new file mode 100644 index 0000000000..8a1e22e8d5 --- /dev/null +++ b/diskann-inmem/src/repr/internal/quantization/rerank.rs @@ -0,0 +1,344 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_vector::distance::{Distance, DistanceProvider}; +use half::f16; + +use crate::{ + counters::LocalCounters, + epoch, + num::Bytes, + repr::{ + self, + internal::{Calf, quantization::Metric}, + }, + store::{ + self, + optional::Optional, + simple::{self, Simple}, + }, +}; + +/// Choose how data is going to be reranked. +#[derive(Debug, Clone, Copy, PartialEq)] +pub(in crate::repr) enum Rerank { + /// No reranking will be performed and no space for higher precision vectors will be + /// allocated. + None, + + /// Use 16-bit floating point numbers to store the higher precision representation. + /// These will be used automatically during search to rerank candidates. + F16, +} + +/// Internal representation of [`Rerank`]. +/// +/// This is used for computing distances among the raw values in the auxiliary store. +#[derive(Debug)] +pub(in crate::repr) enum Reranker { + None, + F16(Distance), +} + +impl Reranker { + /// Construct a new [`Reranker`] and a [`store::slots::SlotsConfig`] for the auxiliary + /// store. + pub(in crate::repr) fn new_with_config( + rerank: Rerank, + metric: Metric, + dim: usize, + ) -> (Self, Option) { + let this = match rerank { + Rerank::None => Self::None, + Rerank::F16 => { + let distance = >::distance_comparer( + metric.as_vector_metric(), + Some(dim), + ); + + Self::F16(distance) + } + }; + + let config = match &this { + Self::None => None, + Self::F16(_) => Some(Simple::config(this.bytes_for(dim))), + }; + + (this, config) + } + + #[expect( + clippy::expect_used, + reason = "the arithmetic should not overflow for the feasible `dim` values" + )] + fn bytes_for(&self, dim: usize) -> Bytes { + match self { + Self::None => Bytes::new(0), + Self::F16(_) => Bytes::new(dim.checked_mul(2).expect("f16 is smaller than the f32")), + } + } + + /// Create a [`repr::PostProcess`]. + /// + /// This assumes that `simple` has the same dimensions as `self`'s contained distance + /// computation and that `guard` belongs to `simple`. + /// + /// # Pre-conditions + /// + /// This requires that `slots` is the [`store::slots::Slots`] created from the + /// configuration returned in [`Self::new_with_config`]. + #[expect( + clippy::panic, + reason = "this is an internal method that must be set up correctly" + )] + pub(in crate::repr) fn post_process<'a>( + &'a self, + query: &'a [f32], + guard: &epoch::Guard<'a>, + slots: &'a Optional, + counters: &LocalCounters<'a>, + ) -> Option> { + match (self, slots.slots()) { + (Self::None, None) => None, + (Self::F16(distance), Some(simple)) => { + let distance = repr::full::QueryDistance::new(Calf::Borrowed(query), *distance); + let reader = simple.reader(guard.share()); + let post_process = + repr::internal::simple::Reranker::new(reader, distance, counters.fork()); + Some(Box::new(post_process)) + } + _ => panic!("invalid combination of arguments"), + } + } + + /// Store the vector `v` into the raw buffer `buf`. + /// + /// # Pre-conditions + /// + /// `buf` must be consistent with the configuration returned from [`Self::new_with_config`], + /// and may only be `None` if that configuration was `None`. + /// + /// If it is `Some`, this function may panic if its length is not consistent with the + /// original configuration. + #[expect( + clippy::panic, + reason = "this is an internal method that must be set up correctly" + )] + pub(in crate::repr) fn store(&self, v: &[f32], buf: &mut Option>) { + match (self, buf) { + (Self::None, None) => {} + (Self::F16(_), Some(exclusive)) => { + use diskann_vector::conversion::CastFromSlice; + bytemuck::cast_slice_mut::(exclusive.as_mut_slice()).cast_from_slice(v); + } + _ => panic!("invalid combination of arguments"), + } + } +} + +/// Test that [`repr::PostProcess`] reranks correctly. +/// +/// Pass all `ids` to [`repr::PostProcess::post_process`]. Verify that all ids not +/// present in `distances` have been removed and the remaining ids are present, sorted, +/// and have distance values matching those in `distances`. +#[cfg(test)] +pub(in crate::repr) fn test_rerank( + post_process: &mut dyn repr::PostProcess, + distances: hashbrown::HashMap, + ids: &[crate::num::SlotId], + ctx: &dyn std::fmt::Display, +) { + use diskann::neighbor::Neighbor; + + use crate::num::SlotId; + + let mut buffer: Vec<_> = ids + .iter() + .map(|slot_id| Neighbor::new(slot_id.value(), 0.0)) + .collect(); + + post_process.post_process(&mut buffer).unwrap(); + let mut previous = f32::NEG_INFINITY; + assert_eq!(buffer.len(), distances.len(), "{ctx}"); + for (pos, neighbor) in buffer.iter().enumerate() { + let current = *neighbor.distance(); + assert_eq!( + current, + distances[&SlotId(*neighbor.id())], + "failed in position {} of {:?} -- {ctx}", + pos, + buffer + ); + + assert!( + current >= previous, + "distances is not monotonically increasing, previous = {}, current = {} -- {}", + previous, + current, + ctx, + ); + + previous = current; + } + + assert!( + previous > f32::NEG_INFINITY, + "previous = {} -- {}", + previous, + ctx + ); +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::assert_matches; + + use diskann::utils::IntoUsize; + + use crate::{ + counters::Counters, + num::{Capacity, LogicalId, MaxDegree, SlotId}, + repr::test::Reference, + store::Store, + }; + + const TEST_DIM: usize = 1; + + /// Create a reranker of dim 5 for the given config and metric. + fn make_store(rerank: Rerank, metric: Metric) -> (Reranker, Store>) { + let (reranker, config) = Reranker::new_with_config(rerank, metric, TEST_DIM); + + let store = Store::new( + store::Layout::new(Capacity::new(10), MaxDegree::new(0), 0), + store::Config::new(), + config, + ) + .unwrap(); + + (reranker, store) + } + + #[test] + fn test_disabled_rerank() { + let (reranker, store) = make_store(Rerank::None, Metric::SquaredL2); + assert_matches!(reranker, Reranker::None); + assert!( + store.slots().slots().is_none(), + "Rerank::None should create empty optional slots" + ); + + let query = std::slice::from_ref(&1.0); + let counters = Counters::new(); + + let post_processor = store + .guard(|slots, guard| reranker.post_process(query, &guard, slots, &counters.local())) + .unwrap(); + + assert!(post_processor.is_none()); + } + + #[test] + fn test_reranker_f16() { + for metric in Metric::all() { + let (reranker, store) = make_store(Rerank::F16, metric); + let mut map = Reference::new(TEST_DIM); + + // Insert the following values by logical id: + // + // LogicalID Value Deleted + // 5 2.5 No + // 4 1.5 No + // 3 0.5 Yes + // 2 -0.5 Yes + // 1 -1.5 No + // 0 -2.5 Yes + // + // The value -1.0 is used as the query, which generates the following values for + // the metrics being tested (note, we're using similarity scores here): + // + // LogicalId Value SquaredL2 InnerProduct Cosine + // 5 2.5 12.25 2.5 2.0 + // 4 1.5 6.25 1.5 2.0 + // 1 -1.5 0.25 -1.5 0.0 + + let distances = match metric { + Metric::SquaredL2 => [ + (LogicalId(5), 12.25), + (LogicalId(4), 6.25), + (LogicalId(1), 0.25), + ], + Metric::InnerProduct => [ + (LogicalId(5), 2.5), + (LogicalId(4), 1.5), + (LogicalId(1), -1.5), + ], + Metric::Cosine => [ + (LogicalId(5), 2.0), + (LogicalId(4), 2.0), + (LogicalId(1), 0.0), + ], + }; + + // Insert + for i in (0..=5).rev() { + let v = (i as f32) - 2.5; + + let mut exclusive = store.acquire().unwrap(); + + map.insert( + LogicalId(i), + SlotId(exclusive.slot()), + std::slice::from_ref(&v), + ); + + reranker.store(std::slice::from_ref(&v), exclusive.data()); + + exclusive.publish(); + } + + // Delete + store + .retire(map.slot_id_for(LogicalId(3)).value().into_usize()) + .unwrap(); + store + .retire(map.slot_id_for(LogicalId(2)).value().into_usize()) + .unwrap(); + store + .retire(map.slot_id_for(LogicalId(0)).value().into_usize()) + .unwrap(); + + let counters = Counters::new(); + let query = std::slice::from_ref(&-1.0); + + // Create the post-processor. + let mut post_process = store + .guard(|slots, guard| { + reranker + .post_process(query, &guard, slots, &counters.local()) + .unwrap() + }) + .unwrap(); + + let expected: hashbrown::HashMap<_, _> = distances + .map(|(logical_id, distance)| (map.slot_id_for(logical_id), distance)) + .into_iter() + .collect(); + + test_rerank( + &mut *post_process, + expected, + &[6, 5, 4, 3, 2, 1, 0].map(SlotId), + &format_args!("metric = {:?}", metric), + ); + } + } +} diff --git a/diskann-inmem/src/repr/mod.rs b/diskann-inmem/src/repr/mod.rs index e20e41a564..2c3bd13ab9 100644 --- a/diskann-inmem/src/repr/mod.rs +++ b/diskann-inmem/src/repr/mod.rs @@ -26,6 +26,14 @@ mod internal; pub mod full; pub use full::{Full, FullPrecision}; +#[cfg(feature = "quantization")] +#[cfg_attr(docsrs, doc(cfg(feature = "quantization")))] +pub mod product; + +#[cfg(feature = "quantization")] +#[cfg_attr(docsrs, doc(cfg(feature = "quantization")))] +pub use product::Product; + #[cfg(feature = "quantization")] #[cfg_attr(docsrs, doc(cfg(feature = "quantization")))] pub mod spherical; diff --git a/diskann-inmem/src/repr/product.rs b/diskann-inmem/src/repr/product.rs new file mode 100644 index 0000000000..a25cfa1657 --- /dev/null +++ b/diskann-inmem/src/repr/product.rs @@ -0,0 +1,1206 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::num::NonZeroUsize; + +use diskann::{ANNError, ANNResult, error::ErrorContext, utils::IntoUsize}; +use diskann_quantization::{distances as quant_distances, product::tables}; +use diskann_utils::{ + lazy_format, + object_pool::{self, ObjectPool}, + views::rowmajor::{self, Matrix, MatrixMut}, +}; +use thiserror::Error; + +use crate::{ + counters::LocalCounters, + num::{Bytes, Capacity, IdLimit, MaxDegree}, + prefetch, repr, + store::{ + self, Store, + cons::{self, Cons}, + intrusive::{self, Intrusive}, + optional::Optional, + simple::{self, Simple}, + }, +}; + +/// Distance metric to use. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Metric { + SquaredL2, + InnerProduct, + Cosine, +} + +impl Metric { + fn as_internal_metric(&self) -> repr::internal::quantization::Metric { + use repr::internal::quantization::Metric as IMetric; + match self { + Self::SquaredL2 => IMetric::SquaredL2, + Self::InnerProduct => IMetric::InnerProduct, + Self::Cosine => IMetric::Cosine, + } + } +} + +impl From for Metric { + fn from(m: diskann_vector::distance::Metric) -> Self { + use diskann_vector::distance::Metric as VMetric; + match m { + VMetric::L2 => Self::SquaredL2, + VMetric::InnerProduct => Self::InnerProduct, + VMetric::Cosine => Self::Cosine, + VMetric::CosineNormalized => Self::Cosine, + } + } +} + +/// Choose how data is going to be reranked. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum Rerank { + /// No reranking will be performed and no space for higher precision vectors will be + /// allocated. + None, + + /// Use 16-bit floating point numbers to store the higher precision representation. + /// These will be used automatically during search to rerank candidates. + F16, +} + +impl Rerank { + fn as_internal_rerank(&self) -> repr::internal::quantization::Rerank { + use repr::internal::quantization::Rerank as IRerank; + match self { + Self::None => IRerank::None, + Self::F16 => IRerank::F16, + } + } +} + +/// The configuration for a [`Product`] quantized representation. +#[derive(Debug)] +pub struct Config { + table: tables::BasicTable, + start_points: rowmajor::Owned, + metric: repr::internal::quantization::Metric, + layout: store::Layout, + store: store::Config, + lookahead: Option, + rerank: repr::internal::quantization::Rerank, + thread_hint: Option, +} + +const DEFAULT_LOOKAHEAD: NonZeroUsize = NonZeroUsize::new(16).unwrap(); + +impl Config { + /// Create a new [`Config`]. Parameters will be used as described below: + /// + /// * `table`: The [`tables::BasicTable`] that contains the PQ pivots and chunking strategy. + /// + /// * `metric`: The [`repr::quantization::Metric`] to use for computing distances. + /// + /// * `capacity`: The number points to allocate space for. + /// + /// * `max_degree`: The maximum degree of the internal graph. + /// + /// * `start_points`: The points to use as frozen start points in the index. + /// + /// * `rerank`: Whether or not reranking is enabled and if so, the representation of the + /// higher precision vectors. + /// + /// # Errors + /// + /// Errors under the following conditions: + /// + /// * `start_points.ncols() != table.dim()`: The dimensionality of the start + /// points must agree with the quantizer. + /// + /// * `start_points.nrows() == 0`: Currently, empty start points are not supported. + /// + /// * The number of start points exceeds `u32::MAX`. + pub fn new( + table: tables::BasicTable, + metric: Metric, + capacity: Capacity, + max_degree: MaxDegree, + start_points: rowmajor::Owned, + rerank: Rerank, + ) -> Result { + let dim = table.dim(); + if dim != start_points.ncols() { + return Err(ConfigError::dim_mismatch(dim, start_points.ncols())); + } + + if start_points.nrows() == 0 { + return Err(ConfigError::empty_start_points()); + } + + let num_start_points: u32 = start_points + .nrows() + .try_into() + .map_err(|_| ConfigError::too_many_start_points(start_points.nrows()))?; + + Ok(Self { + table, + start_points, + metric: metric.as_internal_metric(), + layout: store::Layout::new(capacity, max_degree, num_start_points), + store: store::Config::default(), + lookahead: Some(DEFAULT_LOOKAHEAD), + rerank: rerank.as_internal_rerank(), + thread_hint: None, + }) + } + + /// Override the [`store::Config`] for tailoring concurrency details. + pub fn store(mut self, config: store::Config) -> Self { + self.store = config; + self + } + + /// Set the prefetch lookahead. + /// + /// This controls how many iterations ahead in + /// [`diskann::graph::glue::SearchAccessor::expand_beam`] data is prefetched into the CPU + /// cache. Passing `None` disables prefetching. + pub fn prefetch(mut self, lookahead: Option) -> Self { + self.lookahead = lookahead; + self + } + + /// Provide a hint at the number of threads that will be working concurrently. + /// + /// This will be used to pre-allocate distance tables, shortening the time required to + /// begin handling queries. + pub fn thread_hint(mut self, hint: Option) -> Self { + self.thread_hint = hint; + self + } + + /// Build the [`Product`] from `self`. + pub fn build(self) -> ANNResult { + Product::new(self) + } +} + +impl repr::RepresentationConfig for Config { + type Representation = Product; + + fn build(self) -> ANNResult { + ::build(self) + } +} + +/// Errors that can occur during the construction of [`Config`]. +#[derive(Debug, Error)] +#[error(transparent)] +pub struct ConfigError { + inner: ConfigErrorInner, +} + +diskann::convert_error!(ConfigError); + +impl ConfigError { + fn dim_mismatch(quantizer: usize, start_points: usize) -> Self { + Self { + inner: ConfigErrorInner::DimMismatch { + quantizer, + start_points, + }, + } + } + + fn empty_start_points() -> Self { + Self { + inner: ConfigErrorInner::EmptyStartPoints, + } + } + + fn too_many_start_points(num_start_points: usize) -> Self { + Self { + inner: ConfigErrorInner::TooManyStartPoints { num_start_points }, + } + } +} + +#[derive(Debug, Error)] +enum ConfigErrorInner { + #[error( + "quantizer configured for dimension {} but given start points have dimension {}", + quantizer, + start_points + )] + DimMismatch { + quantizer: usize, + start_points: usize, + }, + #[error("at least one start point must be provided")] + EmptyStartPoints, + #[error("{} start points exceeds u32::MAX", num_start_points)] + TooManyStartPoints { num_start_points: usize }, +} + +/// Product quantized data representation. +#[derive(Debug)] +pub struct Product { + store: Store>>, + /// The PQ table representation used for compression and creation of query computers. + transposed: tables::TransposedTable, + /// The PQ table representation used for pruning. + padded: tables::PaddedTable, + metric: repr::internal::quantization::Metric, + lookahead: Option, + reranker: repr::internal::quantization::Reranker, + + /// An object pool for distance lookup tables. + /// + /// These involve non-trivial allocations, so it can be beneficial to hold onto them. + distance_tables: ObjectPool, +} + +impl Product { + pub fn config( + table: tables::BasicTable, + metric: Metric, + capacity: Capacity, + max_degree: MaxDegree, + start_points: rowmajor::Owned, + rerank: Rerank, + ) -> Result { + Config::new(table, metric, capacity, max_degree, start_points, rerank) + } + + fn new(config: Config) -> ANNResult { + let Config { + table, + start_points, + metric, + layout, + store, + lookahead, + rerank, + thread_hint, + } = config; + + let dim = table.dim(); + let (reranker, rerank_config) = + repr::internal::quantization::Reranker::new_with_config(rerank, metric, dim); + + let transposed = tables::TransposedTable::from_parts( + table.view_pivots(), + table.view_offsets().to_owned(), + ) + .map_err(ANNError::new) + .context("this is a broken internal invariant - please report")?; + + let padded = tables::PaddedTable::from_basic(table.as_view()); + + let slots = cons::Config::new( + Intrusive::config(Bytes::new(transposed.nchunks())), + rerank_config, + ); + + let store = Store::new(layout, store, slots)?; + + let this = Self { + store, + transposed, + padded, + metric, + lookahead, + reranker, + distance_tables: ObjectPool::with_capacity(Some( + thread_hint.map(|v| v.get()).unwrap_or(0), + )), + }; + + // Initialize start points. + let num_start_points = start_points.nrows(); + for (i, (slot_index, row)) in + std::iter::zip(this.store.frozen(), start_points.rows()).enumerate() + { + #[expect( + clippy::expect_used, + reason = "failing this is an internal, unrecoverable bug" + )] + let mut slot = this + .store + .slot(slot_index) + .expect("internal store should leave frozen points available for writing"); + + this.set(row, slot.data()).with_context(|| { + lazy_format!(move, "on start point {} of {}", i + 1, num_start_points) + })?; + + slot.freeze(); + } + + Ok(this) + } + + /// Return the dimension of the data held within `self`. + pub fn dim(&self) -> usize { + self.transposed.dim() + } + + fn nchunks(&self) -> usize { + self.transposed.nchunks() + } + + fn ncenters(&self) -> usize { + self.transposed.ncenters() + } + + fn metric(&self) -> repr::internal::quantization::Metric { + self.metric + } + + /// * Attempt to compress `v` into the [`cons::Exclusive::first`] position. + /// * If [`cons::Exclusive::second`] is occupied, use `self.reranker` to store data + /// into that slot. + fn set( + &self, + v: &[f32], + slot: &mut cons::Exclusive, Option>>, + ) -> ANNResult<()> { + use diskann_quantization::CompressInto; + + self.transposed + .compress_into(v, slot.first().as_mut_slice()) + .map_err(ANNError::new)?; + + self.reranker.store(v, slot.second()); + + Ok(()) + } + + fn create_accessor<'a>( + &'a self, + query: &'a [f32], + provider: &'a (dyn std::any::Any + Send + Sync), + counters: LocalCounters<'a>, + args: AccessorArgs, + ) -> ANNResult> { + use diskann_vector::{Norm, norm::FastL2Norm}; + + let AccessorArgs { rerank_if_enabled } = args; + + // Create the query computer. + // + // We do this first because this is one of the most likely things to fail since it + // operates on largely untrusted data. If it does fail, we save the work of acquiring + // epoch guards etc. + let query_computer = { + let distance_table_args = DistanceTableArgs { + nchunks: self.nchunks(), + ncenters: self.ncenters(), + metric: self.metric(), + }; + + let mut distance_table = self.distance_tables.get_ref(distance_table_args); + + match &mut *distance_table { + DistanceTable::SquaredL2(table) => { + self.transposed + .process_into::(query, table.as_view_mut()); + } + DistanceTable::InnerProduct(table) => { + self.transposed + .process_into::( + query, + table.as_view_mut(), + ); + } + DistanceTable::Cosine { table, query_norm } => { + self.transposed + .process_into::(query, table.as_view_mut()); + *query_norm = (FastL2Norm).evaluate(query); + } + } + + QueryComputer(distance_table) + }; + + // Computer is good - time to make `ExpandBeam` and `PostProcess` (if requested). + let (expand_beam, post_process) = self.store.guard(|cons, guard| { + // TODO: Tailor prefetching to the number of cachelines. + // + // Inlining distance functions will require work in `diskann-quantization`, so + // we can at least optimize prefetching. + let expand_beam = repr::internal::intrusive::ExpandBeam::new( + cons.first().reader(guard), + query_computer, + prefetch::Loop::new(), + self.lookahead, + ); + + let post_process = rerank_if_enabled + .then(|| { + self.reranker + .post_process(query, expand_beam.guard(), cons.second(), &counters) + }) + .flatten(); + + (expand_beam, post_process) + })?; + + Ok(crate::provider::SearchAccessor::new( + self.store.neighbors(), + expand_beam.boxed(), + post_process, + provider, + self.store.frozen(), + counters, + )) + } +} + +#[derive(Debug)] +struct AccessorArgs { + rerank_if_enabled: bool, +} + +repr::internal::macros::representation!(Product); + +repr::internal::macros::set_guard!( + /// A [`repr::Guard`] for [`Product`]. + for<'a> cons::Exclusive, Option>> +); + +impl repr::Set<&[f32]> for Product { + type Guard<'a> = Guard<'a>; + + fn set(&self, v: &[f32]) -> ANNResult> { + // Easy check to reject invalid vectors before acquiring an epoch guard. + let vlen = v.len(); + let dim = self.dim(); + if vlen != dim { + return Err(ANNError::message(lazy_format!( + move, + "vector dim {} does not match quantizer dim {}", + vlen, + dim + ))); + } + + let mut slot = self + .store + .acquire() + .ok_or_else(|| ANNError::message("could not allocate a new slot"))?; + + self.set(v, slot.data())?; + Ok(Guard::new(slot)) + } +} + +impl repr::Search for Product { + type Query<'a> = &'a [f32]; + + fn search_accessor<'a>( + &'a self, + query: &'a [f32], + provider: &'a (dyn std::any::Any + Send + Sync), + counters: LocalCounters<'a>, + ) -> ANNResult> { + self.create_accessor( + query, + provider, + counters, + AccessorArgs { + rerank_if_enabled: true, + }, + ) + } +} + +impl repr::Insert for Product { + fn insert_search_accessor<'a>( + &'a self, + query: Self::Query<'a>, + provider: &'a (dyn std::any::Any + Send + Sync), + counters: LocalCounters<'a>, + ) -> ANNResult> { + self.create_accessor( + query, + provider, + counters, + AccessorArgs { + rerank_if_enabled: false, + }, + ) + } + + fn prune_accessor<'a>( + &'a self, + counters: LocalCounters<'a>, + ) -> ANNResult> { + let distance = Distance::new(&self.padded, self.metric()); + let reader = self + .store + .guard(|slots, guard| slots.first().reader(guard))?; + + let prune = repr::internal::intrusive::Prune::new(reader, distance); + Ok(crate::provider::PruneAccessor::new( + prune.boxed(), + self.store.neighbors(), + counters, + )) + } +} + +//----------------// +// Query Distance // +//----------------// + +/// A populated distance table for query distances. +#[derive(Debug)] +enum DistanceTable { + SquaredL2(rowmajor::Owned), + InnerProduct(rowmajor::Owned), + Cosine { + table: rowmajor::Owned, + query_norm: f32, + }, +} + +#[derive(Debug, Clone, Copy)] +struct DistanceTableArgs { + nchunks: usize, + ncenters: usize, + metric: repr::internal::quantization::Metric, +} + +impl object_pool::AsPooled for DistanceTable { + fn create(args: DistanceTableArgs) -> Self { + use repr::internal::quantization::Metric as IMetric; + + let DistanceTableArgs { + nchunks, + ncenters, + metric, + } = args; + + match metric { + IMetric::SquaredL2 => { + DistanceTable::SquaredL2(rowmajor::Owned::from_element(nchunks, ncenters, 0.0)) + } + + IMetric::InnerProduct => { + DistanceTable::InnerProduct(rowmajor::Owned::from_element(nchunks, ncenters, 0.0)) + } + + IMetric::Cosine => DistanceTable::Cosine { + table: rowmajor::Owned::from_element( + nchunks, + ncenters, + tables::lookup::DotAndNorm::default(), + ), + query_norm: 0.0, + }, + } + } + + fn modify(&mut self, args: DistanceTableArgs) { + use repr::internal::quantization::Metric as IMetric; + + let DistanceTableArgs { + nchunks, + ncenters, + metric, + } = args; + + let sizes_agree = + |nrows: usize, ncols: usize| -> bool { nrows == nchunks && ncols == ncenters }; + + let good = match (&self, metric) { + (Self::SquaredL2(table), IMetric::SquaredL2) => { + sizes_agree(table.nrows(), table.ncols()) + } + (Self::InnerProduct(table), IMetric::InnerProduct) => { + sizes_agree(table.nrows(), table.ncols()) + } + (Self::Cosine { table, .. }, IMetric::Cosine) => { + sizes_agree(table.nrows(), table.ncols()) + } + _ => false, + }; + + if !good { + *self = Self::create(args); + } + } +} + +#[derive(Debug)] +struct QueryComputer<'a>(object_pool::PooledRef<'a, DistanceTable>); + +impl repr::internal::RawQueryDistance for QueryComputer<'_> { + type Error = ANNError; + + fn eval(&self, x: &[u8]) -> Result { + let distance = match &*self.0 { + DistanceTable::SquaredL2(table) | DistanceTable::InnerProduct(table) => { + tables::lookup::lookup_single(tables::lookup::Sum, table.as_view(), x) + .map_err(ANNError::new)? + } + DistanceTable::Cosine { table, query_norm } => { + let sum = tables::lookup::lookup_single(tables::lookup::Sum, table.as_view(), x) + .map_err(ANNError::new)?; + + sum.finish_cosine(*query_norm).into_inner() + } + }; + + Ok(distance) + } +} + +//----------// +// Distance // +//----------// + +#[derive(Debug)] +struct Distance<'a> { + table: &'a tables::PaddedTable, + vtable: tables::padded::VTable, +} + +impl<'a> Distance<'a> { + fn new(table: &'a tables::PaddedTable, metric: repr::internal::quantization::Metric) -> Self { + Self { + table, + vtable: table.vtable(metric.into()), + } + } +} + +impl repr::internal::RawDistance for Distance<'_> { + type Error = ANNError; + + fn eval(&self, x: &[u8], y: &[u8]) -> Result { + self.vtable + .self_distance(self.table, x, y) + .map_err(ANNError::new) + } +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use diskann::graph::test::synthetic::Grid; + use diskann_utils::{assert_contains, views::rowmajor::MatrixMut}; + use diskann_vector::distance::DistanceProvider; + use hashbrown::HashMap; + + use crate::{ + counters::Counters, + num::{LogicalId, SlotId}, + repr::test::{Reference, test_expand_beam, test_prune}, + }; + + fn train_quantizer( + data: rowmajor::Ref<'_, f32>, + chunks: usize, + centers: usize, + ) -> tables::BasicTable { + use diskann_quantization::{ + Parallelism, + cancel::DontCancel, + product::{self, train::TrainQuantizer}, + random, + views::ChunkOffsets, + }; + + let trainer = product::train::LightPQTrainingParameters::new(centers, 2); + trainer + .train( + data, + ChunkOffsets::partition( + NonZeroUsize::new(data.ncols()).unwrap(), + NonZeroUsize::new(chunks).unwrap(), + ) + .unwrap() + .as_view(), + Parallelism::Sequential, + &random::StdRngBuilder::new(0), + &DontCancel, + ) + .unwrap() + } + + /// See the description in [`make_test_repr`]. + const TEST_LIMIT: IdLimit = IdLimit::new(11); + + // Use the canonical grid layout, but center the data around the origin. + // + // This allows cosine distances to return reasonable results as the data is distributed + // around the origin. + // + // To keep computation mostly tractable, we only use a 2d grid with 9 points. So the + // coordinates are as follows: + // + // 0: [-1, -1] + // 1: [-1, 0] + // 2: [-1, +1] + // + // 3: [ 0, -1] + // 4: [ 0, 0] + // 5: [ 0, +1] + // + // 6: [+1, -1] + // 7: [+1, 0] + // 8: [+1, +1] + // + // We put two start points at `[-2, -2]` and `[+2, +2]`. + fn make_test_repr(metric: Metric, rerank: Rerank, fill: bool) -> (Product, Reference) { + let grid = Grid::Two; + let mut data = grid.data(3); + let offset = 1.5; + data.as_mut_slice().iter_mut().for_each(|v| *v -= offset); + + let mut start_points = rowmajor::Owned::from_element(2, data.ncols(), 0.0); + start_points.row_mut(0).fill(-2.0); + start_points.row_mut(1).fill(2.0); + + // Train with 2 chunks and 5 centers. This should be sufficient to exactly represent + // all the points in the grid. + // + // To ensure the results are exact, we append the start points to the training data. + let table = { + let train_data = + rowmajor::Owned::from_fn(data.nrows() + start_points.nrows(), data.ncols(), |rc| { + if let Some(start_point_row) = rc.row.checked_sub(data.nrows()) { + *start_points.element(start_point_row, rc.col) + } else { + *data.element(rc.row, rc.col) + } + }); + + train_quantizer(train_data.as_view(), 2, 5) + }; + + let config = Product::config( + table, + metric, + Capacity::new(data.nrows()), + MaxDegree::new(0), + start_points.clone(), + rerank, + ) + .unwrap() + .thread_hint(NonZeroUsize::new(1)); + + let product = config.build().unwrap(); + + assert_eq!(repr::Representation::id_limit(&product), TEST_LIMIT); + assert_eq!(product.dim(), grid.dim().into()); + + let mut reference = Reference::new(grid.dim().into()); + + if fill { + for (i, row) in data.rows().enumerate() { + let guard = repr::Set::set(&product, row).unwrap(); + let id = repr::Guard::id(&guard); + + reference.insert(LogicalId(i), SlotId(id), row); + repr::Guard::publish(guard); + } + } + + // Insert frozen points. + for (slot, point) in product.store.frozen().zip(start_points.rows()) { + reference.insert(LogicalId(slot.into_usize()), SlotId(slot), point); + } + + (product, reference) + } + + /// Here - we don't test the whole `expand_beam` loop. That would be a waste of time and + /// is already tested by other code. + /// + /// Instead we use: + /// + /// * `ExpandBeam::evaluate` to verify that the correct type of distance computer is + /// created. + /// + /// * `Prune` to validate that the correct distance computer is made. + /// + /// * If reranking exists, that reranking works as expected. + fn test_distances( + product: &Product, + reference: &mut Reference, + metric: Metric, + rerank: Rerank, + ctx: &dyn std::fmt::Display, + ) { + assert_eq!(repr::Representation::id_limit(product), TEST_LIMIT, "{ctx}"); + + let query = [10.0, -10.0]; + + // Generate a couple of data points in the dataset. + let i0 = LogicalId(2); + let s0 = reference.slot_id_for(i0); + + let i1 = LogicalId(5); + let s1 = reference.slot_id_for(i1); + + let i2 = LogicalId(10); + let s2 = reference.slot_id_for(i2); + + // Perform Deletes // + let i3_deleted = LogicalId(8); + let s3_deleted = reference.slot_id_for(i3_deleted); + + repr::Representation::retire(product, s3_deleted.value()).unwrap(); + reference.delete(i3_deleted); + + let i4_deleted = LogicalId(4); + let s4_deleted = reference.slot_id_for(i4_deleted); + + repr::Representation::retire(product, s4_deleted.value()).unwrap(); + reference.delete(i4_deleted); + + // Extract Values. + let v0 = &reference[i0]; + let v1 = &reference[i1]; + let v2 = &reference[i2]; + + // For computing distances, we rely on the PQ representation being exact. + let f = >::distance_comparer( + metric.as_internal_metric().as_vector_metric(), + None, + ); + + let distances = HashMap::from_iter([ + (s0, f.call(&query, v0)), + (s1, f.call(&query, v1)), + (s2, f.call(&query, v2)), + ]); + + // Insert + { + let counters = Counters::new(); + let mut sa = repr::Insert::insert_search_accessor( + product, + query.as_slice(), + &(), + counters.local(), + ) + .unwrap(); + + assert!( + sa.get_post_process().is_none(), + "insert accessors should not build a post-processor -- {}", + ctx, + ); + + test_expand_beam( + sa.get_expand_beam(), + TEST_LIMIT, + distances.clone(), + &[s0, s1, s3_deleted, s2, s4_deleted], + ctx, + ); + } + + // Search + { + let counters = Counters::new(); + let mut sa = + repr::Search::search_accessor(product, query.as_slice(), &(), counters.local()) + .unwrap(); + + test_expand_beam( + sa.get_expand_beam(), + TEST_LIMIT, + distances.clone(), + &[s0, s1, s3_deleted, s2, s4_deleted], + ctx, + ); + + if rerank == Rerank::None { + assert!( + sa.get_post_process().is_none(), + "search accessors should not build a post-processor with reranking disabled -- {}", + ctx, + ); + } else { + let post_process = match sa.get_post_process() { + Some(post_process) => post_process, + None => panic!("expected a post processor -- {}", ctx), + }; + + repr::internal::quantization::rerank::test_rerank( + post_process, + distances, + &[s0, s1, s3_deleted, s2, s4_deleted], + ctx, + ); + } + } + + // Prune + { + let v00 = f.call(v0, v0); + let v01 = f.call(v0, v1); + let v02 = f.call(v0, v2); + + let v10 = v01; + let v11 = f.call(v1, v1); + let v12 = f.call(v1, v2); + + let v20 = v02; + let v21 = v12; + let v22 = f.call(v2, v2); + + let distances = HashMap::from_iter([ + ((s0, s0), v00), + ((s0, s1), v01), + ((s0, s2), v02), + ((s1, s0), v10), + ((s1, s1), v11), + ((s1, s2), v12), + ((s2, s0), v20), + ((s2, s1), v21), + ((s2, s2), v22), + ]); + + let counters = Counters::new(); + let mut pa = repr::Insert::prune_accessor(product, counters.local()).unwrap(); + + test_prune( + pa.get_prune(), + distances, + &[s0, s1, s3_deleted, s2, s4_deleted], + ctx, + ); + } + } + + /// The happy-patch entry point. + /// + /// Note that this method does a grid of the supported parameters. Adding new values + /// to these parameters has a multiplicative effect on runtime. + /// + /// This mainly matters for Miri tests. Currently, the Miri test for this function takes + /// about 30 seconds. If it starts to take much longer, this test should be split into + /// multiple entry points for parallelism. + #[test] + fn test_product() { + let metrics = [Metric::SquaredL2, Metric::InnerProduct, Metric::Cosine]; + + let rerank = [Rerank::None, Rerank::F16]; + + for metric in metrics { + for rerank in rerank { + let (spherical, mut reference) = make_test_repr(metric, rerank, true); + + test_distances( + &spherical, + &mut reference, + metric, + rerank, + &format_args!("metric = {:?}, rerank = {:?}", metric, rerank), + ); + } + } + } + + //-------------// + // Error Paths // + //-------------// + + #[test] + fn test_config_dim_mismatch() { + let data = rowmajor::Owned::from_element(2, 5, 1.0f32); + let quantizer = train_quantizer(data.as_view(), 2, 2); + + let start_points = rowmajor::Owned::from_element(1, 6, 0.0f32); // Wrong number of columns + let err = Product::config( + quantizer, + Metric::SquaredL2, + Capacity::new(10), + MaxDegree::new(0), + start_points, + Rerank::None, + ) + .unwrap_err(); + + let msg = err.to_string(); + assert_contains!( + msg, + "quantizer configured for dimension 5 but given start points have dimension 6" + ); + } + + #[test] + fn test_empty_start_points() { + let data = rowmajor::Owned::from_element(2, 5, 1.0f32); + let quantizer = train_quantizer(data.as_view(), 2, 2); + + let start_points = rowmajor::Owned::from_element(0, 5, 0.0f32); // Empty + let err = Product::config( + quantizer, + Metric::SquaredL2, + Capacity::new(10), + MaxDegree::new(0), + start_points, + Rerank::None, + ) + .unwrap_err(); + + let msg = err.to_string(); + assert_contains!(msg, "at least one start point must be provided",); + } + + #[test] + fn test_build_error_uncompressible_query() { + let data = rowmajor::Owned::from_element(2, 5, 1.0f32); + let quantizer = train_quantizer(data.as_view(), 2, 2); + + let start_points = rowmajor::Owned::from_element(1, 5, f32::INFINITY); // Wrong number of columns + let config = Product::config( + quantizer, + Metric::SquaredL2, + Capacity::new(10), + MaxDegree::new(0), + start_points, + Rerank::None, + ) + .unwrap(); + + let err = repr::RepresentationConfig::build(config).unwrap_err(); + let msg = err.to_string(); + assert_contains!( + msg, + "a value of infinity or NaN was observed", + "we tried to compress a start point with infinites in it - this should error" + ); + + assert_contains!( + msg, + "1 of 1", + "error message should contain which query errored", + ); + } + + #[test] + fn test_set_capacity_exhaustion() { + let (product, _) = make_test_repr(Metric::SquaredL2, Rerank::None, true); + + let err = repr::Set::set(&product, &[1.0, 2.0]).unwrap_err(); + assert_contains!(err.to_string(), "could not allocate a new slot",); + } + + #[test] + fn test_set_dim_mismatch() { + let (product, _) = make_test_repr(Metric::SquaredL2, Rerank::None, false); + + let err = repr::Set::set(&product, &[1.0, 2.0, 3.0]).unwrap_err(); + assert_contains!( + err.to_string(), + "vector dim 3 does not match quantizer dim 2", + ); + } + + #[test] + fn test_insert_slot_recovery() { + let (product, _) = make_test_repr(Metric::SquaredL2, Rerank::F16, true); + + // Free up one slot. + repr::Representation::retire(&product, 0).unwrap(); + + // Insert something that is incompressible. + let err = repr::Set::set(&product, &[f32::INFINITY, f32::INFINITY]).unwrap_err(); + assert_contains!(err.to_string(), "infinity"); + + // If we insert again, this should succeed. + let guard = repr::Set::set(&product, &[1.0, 2.0]).unwrap(); + assert_eq!(repr::Guard::id(&guard), 0); + } + + //-------------// + // Object Pool // + //-------------// + + #[test] + fn test_distance_table_as_pooled() { + use object_pool::AsPooled; + use repr::internal::quantization::Metric as IMetric; + + let args = DistanceTableArgs { + nchunks: 10, + ncenters: 4, + metric: IMetric::SquaredL2, + }; + let mut table = DistanceTable::create(args); + + let ptr = if let DistanceTable::SquaredL2(ref table) = table { + assert_eq!(table.nrows(), 10); + assert_eq!(table.ncols(), 4); + table.as_ptr() + } else { + panic!("Unexpected table: {:?}", table); + }; + + // Modify should leave the allocation untouched if it matches. + table.modify(args); + + if let DistanceTable::SquaredL2(ref table) = table { + assert_eq!(table.nrows(), 10); + assert_eq!(table.ncols(), 4); + assert_eq!(table.as_ptr(), ptr); + } else { + panic!("Unexpected table: {:?}", table); + }; + + // Modify works when changing sizes. + let args = DistanceTableArgs { + nchunks: 9, + ncenters: 5, + metric: IMetric::SquaredL2, + }; + table.modify(args); + if let DistanceTable::SquaredL2(ref table) = table { + assert_eq!(table.nrows(), 9); + assert_eq!(table.ncols(), 5); + } else { + panic!("Unexpected table: {:?}", table); + }; + + // Modify changes table type. + let args = DistanceTableArgs { + nchunks: 10, + ncenters: 6, + metric: IMetric::Cosine, + }; + table.modify(args); + if let DistanceTable::Cosine { ref table, .. } = table { + assert_eq!(table.nrows(), 10); + assert_eq!(table.ncols(), 6); + } else { + panic!("Unexpected table: {:?}", table); + }; + + let args = DistanceTableArgs { + nchunks: 10, + ncenters: 6, + metric: IMetric::InnerProduct, + }; + table.modify(args); + if let DistanceTable::InnerProduct(ref table) = table { + assert_eq!(table.nrows(), 10); + assert_eq!(table.ncols(), 6); + } else { + panic!("Unexpected table: {:?}", table); + }; + } +} diff --git a/diskann-inmem/src/repr/spherical.rs b/diskann-inmem/src/repr/spherical.rs index dde4ae6946..540521d00e 100644 --- a/diskann-inmem/src/repr/spherical.rs +++ b/diskann-inmem/src/repr/spherical.rs @@ -10,22 +10,18 @@ use std::num::NonZeroUsize; use diskann::{ANNError, ANNResult, error::ErrorContext, utils::IntoUsize}; use diskann_quantization::{ alloc::{GlobalAllocator, Poly, ScopedAllocator}, - spherical::{SupportedMetric, iface}, + spherical::iface, }; use diskann_utils::{ lazy_format, views::rowmajor::{self, Matrix}, }; -use diskann_vector::distance::{Distance, DistanceProvider}; -use half::f16; use thiserror::Error; use crate::{ counters::LocalCounters, - epoch, num::{Bytes, Capacity, IdLimit, MaxDegree}, - prefetch, - repr::{self, internal::Calf}, + prefetch, repr, store::{ self, Store, cons::{self, Cons}, @@ -35,6 +31,28 @@ use crate::{ }, }; +/// Choose how data is going to be reranked. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum Rerank { + /// No reranking will be performed and no space for higher precision vectors will be + /// allocated. + None, + + /// Use 16-bit floating point numbers to store the higher precision representation. + /// These will be used automatically during search to rerank candidates. + F16, +} + +impl Rerank { + fn as_internal_rerank(&self) -> repr::internal::quantization::Rerank { + use repr::internal::quantization::Rerank as IRerank; + match self { + Self::None => IRerank::None, + Self::F16 => IRerank::F16, + } + } +} + /// The configuration for a [`Spherical`] representation. #[derive(Debug)] pub struct Config { @@ -45,7 +63,7 @@ pub struct Config { layout: store::Layout, store: store::Config, lookahead: Option, - rerank: Rerank, + rerank: repr::internal::quantization::Rerank, } const DEFAULT_LOOKAHEAD: NonZeroUsize = NonZeroUsize::new(16).unwrap(); @@ -105,7 +123,7 @@ impl Config { layout: store::Layout::new(capacity, max_degree, num_start_points), store: store::Config::default(), lookahead: Some(DEFAULT_LOOKAHEAD), - rerank, + rerank: rerank.as_internal_rerank(), }) } @@ -188,129 +206,6 @@ impl repr::RepresentationConfig for Config { } } -/// Choose how data is going to be reranked. -#[derive(Debug, Clone, Copy, PartialEq)] -pub enum Rerank { - /// No reranking will be performed and no space for higher precision vectors will be - /// allocated in [`Spherical`]. - None, - - /// Use 16-bit floating point numbers to store the higher precision representation. - /// These will be used automatically during search to rerank candidates. - F16, -} - -/// Internal representation of [`Rerank`]. -/// -/// This is used for computing distances among the raw values in the auxiliary store. -#[derive(Debug)] -enum Reranker { - None, - F16(Distance), -} - -fn convert_metric(metric: SupportedMetric) -> diskann_vector::distance::Metric { - use diskann_vector::distance::Metric; - - match metric { - SupportedMetric::SquaredL2 => Metric::L2, - SupportedMetric::InnerProduct => Metric::InnerProduct, - SupportedMetric::Cosine => Metric::Cosine, - } -} - -impl Reranker { - /// Construct a new [`Reranker`] and a [`store::slots::SlotsConfig`] for the auxiliary - /// store. - fn new_with_config( - rerank: Rerank, - metric: SupportedMetric, - dim: usize, - ) -> (Self, Option) { - let this = match rerank { - Rerank::None => Self::None, - Rerank::F16 => { - let distance = >::distance_comparer( - convert_metric(metric), - Some(dim), - ); - - Self::F16(distance) - } - }; - - let config = match &this { - Self::None => None, - Self::F16(_) => Some(Simple::config(this.bytes_for(dim))), - }; - - (this, config) - } - - #[expect( - clippy::expect_used, - reason = "the arithmetic should not overflow for the feasible `dim` values" - )] - fn bytes_for(&self, dim: usize) -> Bytes { - match self { - Self::None => Bytes::new(0), - Self::F16(_) => Bytes::new( - dim.checked_mul(2) - .expect("f16 is smaller than the f32 in the quantizer"), - ), - } - } - - /// Create a [`repr::PostProcess`]. - /// - /// This assumes that `simple` has the same dimensions as `self`'s contained distance - /// computation and that `guard` belongs to `simple`. - /// - /// # Pre-conditions - /// - /// This requires that `slots` is the [`store::slots::Slots`] created from the - /// configuration returned in [`Self::new_with_config`]. - fn post_process<'a>( - &'a self, - query: &'a [f32], - guard: &epoch::Guard<'a>, - slots: &'a Optional, - counters: &LocalCounters<'a>, - ) -> Option> { - match (self, slots.slots()) { - (Self::None, None) => None, - (Self::F16(distance), Some(simple)) => { - let distance = repr::full::QueryDistance::new(Calf::Borrowed(query), *distance); - let reader = simple.reader(guard.share()); - let post_process = - repr::internal::simple::Reranker::new(reader, distance, counters.fork()); - Some(Box::new(post_process)) - } - _ => unreachable!("invalid combination of arguments"), - } - } - - /// Store the vector `v` into the raw buffer `buf`. - /// - /// # Pre-conditions - /// - /// `buf` must be consistent with the configuration returned from [`Self::new_with_config`], - /// and may only be `None` if that configuration was `None`. - /// - /// If it is `Some`, this function may panic if its length is not consistent with the - /// original configuration. - fn store(&self, v: &[f32], buf: &mut Option>) { - match (self, buf) { - (Self::None, None) => {} - (Self::F16(_), Some(exclusive)) => { - use diskann_vector::conversion::CastFromSlice; - bytemuck::cast_slice_mut::(exclusive.as_mut_slice()).cast_from_slice(v); - } - _ => unreachable!("invalid combination of arguments"), - } - } -} - /// Spherically quantized data representation. #[derive(Debug)] pub struct Spherical { @@ -320,7 +215,7 @@ pub struct Spherical { // trait-object function call when accessing. full_dim: usize, lookahead: Option, - reranker: Reranker, + reranker: repr::internal::quantization::Reranker, } impl Spherical { @@ -355,8 +250,11 @@ impl Spherical { } = config; let full_dim = quantizer.full_dim(); - let (reranker, rerank_config) = - Reranker::new_with_config(rerank, quantizer.metric(), full_dim); + let (reranker, rerank_config) = repr::internal::quantization::Reranker::new_with_config( + rerank, + quantizer.metric().into(), + full_dim, + ); let slots = cons::Config::new( Intrusive::config(Bytes::new(quantizer.bytes())), @@ -375,14 +273,16 @@ impl Spherical { // Initialize start points. let num_start_points = start_points.nrows(); - for (i, row) in std::iter::zip(this.store.frozen(), start_points.rows()) { + for (i, (slot_index, row)) in + std::iter::zip(this.store.frozen(), start_points.rows()).enumerate() + { #[expect( clippy::expect_used, reason = "failing this is an internal, unrecoverable bug" )] let mut slot = this .store - .slot(i) + .slot(slot_index) .expect("internal store should leave frozen points available for writing"); this.set(row, slot.data()).with_context(|| { @@ -639,14 +539,16 @@ impl repr::internal::RawDistance for &dyn iface::DynDistanceComputer { mod tests { use super::*; - use diskann::{graph::test::synthetic::Grid, neighbor::Neighbor}; + use diskann::graph::test::synthetic::Grid; + use diskann_quantization::spherical::SupportedMetric; use diskann_utils::{assert_contains, views::rowmajor::MatrixMut}; + use diskann_vector::distance::DistanceProvider; use hashbrown::HashMap; use crate::{ counters::Counters, num::{LogicalId, SlotId}, - repr::test::Reference, + repr::test::{Reference, test_expand_beam, test_prune}, }; #[derive(Debug, Clone, Copy)] @@ -757,183 +659,6 @@ mod tests { (spherical, reference) } - /// Performs the following set of tests: - /// - /// * [`ExpandBeam::evaluate`]: For each id in `ids` - attempt to evaluate the distance - /// through [`ExpandBeam::evaluate`]. If the id is present in `distances`, assert that - /// the value in `distances` agrees with the result of the `EpandBeam method. - /// - /// Otherwise, assert that `ExpandBeam` returns `None`. - /// - /// * [`ExpandBeam::expand_beam`]: Provid all `ids` to `expand_beam`. Verify that ids not - /// present in `distances` get removed and all remaining ids are present and have a - /// distance value equal to the corresponding entry in `distances`. - fn test_expand_beam( - accessor: &dyn repr::ExpandBeam, - distances: HashMap, - ids: &[SlotId], - ctx: &dyn std::fmt::Display, - ) { - assert_eq!(accessor.id_limit(), TEST_LIMIT, "{ctx}"); - - for slot_id in ids { - if let Some(distance) = distances.get(slot_id) { - assert_eq!( - accessor.evaluate(slot_id.value()).unwrap(), - Some(*distance), - "failed on slot id {} -- {}", - slot_id, - ctx, - ); - } else { - assert!( - accessor.evaluate(slot_id.value()).unwrap().is_none(), - "failed on slot id {} -- {}", - slot_id, - ctx - ); - } - } - - // Test via `expand_beam`. - let list: Vec = ids.iter().map(|slot_id| slot_id.value()).collect(); - let mut buffer = vec![Neighbor::default(); list.len()]; - let len = repr::safe_expand_beam(accessor, &list, &mut buffer).unwrap(); - - let expected: Vec> = ids - .iter() - .filter_map(|slot_id| { - distances - .get(slot_id) - .map(|distance| Neighbor::new(slot_id.value(), *distance)) - }) - .collect(); - - assert_eq!( - expected.len(), - len, - "`expand_beam` returned the incorrect number of items -- {}", - ctx, - ); - - for (i, (got, expected)) in std::iter::zip(buffer.iter(), expected.iter()).enumerate() { - assert_eq!( - got.id(), - expected.id(), - "failed on entry {} of {} -- {}", - i, - len, - ctx - ); - assert_eq!( - got.distance(), - expected.distance(), - "failed on entry {} of {} -- {}", - i, - len, - ctx, - ); - } - } - - /// Test that the [`repr::Prune`] computes distances according to the ground truth in - /// `distances`. - /// - /// This assumes that `distances` contains all valid (i.e., between undeleted) entries - /// in `ids` - including self distances. - /// - /// For example, if `ids` contains `[0, 1, 2, 3(deleted)]`, then `distances` should contain - /// the keys: - /// - /// (0, 0), (0, 1), (0, 2) - /// (1, 0), (1, 1), (1, 2) - /// (2, 0), (2, 1), (2, 2) - fn test_prune( - accessor: &mut dyn repr::Prune, - distances: HashMap<(SlotId, SlotId), f32>, - ids: &[SlotId], - ctx: &dyn std::fmt::Display, - ) { - let num_present_ids = ids - .iter() - .filter(|&&slot_id| distances.contains_key(&(slot_id, slot_id))) - .count(); - - let mut items: HashMap> = - ids.iter().map(|slot_id| (slot_id.value(), None)).collect(); - - let count = accessor.prepare(items.iter_mut()).unwrap(); - assert_eq!(count, num_present_ids, "{ctx}"); - - let mut visited = 0; - for slot_id0 in ids.iter() { - if let Some(key0) = items[&slot_id0.value()] { - for slot_id1 in ids.iter() { - if let Some(key1) = items[&slot_id1.value()] { - let d = accessor.evaluate(key0, key1); - let expected = distances[&(*slot_id0, *slot_id1)]; - assert_eq!( - d, expected, - "failed for {} x {} -- {}", - slot_id0, slot_id1, ctx - ); - - visited += 1; - } - } - } - } - - assert_eq!( - visited, - distances.len(), - "not all distances were visited -- {}", - ctx - ); - } - - /// Test that [`repr::PostProcess`] reranks correctly. - /// - /// Pass all `ids` to [`repr::PostProcess::post_process`]. Verify that all ids not - /// present in `distances` have been removed and the remaining ids are present, sorted, - /// and have distance values matching those in `distances`. - fn test_rerank( - post_process: &mut dyn repr::PostProcess, - distances: HashMap, - ids: &[SlotId], - ctx: &dyn std::fmt::Display, - ) { - let mut buffer: Vec<_> = ids - .iter() - .map(|slot_id| Neighbor::new(slot_id.value(), 0.0)) - .collect(); - - post_process.post_process(&mut buffer).unwrap(); - let mut previous = f32::NEG_INFINITY; - assert_eq!(buffer.len(), distances.len(), "{ctx}"); - for neighbor in buffer.iter() { - let current = *neighbor.distance(); - assert_eq!(current, distances[&SlotId(*neighbor.id())], "{ctx}",); - - assert!( - current >= previous, - "distances is not monotonically increasing, previous = {}, current = {} -- {}", - previous, - current, - ctx, - ); - - previous = current; - } - - assert!( - previous > f32::NEG_INFINITY, - "previous = {} -- {}", - previous, - ctx - ); - } - /// Here - we don't test the whole `expand_beam` loop. That would be a waste of time and /// is already tested by other code. /// @@ -1046,6 +771,7 @@ mod tests { test_expand_beam( sa.get_expand_beam(), + TEST_LIMIT, distances, &[s0, s1, s3_deleted, s2, s4_deleted], ctx, @@ -1077,6 +803,7 @@ mod tests { test_expand_beam( sa.get_expand_beam(), + TEST_LIMIT, distances, &[s0, s1, s3_deleted, s2, s4_deleted], ctx, @@ -1094,8 +821,10 @@ mod tests { None => panic!("expected a post processor -- {}", ctx), }; - let f = - >::distance_comparer(convert_metric(metric), None); + let f = >::distance_comparer( + repr::internal::quantization::Metric::from(metric).as_vector_metric(), + None, + ); let distances = HashMap::from_iter([ (s0, f.call(&query, v0)), @@ -1103,7 +832,7 @@ mod tests { (s2, f.call(&query, v2)), ]); - test_rerank( + repr::internal::quantization::rerank::test_rerank( post_process, distances, &[s0, s1, s3_deleted, s2, s4_deleted], diff --git a/diskann-inmem/src/repr/test/mod.rs b/diskann-inmem/src/repr/test/mod.rs index e8167babc3..5e968a54ab 100644 --- a/diskann-inmem/src/repr/test/mod.rs +++ b/diskann-inmem/src/repr/test/mod.rs @@ -10,6 +10,12 @@ use crate::{ repr, }; +#[cfg(feature = "quantization")] +use diskann::neighbor::Neighbor; + +#[cfg(feature = "quantization")] +use crate::num::IdLimit; + /// A test distance that simply sums scalar floating point values. #[derive(Debug)] pub(super) struct TestDistance; @@ -170,3 +176,151 @@ impl ReferenceLookup for LogicalId { ) } } + +//-------------// +// Expand Beam // +//-------------// + +/// Performs the following set of tests: +/// +/// * [`ExpandBeam::id_limit`] is equal to `id_limit`. +/// +/// * [`ExpandBeam::evaluate`]: For each id in `ids` - attempt to evaluate the distance +/// through [`ExpandBeam::evaluate`]. If the id is present in `distances`, assert that +/// the value in `distances` agrees with the result of the `EpandBeam method. +/// +/// Otherwise, assert that `ExpandBeam` returns `None`. +/// +/// * [`ExpandBeam::expand_beam`]: Provide all `ids` to `expand_beam`. Verify that ids not +/// present in `distances` get removed and all remaining ids are present and have a +/// distance value equal to the corresponding entry in `distances`. +#[cfg(feature = "quantization")] +pub(super) fn test_expand_beam( + accessor: &dyn repr::ExpandBeam, + id_limit: IdLimit, + distances: HashMap, + ids: &[SlotId], + ctx: &dyn std::fmt::Display, +) { + assert_eq!(accessor.id_limit(), id_limit, "{ctx}"); + + for slot_id in ids { + if let Some(distance) = distances.get(slot_id) { + assert_eq!( + accessor.evaluate(slot_id.value()).unwrap(), + Some(*distance), + "failed on slot id {} -- {}", + slot_id, + ctx, + ); + } else { + assert!( + accessor.evaluate(slot_id.value()).unwrap().is_none(), + "failed on slot id {} -- {}", + slot_id, + ctx + ); + } + } + + // Test via `expand_beam`. + let list: Vec = ids.iter().map(|slot_id| slot_id.value()).collect(); + let mut buffer = vec![Neighbor::default(); list.len()]; + let len = repr::safe_expand_beam(accessor, &list, &mut buffer).unwrap(); + + let expected: Vec> = ids + .iter() + .filter_map(|slot_id| { + distances + .get(slot_id) + .map(|distance| Neighbor::new(slot_id.value(), *distance)) + }) + .collect(); + + assert_eq!( + expected.len(), + len, + "`expand_beam` returned the incorrect number of items -- {}", + ctx, + ); + + for (i, (got, expected)) in std::iter::zip(buffer.iter(), expected.iter()).enumerate() { + assert_eq!( + got.id(), + expected.id(), + "failed on entry {} of {} -- {}", + i, + len, + ctx + ); + assert_eq!( + got.distance(), + expected.distance(), + "failed on entry {} of {} -- {}", + i, + len, + ctx, + ); + } +} + +//-------// +// Prune // +//-------// + +/// Test that the [`repr::Prune`] computes distances according to the ground truth in +/// `distances`. +/// +/// This assumes that `distances` contains all valid (i.e., between undeleted) entries +/// in `ids` - including self distances. +/// +/// For example, if `ids` contains `[0, 1, 2, 3(deleted)]`, then `distances` should contain +/// the keys: +/// +/// (0, 0), (0, 1), (0, 2) +/// (1, 0), (1, 1), (1, 2) +/// (2, 0), (2, 1), (2, 2) +#[cfg(feature = "quantization")] +pub(super) fn test_prune( + accessor: &mut dyn repr::Prune, + distances: HashMap<(SlotId, SlotId), f32>, + ids: &[SlotId], + ctx: &dyn std::fmt::Display, +) { + let num_present_ids = ids + .iter() + .filter(|&&slot_id| distances.contains_key(&(slot_id, slot_id))) + .count(); + + let mut items: HashMap> = + ids.iter().map(|slot_id| (slot_id.value(), None)).collect(); + + let count = accessor.prepare(items.iter_mut()).unwrap(); + assert_eq!(count, num_present_ids, "{ctx}"); + + let mut visited = 0; + for slot_id0 in ids.iter() { + if let Some(key0) = items[&slot_id0.value()] { + for slot_id1 in ids.iter() { + if let Some(key1) = items[&slot_id1.value()] { + let d = accessor.evaluate(key0, key1); + let expected = distances[&(*slot_id0, *slot_id1)]; + assert_eq!( + d, expected, + "failed for {} x {} -- {}", + slot_id0, slot_id1, ctx + ); + + visited += 1; + } + } + } + } + + assert_eq!( + visited, + distances.len(), + "not all distances were visited -- {}", + ctx + ); +} diff --git a/diskann-quantization/src/product/train.rs b/diskann-quantization/src/product/train.rs index 9b64cf777b..c559d2782f 100644 --- a/diskann-quantization/src/product/train.rs +++ b/diskann-quantization/src/product/train.rs @@ -19,7 +19,7 @@ use crate::{ algorithms::kmeans::{self, common::square_norm}, cancel::Cancelation, multi_vector::BlockTransposed, - product::tables::BasicTable, + product::BasicTable, random::{BoxedRngBuilder, RngBuilder}, }; diff --git a/diskann-utils/src/object_pool.rs b/diskann-utils/src/object_pool.rs index 7ed447f02a..4db9c964f3 100644 --- a/diskann-utils/src/object_pool.rs +++ b/diskann-utils/src/object_pool.rs @@ -53,6 +53,14 @@ impl ObjectPool { } } + /// Create an empty [`ObjectPool`] with the requested capacity. + pub fn with_capacity(capacity: Option) -> Self { + Self { + queue: Mutex::new(VecDeque::new()), + capacity, + } + } + /// Create an object pool consisting of `initial_size` object initialized using /// [`TryAsPooled::try_create`]. /// @@ -572,6 +580,35 @@ mod tests { assert_eq!(*item.value, 100); } + #[test] + fn test_pool_with_capacity() { + let pool = ObjectPool::::with_capacity(Some(4)); + assert!(pool.is_empty()); + + let _ = pool.get_ref(100); + assert_eq!(pool.len(), 1); + + { + let _v0 = pool.get_ref(100); + let _v1 = pool.get_ref(100); + } + assert_eq!(pool.len(), 2); + + { + let _v0 = pool.get_ref(100); + let _v1 = pool.get_ref(100); + let _v2 = pool.get_ref(100); + let _v3 = pool.get_ref(100); + let _v4 = pool.get_ref(100); + let _v5 = pool.get_ref(100); + } + assert_eq!( + pool.len(), + 4, + "pool should max out at the requested capacity." + ); + } + #[test] fn test_pool_basic_tests_with_try() { // Create a pool with negative initial size to test error case