use core::fmt::Debug;
use diskann::{ANNError, ANNResult};
use diskann_providers::model::FixedChunkPQTable;
use diskann_quantization::product::TransposedTable;
use diskann_utils::views::Matrix;
#[derive(Debug)]
pub enum PQTable {
Transposed(TransposedTable),
Fixed(FixedChunkPQTable),
}
#[derive(Debug)]
pub struct PQData {
pq_pivot_table: PQTable,
pq_compressed_data: Matrix<u8>,
}
impl PQData {
pub fn new(
pq_pivot_table: FixedChunkPQTable,
pq_compressed_data: Matrix<u8>,
) -> ANNResult<Self> {
let centroid_is_zero = pq_pivot_table.get_centroids().iter().all(|i| *i == 0.0);
let pq_pivot_table = if centroid_is_zero {
let transposed = TransposedTable::from_parts(
pq_pivot_table.view_pivots(),
pq_pivot_table.view_offsets().to_owned(),
)
.map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?;
PQTable::Transposed(transposed)
} else {
PQTable::Fixed(pq_pivot_table)
};
Ok(Self {
pq_pivot_table,
pq_compressed_data,
})
}
pub fn pq_table(&self) -> &PQTable {
&self.pq_pivot_table
}
pub fn get_num_chunks(&self) -> usize {
match &self.pq_pivot_table {
PQTable::Transposed(table) => table.nchunks(),
PQTable::Fixed(table) => table.get_num_chunks(),
}
}
pub fn get_num_centers(&self) -> usize {
match &self.pq_pivot_table {
PQTable::Transposed(table) => table.ncenters(),
PQTable::Fixed(table) => table.get_num_centers(),
}
}
pub fn pq_compressed_data(&self) -> &Matrix<u8> {
&self.pq_compressed_data
}
pub fn get_compressed_vector(&self, vector_id: usize) -> ANNResult<&[u8]> {
self.pq_compressed_data.get_row(vector_id).ok_or_else(|| {
ANNError::log_index_error("Vector id is out of boundary in the compressed dataset.")
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn create_pq_data() -> ANNResult<PQData> {
let dim = 2;
let pq_pivot_table = FixedChunkPQTable::new(
dim,
Box::new([0.0, 0.0, 1.0, 1.0]),
Box::new([0.0, 0.0]),
Box::new([0, 2]),
)
.unwrap();
let pq_compressed_data = Matrix::try_from(Box::new([123u8, 111, 255]) as Box<[u8]>, 3, 1)
.expect("valid matrix shape");
PQData::new(pq_pivot_table, pq_compressed_data)
}
#[test]
fn test_get_compressed_vector() {
let dataset = create_pq_data().unwrap();
let vector_id = 0;
let result = dataset.get_compressed_vector(vector_id).unwrap();
assert_eq!(result, &[123]);
let vector_id = 1;
let result = dataset.get_compressed_vector(vector_id).unwrap();
assert_eq!(result, &[111]);
let vector_id = 2;
let result = dataset.get_compressed_vector(vector_id).unwrap();
assert_eq!(result, &[255]);
}
#[test]
fn test_get_num_chunks() {
let dataset = create_pq_data().unwrap();
assert_eq!(dataset.get_num_chunks(), 1);
}
#[test]
fn test_get_num_centers() {
let dataset = create_pq_data().unwrap();
assert_eq!(dataset.get_num_centers(), 2);
}
}