use std::ffi::CString;
use std::io::{Write, stderr};
use std::path::Path;
use super::{IndexParams, VamanaError};
use crate::dlpack::AsDlTensor;
use crate::error::check_cuvs;
use crate::resources::Resources;
type Result<T> = std::result::Result<T, VamanaError>;
#[derive(Debug)]
pub struct Index(ffi::cuvsVamanaIndex_t);
impl Index {
pub fn build<T>(res: &Resources, params: &IndexParams, dataset: &T) -> Result<Index>
where
T: AsDlTensor + ?Sized,
{
let dataset = dataset.as_dl_tensor()?;
let index = Index::create_handle()?;
unsafe {
check_cuvs(ffi::cuvsVamanaBuild(
res.handle(),
params.handle(),
dataset.to_c().as_mut_ptr(),
index.0,
))?;
}
Ok(index)
}
fn create_handle() -> Result<Index> {
unsafe {
let mut index = std::mem::MaybeUninit::<ffi::cuvsVamanaIndex_t>::uninit();
check_cuvs(ffi::cuvsVamanaIndexCreate(index.as_mut_ptr()))?;
Ok(Index(index.assume_init()))
}
}
pub fn serialize(
&self,
res: &Resources,
filename: impl AsRef<Path>,
include_dataset: bool,
) -> Result<()> {
let c_filename = CString::new(filename.as_ref().as_os_str().as_encoded_bytes())?;
check_cuvs(unsafe {
ffi::cuvsVamanaSerialize(res.handle(), c_filename.as_ptr(), self.0, include_dataset)
})?;
Ok(())
}
}
impl Drop for Index {
fn drop(&mut self) {
if let Err(e) = check_cuvs(unsafe { ffi::cuvsVamanaIndexDestroy(self.0) }) {
write!(stderr(), "failed to call cuvsVamanaIndexDestroy {:?}", e)
.expect("failed to write to stderr");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::DeviceTensor;
use ndarray_rand::RandomExt;
use ndarray_rand::rand_distr::Uniform;
#[test]
fn test_vamana() {
let build_params = IndexParams::builder().build().unwrap();
let res = Resources::new().unwrap();
let n_datapoints = 1024;
let n_features = 16;
let dataset = ndarray::Array::<f32, _>::random(
(n_datapoints, n_features),
Uniform::new(0., 1.0).unwrap(),
);
let dataset_device = DeviceTensor::from_host(&res, &dataset).unwrap();
let _index = Index::build(&res, &build_params, &dataset_device)
.expect("failed to create vamana index");
}
}