use crate::data_model::GraphDataType;
use diskann::{ANNError, ANNResult};
use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
use diskann_providers::{
index::diskann_async::train_pq,
model::{
graph::provider::async_::{common::NoStore, inmem::WithBits},
FixedChunkPQTable, IndexConfiguration, MAX_PQ_TRAINING_SET_SIZE,
},
storage::{PQStorage, SQStorage},
utils::{create_thread_pool, gen_random_slice, BridgeErr, PQPathNames},
};
use diskann_quantization::{
algorithms::transforms::{TargetDim, TransformKind},
alloc::GlobalAllocator,
scalar::train::ScalarQuantizationParameters,
spherical::{PreScale, SphericalQuantizer, SupportedMetric},
};
use diskann_utils::views::MatrixView;
use tracing::info;
use crate::{
error::{diskann_error, ErrorKind},
QuantizationType, SphericalBits,
};
pub enum BuildQuantizer {
NoQuant(NoStore),
Scalar1Bit(WithBits<1>),
Spherical1Bit(SphericalQuantizer),
PQ(FixedChunkPQTable),
}
impl BuildQuantizer {
pub fn train<Data, StorageProvider>(
build_quantization_type: &QuantizationType,
index_path_prefix: &str,
index_configuration: &IndexConfiguration,
data_path: &str,
storage_provider: &StorageProvider,
) -> ANNResult<Self>
where
Data: GraphDataType<VectorIdType = u32>,
StorageProvider: StorageReadProvider + StorageWriteProvider,
{
let num_points = index_configuration.max_points;
let p_val = MAX_PQ_TRAINING_SET_SIZE / (num_points as f64);
match *build_quantization_type {
QuantizationType::FP => Ok(Self::NoQuant(NoStore)),
QuantizationType::PQ { num_chunks } => {
let table = {
let seed = index_configuration.random_seed;
let mut rnd =
diskann_providers::utils::create_rnd_provider_from_optional_seed(seed)
.create_rnd();
let (train_data, train_size, train_dim) =
gen_random_slice::<Data::VectorDataType, _>(
data_path,
p_val,
storage_provider,
&mut rnd,
)?;
train_pq(
MatrixView::try_from(&train_data, train_size, train_dim).bridge_err()?,
num_chunks,
&mut rnd,
create_thread_pool(index_configuration.num_threads)?.as_ref(),
)?
};
let pq_paths = PQPathNames::new(index_path_prefix);
let pq_build_storage =
PQStorage::new(&pq_paths.pivots, &pq_paths.compressed_data, None);
pq_build_storage.write_pivot_data(
table.get_pq_table(),
None,
table.get_chunk_offsets(),
table.get_num_centers(),
table.get_dim(),
storage_provider,
)?;
Ok(Self::PQ(table))
}
QuantizationType::SQ {
nbits,
standard_deviation,
} => {
if nbits != 1 {
return Err(diskann_error!(
ErrorKind::IndexConfigError("build_quantization_type"),
"SQ quantization is only supported for 1 bit",
));
}
let rng = diskann_providers::utils::create_rnd_provider_from_optional_seed(
index_configuration.random_seed,
);
let (train_data_vector, train_size, train_dim) =
gen_random_slice::<Data::VectorDataType, _>(
data_path,
p_val,
storage_provider,
&mut rng.create_rnd(),
)?;
let quantizer_params = if let Some(std_dev) = standard_deviation {
ScalarQuantizationParameters::new(std_dev)
} else {
ScalarQuantizationParameters::default()
};
let quantizer = quantizer_params.train(
MatrixView::try_from(&train_data_vector, train_size, train_dim).bridge_err()?,
);
info!("Now quantizer is trained and saving to file");
let sq_storage = SQStorage::new(index_path_prefix);
sq_storage.save_quantizer(&quantizer, storage_provider)?;
Ok(Self::Scalar1Bit(WithBits::<1>::new(quantizer)))
}
QuantizationType::Spherical(SphericalBits::One) => {
let metric: SupportedMetric =
index_configuration.dist_metric.try_into().bridge_err()?;
let rng = diskann_providers::utils::create_rnd_provider_from_optional_seed(
index_configuration.random_seed,
);
let mut rnd = rng.create_rnd();
let (train_data, train_size, train_dim) =
gen_random_slice::<Data::VectorDataType, _>(
data_path,
p_val,
storage_provider,
&mut rnd,
)?;
let train_data =
MatrixView::try_from(&train_data, train_size, train_dim).bridge_err()?;
let quantizer = SphericalQuantizer::train(
train_data,
TransformKind::DoubleHadamard {
target_dim: TargetDim::Natural,
},
metric,
PreScale::ReciprocalMeanNorm,
&mut rnd,
GlobalAllocator,
)
.map_err(|err| ANNError::new(err).context("Failed to train spherical quantizer"))?;
Ok(Self::Spherical1Bit(quantizer))
}
}
}
}