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
3 changes: 2 additions & 1 deletion diskann-benchmark/example/product-exhaustive.json
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,8 @@
"compression_threads": 1,
"seed": 7831252621480178695,
"num_pq_chunks": 16,
"num_pq_centers": 16
"num_pq_centers": 16,
"table_style": "transposed"
}
}
]
Expand Down
260 changes: 237 additions & 23 deletions diskann-benchmark/src/exhaustive/product.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,8 +34,13 @@ mod imp {
utils::{percentiles, MicroSeconds},
Benchmark, Output,
};
use diskann_quantization::{product::train::TrainQuantizer, CompressInto};
use diskann_providers::model::pq::FixedChunkPQTable;
use diskann_quantization::{
product::{tables, train::TrainQuantizer},
CompressInto,
};
use diskann_utils::views::rowmajor::{self, Matrix, MatrixMut};
use diskann_vector::distance::Metric;
use indicatif::{ProgressBar, ProgressStyle};
use rayon::iter::{IndexedParallelIterator, ParallelIterator};
use serde::Serialize;
Expand Down Expand Up @@ -94,7 +99,7 @@ mod imp {
let offsets =
diskann_quantization::views::ChunkOffsets::partition(dim, input.num_pq_chunks)?;

let base = {
let table = {
let threadpool = rayon::ThreadPoolBuilder::new()
.num_threads(input.compression_threads.get())
.build()?;
Expand All @@ -109,12 +114,6 @@ mod imp {
})?
};

let quantizer = diskann_providers::model::pq::FixedChunkPQTable::new(
data.ncols(),
base.flatten().into(),
offsets.as_slice().into(),
)?;

let training_time: MicroSeconds = start.elapsed().into();

// Compressing
Expand All @@ -126,8 +125,14 @@ mod imp {

let compression_progress =
make_progress_bar("compressing", data.nrows(), output.draw_target())?;
let store = threadpool
.install(|| Store::new(data.as_view(), quantizer, &compression_progress))?;
let store = threadpool.install(|| {
Store::new(
data.as_view(),
table,
input.table_style,
&compression_progress,
)
})?;
compression_progress.finish();
store
};
Expand Down Expand Up @@ -321,32 +326,246 @@ mod imp {
}
}

//------------//
// Compressor //
//------------//

#[derive(Debug)]
enum Compressor {
FixedChunk(FixedChunkPQTable),
Transposed(tables::TransposedTable),
}

impl Compressor {
fn new(
table: tables::BasicTable,
style: inputs::exhaustive::PQTableStyle,
) -> anyhow::Result<Self> {
use inputs::exhaustive::PQTableStyle;
match style {
PQTableStyle::FixedChunk => Ok(Self::FixedChunk(table.try_into()?)),
PQTableStyle::Padded | PQTableStyle::Transposed => {
let table = tables::TransposedTable::from_parts(
table.view_pivots(),
table.view_offsets().to_owned(),
)?;

Ok(Self::Transposed(table))
}
}
}

fn compress(&self, storage: &mut [u8], data: &[f32]) -> anyhow::Result<()> {
match self {
Self::FixedChunk(table) => table.compress_into(data, storage)?,
Self::Transposed(table) => table.compress_into(data, storage)?,
}
Ok(())
}
}

//-----------//
// Distances //
//-----------//

trait ComputerImpl: std::fmt::Debug {
fn evaluate(&self, x: &[u8]) -> anyhow::Result<f32>;
}

#[derive(Debug)]
struct Computer<'a>(Box<dyn ComputerImpl + 'a>);

impl<'a> Computer<'a> {
fn new<C>(inner: C) -> Self
where
C: ComputerImpl + 'a,
{
Self(Box::new(inner))
}
}

impl diskann_vector::PreprocessedDistanceFunction<&[u8], f32> for Computer<'_> {
fn evaluate_similarity(&self, x: &[u8]) -> f32 {
match self.0.evaluate(x) {
Ok(v) => v,
Err(err) => panic!("distance failed with {:#}", err),
}
}
}

//-----------//
// Computers //
//-----------//

impl ComputerImpl for diskann_providers::model::pq::distance::QueryComputer<'_> {
fn evaluate(&self, x: &[u8]) -> anyhow::Result<f32> {
Ok(<Self as diskann_vector::PreprocessedDistanceFunction<
&[u8],
f32,
>>::evaluate_similarity(self, x))
}
}

#[derive(Debug)]
struct PaddedComputer<'a> {
table: &'a tables::PaddedTable,
vtable: tables::padded::VTable,
query: Vec<f32>,
}

impl ComputerImpl for PaddedComputer<'_> {
fn evaluate(&self, x: &[u8]) -> anyhow::Result<f32> {
Ok(self.vtable.distance(self.table, &self.query, x)?)
}
}

#[derive(Debug)]
struct LookupTable {
lookup: rowmajor::Owned<f32>,
}

