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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
240 changes: 240 additions & 0 deletions diskann-benchmark/src/index/inmem2.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand All @@ -56,7 +57,9 @@ pub(crate) fn register_benchmarks(registry: &mut Registry) -> anyhow::Result<()>
registry.register("inmem2-f32", Build::<f32>::new())?;
registry.register("inmem2-f16", Build::<f16>::new())?;
registry.register("inmem2-u8", Build::<u8>::new())?;

registry.register("inmem2-spherical", SphericalBuild)?;
registry.register("inmem2-pq", ProductBuild)?;

registry.register("inmem2-f32-stream", StreamingBenchmark::<f32>::new())?;
Ok(())
Expand Down Expand Up @@ -109,6 +112,7 @@ mod dto {
#[serde(rename_all = "kebab-case")]
pub(super) enum Quantization {
Spherical(Spherical),
Product(Product),
}

#[derive(Debug, Serialize, Deserialize)]
Expand All @@ -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 //
//-----------//
Expand Down Expand Up @@ -357,6 +368,7 @@ impl Display for BuildParams {
enum Quantization {
None,
Spherical(Spherical),
Product(Product),
}

impl Quantization {
Expand All @@ -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
Expand All @@ -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 {
Expand All @@ -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)
}
}
}
}
Expand Down Expand Up @@ -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<BasicTable> {
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 //
//-------------//
Expand Down Expand Up @@ -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<rowmajor::Owned<f32>> = 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::<u32>::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<rowmajor::Owned<f32>> = 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<u32>,
groundtruth: &dyn benchmark_core::recall::Rows<u32>,
Expand Down
1 change: 1 addition & 0 deletions diskann-inmem/integration/index/object.rs
Original file line number Diff line number Diff line change
Expand Up @@ -237,3 +237,4 @@ macro_rules! index {

index!({ T } repr::Full<T> where T: repr::FullPrecision + FromSlice + AsDataType);
index!(repr::Spherical);
index!(repr::Product);
Loading
Loading