use diskann::ANNResult;
use diskann_vector::distance::Metric;
use diskann_providers::model::compute_pq_distance;
use diskann_providers::utils::BridgeErr;
use super::{PQData, PQScratch};
use crate::storage::quant::pq::pq_dataset::PQTable;
pub fn quantizer_preprocess(
pq_scratch: &mut PQScratch,
pq_data: &PQData,
metric: Metric,
id_to_calculate_pq_distance: &[u32],
) -> ANNResult<()> {
match &pq_data.pq_table() {
PQTable::Transposed(table) => {
let dim = table.dim();
let expected_len = table.ncenters() * table.nchunks();
let dst = diskann_utils::views::MutMatrixView::try_from(
&mut (*pq_scratch.aligned_pqtable_dist_scratch)[..expected_len],
table.nchunks(),
table.ncenters(),
)
.bridge_err()?;
match metric {
Metric::L2 | Metric::Cosine | Metric::CosineNormalized => {
table.process_into::<diskann_quantization::distances::SquaredL2>(
&pq_scratch.rotated_query[..dim],
dst,
);
}
Metric::InnerProduct => {
table.process_into::<diskann_quantization::distances::InnerProduct>(
&pq_scratch.rotated_query[..dim],
dst,
);
}
}
}
PQTable::Fixed(table) => {
match metric {
Metric::L2 | Metric::Cosine | Metric::CosineNormalized => {
let dim = table.get_dim();
table.preprocess_query(&mut pq_scratch.rotated_query[..dim]);
table.populate_chunk_distances(
pq_scratch.rotated_query.as_slice(),
&mut pq_scratch.aligned_pqtable_dist_scratch,
)?;
}
Metric::InnerProduct => {
table.populate_chunk_inner_products(
pq_scratch.rotated_query.as_slice(),
&mut pq_scratch.aligned_pqtable_dist_scratch,
)?;
}
}
}
}
compute_pq_distance(
id_to_calculate_pq_distance,
pq_data.get_num_chunks(),
&pq_scratch.aligned_pqtable_dist_scratch,
pq_data.pq_compressed_data().as_slice(),
&mut pq_scratch.aligned_pq_coord_scratch,
&mut pq_scratch.aligned_dist_scratch,
)?;
Ok(())
}