impl ComputerImpl for LookupTable {
fn evaluate(&self, x: &[u8]) -> anyhow::Result<f32> {
Ok(tables::lookup::lookup_single(
tables::lookup::Sum,
self.lookup.as_view(),
x,
)?)
}
}

#[derive(Debug)]
struct CosineLookupTable {
lookup: rowmajor::Owned<tables::lookup::DotAndNorm>,
query_norm: f32,
}

impl ComputerImpl for CosineLookupTable {
fn evaluate(&self, x: &[u8]) -> anyhow::Result<f32> {
let partial =
tables::lookup::lookup_single(tables::lookup::Sum, self.lookup.as_view(), x)?;
Ok(partial.finish_cosine(self.query_norm).into_inner())
}
}

#[derive(Debug)]
enum Distance {
FixedChunk(FixedChunkPQTable),
Padded(tables::PaddedTable),
Transposed(tables::TransposedTable),
}

impl Distance {
fn new(
basic: tables::BasicTable,
style: inputs::exhaustive::PQTableStyle,
) -> anyhow::Result<Self> {
use inputs::exhaustive::PQTableStyle;
match style {
PQTableStyle::FixedChunk => Ok(Self::FixedChunk(basic.try_into()?)),
PQTableStyle::Padded => Ok(Self::Padded(tables::PaddedTable::from_basic(
basic.as_view(),
))),
PQTableStyle::Transposed => {
Ok(Self::Transposed(tables::TransposedTable::from_parts(
basic.view_pivots(),
basic.view_offsets().to_owned(),
)?))
}
}
}

fn computer(&self, query: &[f32], metric: Metric) -> anyhow::Result<Computer<'_>> {
match self {
Self::FixedChunk(table) => {
let inner = diskann_providers::model::pq::distance::QueryComputer::new(
table.into(),
metric,
query,
None,
)?;
Ok(Computer::new(inner))
}
Self::Padded(table) => {
let inner = PaddedComputer {
table,
vtable: table.vtable(metric.into()),
query: query.into(),
};

Ok(Computer::new(inner))
}
Self::Transposed(table) => match metric {
Metric::L2 => {
let mut lookup =
rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0);
table.process_into::<diskann_quantization::distances::SquaredL2, _>(
query,
lookup.as_view_mut(),
);
Ok(Computer::new(LookupTable { lookup }))
}
Metric::InnerProduct => {
let mut lookup =
rowmajor::Owned::from_element(table.nchunks(), table.ncenters(), 0.0);
table.process_into::<diskann_quantization::distances::InnerProduct, _>(
query,
lookup.as_view_mut(),
);
Ok(Computer::new(LookupTable { lookup }))
}
Metric::Cosine | Metric::CosineNormalized => {
let mut lookup = rowmajor::Owned::from_element(
table.nchunks(),
table.ncenters(),
tables::lookup::DotAndNorm::default(),
);

let query_norm = <_ as diskann_vector::Norm<&[f32]>>::evaluate(
&diskann_vector::norm::FastL2Norm,
query,
);

table.process_into::<diskann_quantization::distances::Cosine, _>(
query,
lookup.as_view_mut(),
);
Ok(Computer::new(CosineLookupTable { lookup, query_norm }))
}
},
}
}
}

/// A store for quantized data.
pub(super) struct Store {
data: rowmajor::Owned<u8>,
quantizer: diskann_providers::model::pq::FixedChunkPQTable,
distance: Distance,
}

impl Store {
fn new(
input: rowmajor::Ref<f32>,
quantizer: diskann_providers::model::pq::FixedChunkPQTable,
table: tables::BasicTable,
style: inputs::exhaustive::PQTableStyle,
progress: &ProgressBar,
) -> anyhow::Result<Self> {
let mut data =
rowmajor::Owned::try_from_element(input.nrows(), quantizer.get_num_chunks(), 0)?;
let mut data = rowmajor::Owned::try_from_element(input.nrows(), table.nchunks(), 0)?;

let compressor = Compressor::new(table.clone(), style)?;

// Compress the data.
#[expect(clippy::disallowed_methods)]
data.par_rows_mut().zip(input.par_rows()).try_for_each(
|(d, i)| -> anyhow::Result<()> {
quantizer.compress_into(i, d)?;
compressor.compress(d, i)?;
progress.inc(1);
Ok(())
},
)?;

Ok(Self { data, quantizer })
let distance = Distance::new(table, style)?;
Ok(Self { data, distance })
}
}

Expand All @@ -366,19 +585,14 @@ mod imp {
}

impl algos::CreateQuantComputer<Store> for Plan {
type Computer<'a> = diskann_providers::model::pq::distance::QueryComputer<'a>;
type Computer<'a> = Computer<'a>;

fn create_quant_computer<'a>(
&self,
store: &'a Store,
query: &[f32],
) -> anyhow::Result<Self::Computer<'a>> {
Ok(diskann_providers::model::pq::distance::QueryComputer::new(
(&store.quantizer).into(),
self.measure.into(),
query,
None,
)?)
store.distance.computer(query, self.measure.into())
}
}
}
Loading
Loading