use std::ptr;
use crate::data_model::GraphDataType;
use byteorder::{ByteOrder, LittleEndian};
use diskann::{ANNError, ANNResult};
use hashbrown::HashMap;
use crate::{
data_model::GraphHeader,
search::{provider::disk_sector_graph::DiskSectorGraph, traits::VertexProvider},
utils::aligned_file_reader::traits::AlignedFileReader,
};
struct Offsets {
vec_idx: usize,
idx: usize,
}
pub struct DiskVertexProvider<Data, AlignedReaderType>
where
Data: GraphDataType<VectorIdType = u32>,
AlignedReaderType: AlignedFileReader,
{
pub centroid_vertex_id: u64,
dim: usize,
fp_vector_len: u64,
sector_graph: DiskSectorGraph<AlignedReaderType>,
vector_buf: Vec<Data::VectorDataType>,
cached_adjacency_list: Vec<Vec<Data::VectorIdType>>,
cached_associated_data: Vec<Data::AssociatedDataType>,
loaded_nodes: HashMap<Data::VectorIdType, Offsets>,
associated_data_size: usize,
node_len: u64,
io_operations: u32,
max_batch_size: usize,
}
impl<Data, AlignedReaderType> VertexProvider<Data> for DiskVertexProvider<Data, AlignedReaderType>
where
Data: GraphDataType<VectorIdType = u32>,
AlignedReaderType: AlignedFileReader,
{
fn get_vector(
&self,
vertex_id: &Data::VectorIdType,
) -> ANNResult<&[<Data as GraphDataType>::VectorDataType]> {
match self.loaded_nodes.get(vertex_id) {
Some(local_offset) => Ok(&self.vector_buf
[local_offset.idx * self.dim..(local_offset.idx * self.dim) + self.dim]),
None => Err(ANNError::log_get_vertex_data_error(
vertex_id.to_string(),
"Vector".to_string(),
)),
}
}
fn get_adjacency_list(
&self,
vertex_id: &Data::VectorIdType,
) -> ANNResult<&[Data::VectorIdType]> {
match self.loaded_nodes.get(vertex_id) {
Some(local_offset) => Ok(&self.cached_adjacency_list[local_offset.vec_idx]),
None => Err(ANNError::log_get_vertex_data_error(
vertex_id.to_string(),
"AdjacencyList".to_string(),
)),
}
}
fn get_associated_data(
&self,
vertex_id: &Data::VectorIdType,
) -> ANNResult<&Data::AssociatedDataType> {
match self.loaded_nodes.get(vertex_id) {
Some(local_offset) => Ok(&self.cached_associated_data[local_offset.vec_idx]),
None => Err(ANNError::log_get_vertex_data_error(
vertex_id.to_string(),
"AssociatedData".to_string(),
)),
}
}
fn load_vertices(&mut self, vertex_ids: &[Data::VectorIdType]) -> ANNResult<()> {
self.clear_before_next_read();
self.fetch_nodes(vertex_ids)?;
self.io_operations += vertex_ids.len() as u32;
Ok(())
}
fn process_loaded_node(
&mut self,
vertex_id: &Data::VectorIdType,
idx: usize,
) -> Result<(), ANNError> {
let fp_vector_buf =
&self.sector_graph.node_disk_buf(idx, *vertex_id)[..self.fp_vector_len as usize];
unsafe {
ptr::copy_nonoverlapping(
fp_vector_buf.as_ptr(),
self.vector_buf[idx * self.dim..(idx * self.dim) + self.dim].as_mut_ptr()
as *mut u8,
fp_vector_buf.len(),
);
}
let neighbor_and_data_buf =
&self.sector_graph.node_disk_buf(idx, *vertex_id)[self.fp_vector_len as usize..];
let num_neighbors = LittleEndian::read_u32(&neighbor_and_data_buf[0..4]) as usize;
match neighbor_and_data_buf.get(4..4 + num_neighbors * 4) {
Some(buf) => {
let mut adjacency_list = vec![Default::default(); num_neighbors];
bytemuck::must_cast_slice_mut::<_, u8>(&mut adjacency_list).copy_from_slice(buf);
self.cached_adjacency_list.push(adjacency_list);
}
None => {
return Err(ANNError::message(
diskann::ANNErrorKind::SerdeError,
format!(
"malformed length for vertex {} \
- reported neighbors is {} ({} bytes) which exceeds the buffer length {}",
vertex_id,
num_neighbors,
4 * num_neighbors + 4,
neighbor_and_data_buf.len()
),
));
}
}
let data_end: usize = (self.node_len - self.fp_vector_len) as usize;
let associated_data = bincode::deserialize(
&neighbor_and_data_buf[data_end - self.associated_data_size..data_end],
)
.map_err(|err| {
ANNError::log_serde_error(
"Error deserializing associated data from bytes".to_string(),
*err,
)
})?;
self.cached_associated_data.push(associated_data);
let vec_idx = self.loaded_nodes.len();
self.loaded_nodes
.insert(*vertex_id, Offsets { vec_idx, idx });
Ok(())
}
fn io_operations(&self) -> u32 {
self.io_operations
}
fn clear(&mut self) {
self.clear_before_next_read();
self.io_operations = 0;
}
fn vertices_loaded_count(&self) -> u32 {
self.io_operations
}
}
impl<Data, AlignedReaderType> DiskVertexProvider<Data, AlignedReaderType>
where
Data: GraphDataType<VectorIdType = u32>,
AlignedReaderType: AlignedFileReader,
{
pub fn new(
header: &GraphHeader,
max_batch_size: usize,
sector_reader: AlignedReaderType,
) -> ANNResult<Self> {
let metadata = header.metadata();
let dim = metadata.dims;
Ok(Self {
centroid_vertex_id: metadata.medoid,
dim,
fp_vector_len: (dim * std::mem::size_of::<Data::VectorDataType>()) as u64,
sector_graph: DiskSectorGraph::new(sector_reader, header, max_batch_size)?,
vector_buf: vec![Data::VectorDataType::default(); max_batch_size * dim],
cached_adjacency_list: Vec::with_capacity(max_batch_size),
cached_associated_data: Vec::with_capacity(max_batch_size),
loaded_nodes: HashMap::with_capacity(max_batch_size),
associated_data_size: metadata.associated_data_length,
node_len: metadata.node_len,
io_operations: 0,
max_batch_size,
})
}
fn reconfigure(&mut self, max_batch_size: usize) -> ANNResult<()> {
if max_batch_size > self.max_batch_size {
self.clear_before_next_read();
self.sector_graph.reconfigure(max_batch_size)?;
self.cached_adjacency_list.reserve(max_batch_size);
self.cached_associated_data.reserve(max_batch_size);
self.loaded_nodes.reserve(max_batch_size);
self.vector_buf = vec![Data::VectorDataType::default(); max_batch_size * self.dim];
self.max_batch_size = max_batch_size;
}
Ok(())
}
fn fetch_nodes(&mut self, nodes_to_fetch: &[Data::VectorIdType]) -> ANNResult<()> {
self.reconfigure(nodes_to_fetch.len())?;
let sectors_to_fetch: Vec<u64> = nodes_to_fetch
.iter()
.map(|&vertex_id| self.sector_graph.node_sector_index(vertex_id))
.collect();
self.sector_graph.read_graph(§ors_to_fetch)?;
Ok(())
}
fn clear_before_next_read(&mut self) {
self.sector_graph.reset();
self.loaded_nodes.clear();
self.cached_adjacency_list.clear();
self.cached_associated_data.clear();
}
}
#[cfg(test)]
mod disk_vertex_provider_tests {
use std::sync::Arc;
use crate::{data_model::GraphDataType, test_utils::GraphDataF32VectorU32Data};
use diskann::{graph::config, utils::ONE};
use diskann_providers::storage::{
StorageReadProvider, StorageWriteProvider, VirtualStorageProvider,
};
use diskann_providers::{
model::IndexConfiguration, storage::get_disk_index_file, utils::load_metadata_from_file,
};
use diskann_utils::test_data_root;
use vfs::OverlayFS;
use crate::{
build::builder::build::DiskIndexBuilder,
data_model::{CachingStrategy, GraphHeader},
disk_index_build_parameter::{
DiskIndexBuildParameters, MemoryBudget, NumPQChunks, DISK_SECTOR_LEN,
},
search::{
provider::disk_vertex_provider_factory::DiskVertexProviderFactory,
traits::{VertexProvider, VertexProviderFactory},
},
storage::DiskIndexWriter,
utils::VirtualAlignedReaderFactory,
QuantizationType,
};
fn generate_disk_index_with_associated_data<StorageProviderType>(
storage_provider: &StorageProviderType,
index_path_prefix: &str,
) where
StorageProviderType: StorageReadProvider + StorageWriteProvider,
<StorageProviderType as StorageReadProvider>::Reader: std::marker::Send,
StorageProviderType: 'static,
{
let max_degree = 4;
let l_build = 50;
let data_path = "/disk_index_search/disk_index_siftsmall_learn_256pts_data.fbin";
let associated_data_path = "/sift/siftsmall_learn_256pts_u32_associated_data.fbin";
let metadata = load_metadata_from_file(storage_provider, data_path).unwrap();
let memory_budget = MemoryBudget::try_from_gb(1.0).unwrap();
let num_pq_chunks = NumPQChunks::new_with(128, metadata.ndims()).unwrap();
let disk_index_build_parameters =
DiskIndexBuildParameters::new(memory_budget, QuantizationType::FP, num_pq_chunks);
let config = config::Builder::new_with(
max_degree,
config::MaxDegree::default_slack(),
l_build,
diskann_vector::distance::Metric::L2.into(),
|b| {
b.saturate_after_prune(true);
},
)
.build()
.unwrap();
let config = IndexConfiguration::new(
diskann_vector::distance::Metric::L2,
metadata.ndims(),
metadata.npoints(),
ONE,
1,
config,
);
let disk_index_writer = DiskIndexWriter::new(
data_path.to_string(),
index_path_prefix.to_string(),
Some(associated_data_path.to_string()),
DISK_SECTOR_LEN,
)
.unwrap();
let mut disk_index: DiskIndexBuilder<GraphDataF32VectorU32Data, StorageProviderType> =
DiskIndexBuilder::<GraphDataF32VectorU32Data, StorageProviderType>::new(
storage_provider,
disk_index_build_parameters,
config,
disk_index_writer,
)
.unwrap();
let mem_index_file_path = format!("{}_mem.index.data", index_path_prefix);
let mem_index_associated_data_path =
format!("{}_mem.index.associated_data", index_path_prefix);
if storage_provider.exists(&mem_index_file_path) {
storage_provider
.delete(&mem_index_file_path)
.expect("Failed to delete mem index file");
}
if storage_provider.exists(&mem_index_associated_data_path) {
storage_provider
.delete(&mem_index_associated_data_path)
.expect("Failed to delete mem index associated data file");
}
disk_index.build().unwrap();
assert!(!storage_provider.exists(&mem_index_file_path));
assert!(!storage_provider.exists(&mem_index_associated_data_path));
storage_provider
.delete(&format!("{}_pq_pivots.bin", index_path_prefix))
.expect("Failed to delete file");
storage_provider
.delete(&format!("{}_pq_compressed.bin", index_path_prefix))
.expect("Failed to delete file");
}
#[test]
fn test_disk_index_with_associated_data() {
let storage_provider = Arc::new(VirtualStorageProvider::new_overlay(test_data_root()));
let index_path_prefix = "/disk_index_search/disk_index_sift_learn_R4_L50_A1.2_test_disk_index_with_associated_data";
generate_disk_index_with_associated_data(storage_provider.as_ref(), index_path_prefix);
{
let vertex_provider_factory = DiskVertexProviderFactory::new(
VirtualAlignedReaderFactory::new(
get_disk_index_file(index_path_prefix).to_string(),
storage_provider.clone(),
),
CachingStrategy::None,
)
.unwrap();
let (mut vertex_provider, header) =
create_disk_provider::<GraphDataF32VectorU32Data>(&vertex_provider_factory);
let nodes = (0..256).map(|i| i as u32).collect::<Vec<u32>>();
VertexProvider::load_vertices(&mut vertex_provider, nodes.as_slice()).unwrap();
for (idx, vertex_id) in nodes.iter().enumerate() {
VertexProvider::process_loaded_node(&mut vertex_provider, vertex_id, idx).unwrap();
}
for vertex_id in 0..header.metadata().num_pts {
let test_vertex_id = vertex_id as u32;
let associated_data =
VertexProvider::get_associated_data(&vertex_provider, &test_vertex_id).unwrap();
assert_eq!(
test_vertex_id,
{ *associated_data },
"vertex_id: {}, associated_data {}",
vertex_id,
{ *associated_data }
);
}
assert_eq!(vertex_provider.io_operations(), 256);
assert_eq!(vertex_provider.vertices_loaded_count(), 256);
}
storage_provider
.delete(&get_disk_index_file(index_path_prefix))
.expect("Failed to delete file");
}
fn create_disk_provider<Data: GraphDataType<VectorIdType = u32>>(
vertex_provider_factory: &DiskVertexProviderFactory<Data, VirtualAlignedReaderFactory<OverlayFS>>,
) -> (
<DiskVertexProviderFactory<Data, VirtualAlignedReaderFactory<OverlayFS>> as VertexProviderFactory<
Data,
>>::VertexProviderType,
GraphHeader,
){
let header = vertex_provider_factory.get_header().unwrap();
let vertex_provider = vertex_provider_factory
.create_vertex_provider(256, &header)
.unwrap();
(vertex_provider, header)
}
}