use std::io::{Write, stderr};
use super::{IndexParams, IvfFlatError, SearchParams};
use crate::dlpack::{AsDlTensor, AsDlTensorMut, DLTensorView, DLTensorViewMut};
use crate::error::check_cuvs;
use crate::neighbors::filters::{Bitset, Filter, with_filter};
use crate::resources::Resources;
type Result<T> = std::result::Result<T, IvfFlatError>;
#[derive(Debug)]
pub struct Index(ffi::cuvsIvfFlatIndex_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::cuvsIvfFlatBuild(
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::cuvsIvfFlatIndex_t>::uninit();
check_cuvs(ffi::cuvsIvfFlatIndexCreate(index.as_mut_ptr()))?;
Ok(Index(index.assume_init()))
}
}
pub fn search<Q, N, D>(
&self,
res: &Resources,
params: &SearchParams,
queries: &Q,
neighbors: &mut N,
distances: &mut D,
) -> Result<()>
where
Q: AsDlTensor + ?Sized,
N: AsDlTensorMut + ?Sized,
D: AsDlTensorMut + ?Sized,
{
let queries = queries.as_dl_tensor()?;
let mut neighbors = neighbors.as_dl_tensor_mut()?;
let mut distances = distances.as_dl_tensor_mut()?;
self.search_impl(res, params, &queries, &mut neighbors, &mut distances, None)
}
pub fn search_filtered<Q, N, D>(
&self,
res: &Resources,
params: &SearchParams,
queries: &Q,
neighbors: &mut N,
distances: &mut D,
filter: &Filter<'_, Bitset>,
) -> Result<()>
where
Q: AsDlTensor + ?Sized,
N: AsDlTensorMut + ?Sized,
D: AsDlTensorMut + ?Sized,
{
let queries = queries.as_dl_tensor()?;
let mut neighbors = neighbors.as_dl_tensor_mut()?;
let mut distances = distances.as_dl_tensor_mut()?;
self.search_impl(res, params, &queries, &mut neighbors, &mut distances, Some(filter))
}
fn search_impl(
&self,
res: &Resources,
params: &SearchParams,
queries: &DLTensorView<'_>,
neighbors: &mut DLTensorViewMut<'_>,
distances: &mut DLTensorViewMut<'_>,
filter: Option<&Filter<'_, Bitset>>,
) -> Result<()> {
with_filter(filter, |prefilter| {
check_cuvs(unsafe {
ffi::cuvsIvfFlatSearch(
res.handle(),
params.handle(),
self.0,
queries.to_c().as_mut_ptr(),
neighbors.to_c().as_mut_ptr(),
distances.to_c().as_mut_ptr(),
prefilter,
)
})?;
Ok(())
})
}
}
impl Drop for Index {
fn drop(&mut self) {
if let Err(e) = check_cuvs(unsafe { ffi::cuvsIvfFlatIndexDestroy(self.0) }) {
write!(stderr(), "failed to call cuvsIvfFlatIndexDestroy {:?}", e)
.expect("failed to write to stderr");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_utils::DeviceTensor;
use ndarray::s;
use ndarray_rand::RandomExt;
use ndarray_rand::rand_distr::Uniform;
#[test]
fn test_ivf_flat() {
let build_params = IndexParams::builder().n_lists(64).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 ivf-flat index");
let n_queries = 4;
let queries = dataset.slice(s![0..n_queries, ..]).to_owned();
let k = 10;
let queries = DeviceTensor::from_host(&res, &queries).unwrap();
let mut neighbors_host = ndarray::Array::<i64, _>::zeros((n_queries, k));
let mut neighbors = DeviceTensor::<i64>::zeros(&res, &[n_queries, k]).unwrap();
let mut distances_host = ndarray::Array::<f32, _>::zeros((n_queries, k));
let mut distances = DeviceTensor::<f32>::zeros(&res, &[n_queries, k]).unwrap();
let search_params = SearchParams::builder().build().unwrap();
index.search(&res, &search_params, &queries, &mut neighbors, &mut distances).unwrap();
distances.copy_to_host(&res, &mut distances_host).unwrap();
neighbors.copy_to_host(&res, &mut neighbors_host).unwrap();
assert_eq!(neighbors_host[[0, 0]], 0);
assert_eq!(neighbors_host[[1, 0]], 1);
assert_eq!(neighbors_host[[2, 0]], 2);
assert_eq!(neighbors_host[[3, 0]], 3);
}
#[test]
fn test_ivf_flat_multiple_searches() {
let build_params = IndexParams::builder().n_lists(64).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 ivf-flat index");
let search_params = SearchParams::builder().build().unwrap();
let k = 5;
for _ in 0..3 {
let n_queries = 4;
let queries = dataset.slice(s![0..n_queries, ..]).to_owned();
let queries = DeviceTensor::from_host(&res, &queries).unwrap();
let mut neighbors_host = ndarray::Array::<i64, _>::zeros((n_queries, k));
let mut neighbors = DeviceTensor::<i64>::zeros(&res, &[n_queries, k]).unwrap();
let mut distances = DeviceTensor::<f32>::zeros(&res, &[n_queries, k]).unwrap();
index
.search(&res, &search_params, &queries, &mut neighbors, &mut distances)
.expect("search failed");
neighbors.copy_to_host(&res, &mut neighbors_host).unwrap();
assert_eq!(neighbors_host[[0, 0]], 0, "first query should find itself");
}
}
}