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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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
264 changes: 244 additions & 20 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 @@ -109,10 +114,14 @@ mod imp {
})?
};

let quantizer = diskann_providers::model::pq::FixedChunkPQTable::new(
data.ncols(),
base.flatten().into(),
offsets.as_slice().into(),
// TODO: Training should return a `BasicTable` directly.
let table = tables::BasicTable::new(
rowmajor::Owned::try_from_data(
base.flatten().into(),
input.num_pq_centers.get(),
data.ncols(),
)?,
offsets,
)?;

let training_time: MicroSeconds = start.elapsed().into();
Expand All @@ -126,8 +135,14 @@ mod imp {

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

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

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

impl Compressor {
fn new(
table: tables::BasicTable,
style: inputs::exhaustive::PQTableStyle,
) -> anyhow::Result<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 +595,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())
}
}
}
21 changes: 21 additions & 0 deletions diskann-benchmark/src/inputs/exhaustive.rs
Original file line number Diff line number Diff line change
Expand Up @@ -200,6 +200,24 @@ impl From<&TransformKind> for diskann_quantization::algorithms::transforms::Tran
// Product Quantization Methods //
//////////////////////////////////

#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
pub(crate) enum PQTableStyle {
FixedChunk,
Padded,
Transposed,
}

impl PQTableStyle {
fn as_str(&self) -> &'static str {
match self {
PQTableStyle::FixedChunk => "fixed-chunk",
PQTableStyle::Padded => "padded",
PQTableStyle::Transposed => "transposed",
}
}
}

#[derive(Debug, Serialize, Deserialize)]
pub(crate) struct Product {
pub(crate) data: InputFile,
Expand All @@ -210,6 +228,7 @@ pub(crate) struct Product {
pub(crate) seed: u64,
pub(crate) num_pq_chunks: NonZeroUsize,
pub(crate) num_pq_centers: NonZeroUsize,
pub(crate) table_style: PQTableStyle,
Comment thread
hildebrandmw marked this conversation as resolved.
}

impl Product {
Expand Down Expand Up @@ -251,6 +270,7 @@ impl Example for Product {
seed: 0x6cae32c479ac3407,
num_pq_chunks: NUM_PQ_CHUNKS,
num_pq_centers: NUM_PQ_CENTERS,
table_style: PQTableStyle::FixedChunk,
}
}
}
Expand All @@ -264,6 +284,7 @@ impl std::fmt::Display for Product {
write_field!(f, "seed", self.seed)?;
write_field!(f, "PQ Chunks", self.num_pq_chunks.get())?;
write_field!(f, "PQ Centers", self.num_pq_centers.get())?;
write_field!(f, "Table Style", self.table_style.as_str())?;
Ok(())
}
}
Expand Down
Loading
Loading