#![warn(missing_debug_implementations, missing_docs)]
use std::num::NonZeroUsize;
use diskann::ANNError;
use thiserror::Error;
use super::QuantizationType;
use crate::error::{diskann_error, ErrorKind};
pub const BYTES_IN_GB: f64 = 1024_f64 * 1024_f64 * 1024_f64;
pub const DISK_SECTOR_LEN: usize = 4096;
const DEFAULT_DATA_COMPRESSION_CHUNK_VECTOR_COUNT: usize = 25_000;
#[derive(Debug, Error, PartialEq)]
#[error("Budget must be greater than zero")]
pub struct InvalidMemBudget;
impl From<InvalidMemBudget> for ANNError {
fn from(value: InvalidMemBudget) -> Self {
diskann_error!(ErrorKind::IndexConfigError("MemoryBudget"), value)
}
}
#[derive(Clone, Copy, PartialEq, Debug)]
pub struct MemoryBudget {
bytes: NonZeroUsize,
}
impl MemoryBudget {
pub fn try_from_gb(gib: f64) -> Result<Self, InvalidMemBudget> {
let bytes_f = (gib * BYTES_IN_GB).round() as usize;
let bytes = NonZeroUsize::new(bytes_f).ok_or(InvalidMemBudget)?;
Ok(Self { bytes })
}
pub fn in_bytes(self) -> usize {
self.bytes.get()
}
}
#[derive(Debug, Error, PartialEq)]
pub enum PQChunksError {
#[error("Dimension must be greater than zero")]
DimensionIsZero,
#[error("Number of PQ chunks must be within [1, {dim}], received {num_chunks}")]
OutOfRange {
num_chunks: usize,
dim: usize,
},
}
impl From<PQChunksError> for ANNError {
fn from(value: PQChunksError) -> Self {
diskann_error!(ErrorKind::IndexConfigError("NumPQChunks"), value)
}
}
#[derive(Clone, Copy, PartialEq, Debug)]
pub struct NumPQChunks(NonZeroUsize);
impl NumPQChunks {
pub fn new_with(num_chunks: usize, dim: usize) -> Result<Self, PQChunksError> {
if dim == 0 {
return Err(PQChunksError::DimensionIsZero);
}
let num_chunks = NonZeroUsize::new(num_chunks).ok_or(PQChunksError::DimensionIsZero)?;
if num_chunks.get() > dim {
return Err(PQChunksError::OutOfRange {
dim,
num_chunks: num_chunks.get(),
});
}
Ok(Self(num_chunks))
}
pub fn get(self) -> usize {
self.0.into()
}
}
#[derive(Clone, Copy, PartialEq, Debug)]
pub struct DiskIndexBuildParameters {
build_memory_limit: MemoryBudget,
search_pq_chunks: NumPQChunks,
build_quantization: QuantizationType,
data_compression_chunk_vector_count: usize,
}
impl DiskIndexBuildParameters {
pub fn new(
build_memory_limit: MemoryBudget,
build_quantization: QuantizationType,
search_pq_chunks: NumPQChunks,
) -> Self {
Self {
build_memory_limit,
search_pq_chunks,
build_quantization,
data_compression_chunk_vector_count: DEFAULT_DATA_COMPRESSION_CHUNK_VECTOR_COUNT,
}
}
pub fn with_data_compression_chunk_vector_count(
mut self,
data_compression_chunk_vector_count: usize,
) -> Self {
self.data_compression_chunk_vector_count = data_compression_chunk_vector_count;
self
}
pub fn build_memory_limit(&self) -> MemoryBudget {
self.build_memory_limit
}
pub fn build_quantization(&self) -> &QuantizationType {
&self.build_quantization
}
pub fn search_pq_chunks(&self) -> NumPQChunks {
self.search_pq_chunks
}
pub fn data_compression_chunk_vector_count(&self) -> usize {
self.data_compression_chunk_vector_count
}
}
#[cfg(test)]
mod dataset_test {
use diskann::ANNError;
use crate::error::{error_kind, ErrorKind};
use super::*;
#[test]
fn memory_budget_converts_units() {
let budget = MemoryBudget::try_from_gb(2.0).unwrap();
assert_eq!(budget.in_bytes() as f64, 2.0 * BYTES_IN_GB);
assert!(MemoryBudget::try_from_gb(0.0).is_err());
}
#[test]
fn build_with_num_of_pq_chunks_should_work() {
let memory_budget = MemoryBudget::try_from_gb(2.0).unwrap();
let num_pq_chunks = NumPQChunks::new_with(20, 128).unwrap();
let result = DiskIndexBuildParameters::new(
memory_budget,
QuantizationType::default(),
num_pq_chunks,
);
assert_eq!(result.search_pq_chunks().get(), num_pq_chunks.get());
assert_eq!(
result.data_compression_chunk_vector_count(),
DEFAULT_DATA_COMPRESSION_CHUNK_VECTOR_COUNT
);
}
#[test]
fn data_compression_chunk_vector_count_can_be_configured() {
let memory_budget = MemoryBudget::try_from_gb(2.0).unwrap();
let num_pq_chunks = NumPQChunks::new_with(20, 128).unwrap();
let result = DiskIndexBuildParameters::new(
memory_budget,
QuantizationType::default(),
num_pq_chunks,
)
.with_data_compression_chunk_vector_count(10_000);
assert_eq!(result.data_compression_chunk_vector_count(), 10_000);
}
#[test]
fn disk_index_build_parameters_try_new_handles_invalid() {
let memory_budget = MemoryBudget::try_from_gb(1.0).unwrap();
let pq_chunks = NumPQChunks::new_with(1, 128).unwrap();
let params =
DiskIndexBuildParameters::new(memory_budget, QuantizationType::default(), pq_chunks);
assert_eq!(
params.build_memory_limit().in_bytes() as f64,
1.0 * BYTES_IN_GB
);
assert!(MemoryBudget::try_from_gb(0.0).is_err());
let err = MemoryBudget::try_from_gb(-1.0)
.map_err(ANNError::from)
.unwrap_err();
assert_eq!(
error_kind(&err),
ErrorKind::IndexConfigError("MemoryBudget")
);
}
#[test]
fn num_pq_chunks_new_rejects_invalid_values() {
assert!(NumPQChunks::new_with(0, 128).is_err());
assert!(NumPQChunks::new_with(129, 128).is_err());
assert!(NumPQChunks::new_with(1, 0).is_err());
}
#[test]
fn num_pq_chunks_new_accepts_valid_values() {
let chunks = NumPQChunks::new_with(64, 128).unwrap();
assert_eq!(chunks.get(), 64);
}
}