use std::marker::PhantomData;
use diskann::{utils::VectorRepr, ANNError};
use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
use diskann_providers::{
model::{
pq::{accum_row_inplace, generate_pq_pivots},
GeneratePivotArguments,
},
storage::PQStorage,
utils::{BridgeErr, RayonThreadPoolRef, Timer},
};
use diskann_quantization::{product::TransposedTable, CompressInto};
use diskann_utils::views::MatrixBase;
use diskann_vector::distance::Metric;
use tracing::info;
use crate::storage::quant::compressor::{CompressionStage, QuantCompressor};
pub struct PQGenerationContext<'a, Storage>
where
Storage: StorageReadProvider + StorageWriteProvider,
{
pub pq_storage: PQStorage,
pub num_chunks: usize,
pub seed: Option<u64>,
pub p_val: f64,
pub storage_provider: &'a Storage,
pub pool: RayonThreadPoolRef<'a>,
pub metric: Metric,
pub dim: usize,
pub max_kmeans_reps: usize,
pub num_centers: usize,
}
pub struct PQGeneration<'a, T, Storage>
where
T: VectorRepr,
Storage: StorageReadProvider + StorageWriteProvider + 'a,
{
table: TransposedTable,
num_chunks: usize,
phantom_data: PhantomData<T>,
phantom_storage: PhantomData<&'a Storage>,
}
impl<'a, T, Storage> QuantCompressor<T> for PQGeneration<'a, T, Storage>
where
T: VectorRepr,
Storage: StorageReadProvider + StorageWriteProvider + 'a,
{
type CompressorContext = PQGenerationContext<'a, Storage>;
fn new_at_stage(
stage: CompressionStage,
context: &Self::CompressorContext,
) -> diskann::ANNResult<Self> {
if context.num_chunks > context.dim {
return Err(ANNError::log_pq_error(
"Error: number of chunks more than dimension.",
));
}
let pivots_exists = context
.pq_storage
.pivot_data_exist(context.storage_provider);
let pool = context.pool;
if !pivots_exists {
if stage == CompressionStage::Resume {
return Err(ANNError::log_pq_error(
"Error: Pivot data does not exist when start_vertex_id is not 0.",
));
}
let timer = Timer::new();
let rng =
diskann_providers::utils::create_rnd_provider_from_optional_seed(context.seed);
let (mut train_data, train_size, train_dim) = context
.pq_storage
.get_random_train_data_slice::<T, Storage>(
context.p_val,
context.storage_provider,
&mut rng.create_rnd(),
)?;
generate_pq_pivots(
GeneratePivotArguments::new(
train_size,
train_dim,
context.num_centers,
context.num_chunks,
context.max_kmeans_reps,
context.metric == Metric::L2,
)?,
&mut train_data,
&context.pq_storage,
context.storage_provider,
rng,
pool,
)?;
info!(
"PQ pivot generation took {} seconds",
timer.elapsed().as_secs_f64()
);
}
let (_, full_dim) = context
.pq_storage
.read_existing_pivot_metadata(context.storage_provider)?;
let num_chunks = context.num_chunks;
let (mut full_pivot_data, centroid, chunk_offsets) =
context.pq_storage.load_existing_pivot_data(
&num_chunks,
&context.num_centers,
&full_dim,
context.storage_provider,
)?;
let mut full_pivot_data_mat = diskann_utils::views::MutMatrixView::try_from(
full_pivot_data.as_mut_slice(),
context.num_centers,
full_dim,
)
.bridge_err()?;
accum_row_inplace(full_pivot_data_mat.as_mut_view(), centroid.as_slice());
let table = TransposedTable::from_parts(
full_pivot_data_mat.as_view(),
diskann_quantization::views::ChunkOffsetsView::new(&chunk_offsets)
.bridge_err()?
.to_owned(),
)
.map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))?;
Ok(Self {
table,
num_chunks,
phantom_data: PhantomData,
phantom_storage: PhantomData,
})
}
fn compress(
&self,
vector: MatrixBase<&[f32]>,
output: MatrixBase<&mut [u8]>,
) -> Result<(), diskann::ANNError> {
self.table
.compress_into(vector, output)
.map_err(|err| ANNError::log_pq_error(diskann_quantization::error::format(&err)))
}
fn compressed_bytes(&self) -> usize {
self.num_chunks
}
}
#[cfg(test)]
mod pq_generation_tests {
use diskann::ANNError;
use diskann_providers::model::pq::generate_pq_pivots;
use diskann_providers::model::GeneratePivotArguments;
use diskann_providers::storage::{
PQStorage, StorageReadProvider, StorageWriteProvider, VirtualStorageProvider,
};
use diskann_providers::utils::{create_thread_pool_for_test, RayonThreadPoolRef};
use diskann_utils::{
io::{read_bin, write_bin},
test_data_root,
views::{MatrixView, MutMatrixView},
};
use diskann_vector::distance::Metric;
use rstest::rstest;
use vfs::FileSystem;
use super::{CompressionStage, PQGeneration, PQGenerationContext};
use crate::storage::quant::compressor::QuantCompressor;
const TEST_PQ_DATA_PATH: &str = "/sift/siftsmall_learn.bin";
const TEST_PQ_PIVOTS_PATH: &str = "/sift/siftsmall_learn_pq_pivots.bin";
const TEST_PQ_COMPRESSED_PATH: &str = "/sift/siftsmall_learn_pq_compressed.bin";
const VALIDATION_DATA: [f32; 40] = [
1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 1.0f32, 2.0f32, 2.0f32, 2.0f32,
2.0f32, 2.0f32, 2.0f32, 2.0f32, 2.0f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32, 2.1f32,
2.1f32, 2.1f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 2.2f32, 100.0f32,
100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32, 100.0f32,
];
#[allow(clippy::too_many_arguments)]
fn create_new_compressor<'a, F: vfs::FileSystem>(
stage: CompressionStage,
provider: &'a VirtualStorageProvider<F>,
dim: usize,
num_chunks: usize,
max_kmeans_reps: usize,
num_centers: usize,
p_val: f64,
pool: RayonThreadPoolRef<'a>,
pivots_path: String,
compressed_path: String,
data_path: Option<&str>,
) -> Result<PQGeneration<'a, f32, VirtualStorageProvider<F>>, ANNError> {
let pq_storage = PQStorage::new(&pivots_path, &compressed_path, data_path);
let context = PQGenerationContext::<'_, _> {
pq_storage,
num_chunks,
num_centers,
seed: Some(42),
p_val,
max_kmeans_reps,
storage_provider: provider,
pool,
metric: Metric::L2,
dim,
};
PQGeneration::<_, _>::new_at_stage(stage, &context)
}
#[rstest]
fn test_create_and_load_pivots_file() {
let storage_provider = VirtualStorageProvider::new_memory();
storage_provider
.filesystem()
.create_dir("/pq_generation_tests")
.expect("Could not create test directory");
let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin";
let pivot_file_name_compressor = "/pq_generation_tests/compressor_pivots_test.bin";
let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin";
let data_path = "/pq_generation_tests/data_path.bin";
let pq_storage: PQStorage =
PQStorage::new(pivot_file_name, compressed_file_name, Some(data_path));
let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5);
let mut train_data: Vec<f32> = VALIDATION_DATA.to_vec();
write_bin(
MatrixView::try_from(train_data.as_slice(), ndata, dim).unwrap(),
&mut storage_provider.create_for_write(data_path).unwrap(),
)
.unwrap();
let pool = create_thread_pool_for_test();
generate_pq_pivots(
GeneratePivotArguments::new(
ndata,
dim,
num_centers,
num_chunks,
max_k_means_reps,
true,
)
.unwrap(),
&mut train_data,
&pq_storage,
&storage_provider,
diskann_providers::utils::create_rnd_provider_from_seed_in_tests(42),
pool.as_ref(),
)
.unwrap();
let compressor = create_new_compressor(
CompressionStage::Start,
&storage_provider,
dim,
num_chunks,
max_k_means_reps,
num_centers,
1.0, pool.as_ref(),
pivot_file_name_compressor.to_string(),
compressed_file_name.to_string(),
Some(data_path),
);
assert!(compressor.is_ok());
let compressor = compressor.unwrap();
assert_eq!(compressor.num_chunks, num_chunks);
assert_eq!(compressor.compressed_bytes(), num_chunks);
assert_eq!(compressor.table.dim(), dim);
assert_eq!(compressor.table.ncenters(), num_centers);
assert_eq!(compressor.table.nchunks(), num_chunks);
assert!(&storage_provider.exists(pivot_file_name_compressor));
let compressor_pivots = read_bin::<u8>(
&mut storage_provider
.open_reader(pivot_file_name_compressor)
.unwrap(),
)
.unwrap();
let true_pivots =
read_bin::<u8>(&mut storage_provider.open_reader(pivot_file_name).unwrap()).unwrap();
assert_eq!(compressor_pivots, true_pivots);
}
#[rstest]
fn throw_error_for_resume_and_no_existing_file() {
let storage_provider = VirtualStorageProvider::new_memory();
storage_provider
.filesystem()
.create_dir("/pq_generation_tests")
.expect("Could not create test directory");
let pivot_file_name = "/pq_generation_tests/generate_pq_pivots_test.bin";
let compressed_file_name = "/pq_generation_tests/compressed_not_used.bin";
let data_path = "/pq_generation_tests/data_path.bin";
let (ndata, dim, num_centers, num_chunks, max_k_means_reps) = (5, 8, 2, 2, 5);
write_bin(
MatrixView::try_from(VALIDATION_DATA.as_slice(), ndata, dim).unwrap(),
&mut storage_provider.create_for_write(data_path).unwrap(),
)
.unwrap();
let pool = create_thread_pool_for_test();
let compressor = create_new_compressor(
CompressionStage::Resume,
&storage_provider,
dim,
num_chunks,
max_k_means_reps,
num_centers,
1.0,
pool.as_ref(),
pivot_file_name.to_string(),
compressed_file_name.to_string(),
Some(data_path),
);
assert!(compressor.is_err());
}
#[rstest]
fn test_pq_end_to_end_with_codebook() {
let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
let pool = create_thread_pool_for_test();
let dim = 128;
let num_chunks = 1;
let max_k_means_reps = 10;
let compressor = create_new_compressor(
CompressionStage::Resume,
&storage_provider,
dim,
num_chunks,
max_k_means_reps,
256,
1.0,
pool.as_ref(),
TEST_PQ_PIVOTS_PATH.to_string(),
"".to_string(),
None,
);
if let Err(x) = compressor.as_ref() {
println!("Error creating compressor: {x}");
};
assert!(compressor.is_ok());
let data_matrix =
read_bin::<f32>(&mut storage_provider.open_reader(TEST_PQ_DATA_PATH).unwrap()).unwrap();
let npts = data_matrix.nrows();
let mut compressed_mat = vec![0_u8; num_chunks * npts];
let result = compressor.unwrap().compress(
data_matrix.as_view(),
MutMatrixView::try_from(&mut compressed_mat, npts, num_chunks).unwrap(),
);
assert!(result.is_ok());
let compressed_gt = read_bin::<u8>(
&mut storage_provider
.open_reader(TEST_PQ_COMPRESSED_PATH)
.unwrap(),
)
.unwrap();
assert_eq!(compressed_gt.as_slice(), &compressed_mat);
}
#[rstest]
#[case(129, 128, 256)] #[case(128, 0, 256)] #[case(128, 128, 0)] fn test_parameter_error_cases(
#[case] dim: usize,
#[case] num_chunks: usize,
#[case] centers: usize,
) {
let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
let pool = create_thread_pool_for_test();
let max_k_means_reps = 10;
let compressor = create_new_compressor(
CompressionStage::Start,
&storage_provider,
dim,
num_chunks,
max_k_means_reps,
centers,
1.0,
pool.as_ref(),
TEST_PQ_PIVOTS_PATH.to_string(),
"".to_string(),
None,
);
assert!(compressor.is_err());
}
}