use rustc_hash::FxHashMap as HashMap;
use std::fmt::{Debug, Display};
use std::sync::Arc;
use instant_distance::Search;
pub mod all_pairs;
mod backend;
mod exact;
pub mod ivf;
pub mod metric;
#[cfg(test)]
pub(crate) mod tests;
use backend::{build_column_dict, Backend, EF_SEARCH};
pub(crate) use backend::{EXACT_THRESHOLD, KNN_SEED};
pub use metric::l2_simd;
pub struct ColumnDict<K> {
backend: Backend,
data_vec: Vec<VecPoint>,
name2index: HashMap<K, usize>,
names: Vec<K>,
}
impl<K> ColumnDict<K>
where
K: Clone + Eq + std::hash::Hash + Debug + Display + std::cmp::PartialEq,
{
pub fn names(&self) -> &Vec<K> {
&self.names
}
pub fn from_ndarray_views<'a>(data: Vec<ndarray::ArrayView1<'a, f32>>, names: Vec<K>) -> Self {
<ColumnDict<K> as ColumnDictOps<K, ndarray::ArrayView1<'a, f32>>>::from_column_views(
data, names,
)
}
pub fn from_ndarray(data: ndarray::Array2<f32>, names: Vec<K>) -> Self {
let views: Vec<_> = data.outer_iter().collect();
Self::from_ndarray_views(views, names)
}
pub fn from_dvector_views(data: Vec<nalgebra::DVectorView<f32>>, names: Vec<K>) -> Self {
<ColumnDict<K> as ColumnDictOps<K, nalgebra::DVectorView<f32>>>::from_column_views(
data, names,
)
}
pub fn from_dmatrix(data: nalgebra::DMatrix<f32>, names: Vec<K>) -> Self {
Self::from_dvector_views(data.column_iter().collect(), names)
}
pub fn empty_ndarray_views() -> Self {
<ColumnDict<K> as ColumnDictOps<K, ndarray::ArrayView1<f32>>>::empty()
}
pub fn empty_dvector_views() -> Self {
<ColumnDict<K> as ColumnDictOps<K, nalgebra::DVectorView<f32>>>::empty()
}
pub fn dim(&self) -> Option<usize> {
self.data_vec.first().map(|p| p.len())
}
pub fn num_points(&self) -> usize {
self.data_vec.len()
}
pub fn points(&self) -> impl Iterator<Item = &[f32]> {
self.data_vec.iter().map(|p| p.as_slice())
}
pub fn search_others(&self, query_name: &K, knn: usize) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
self.search_by_query_name(query_name, knn, true)
}
pub fn search_by_query_name(
&self,
query_name: &K,
knn: usize,
exclude_same: bool,
) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
let self_idx = self.index_of(query_name)?;
let query = &self.data_vec[self_idx];
let exclude = if exclude_same { Some(self_idx) } else { None };
let mut search = Search::default();
let (indices, distances) = self.search_indices(query, knn, exclude, &mut search);
Ok((self.names_of(&indices), distances))
}
pub fn search_others_reuse(
&self,
query_name: &K,
knn: usize,
scratch: &mut SearchScratch,
) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
let self_idx = self.index_of(query_name)?;
let query = &self.data_vec[self_idx];
let (indices, distances) = self.search_indices(query, knn, Some(self_idx), &mut scratch.0);
Ok((self.names_of(&indices), distances))
}
pub fn search_by_query_data(
&self,
query: &[f32],
knn: usize,
) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
if self.dim().unwrap_or(0) != query.len() {
return Err(anyhow::anyhow!("query's dim does not match"));
}
let query = VecPoint {
data: Arc::from(query),
};
let mut search = Search::default();
let (indices, distances) = self.search_indices(&query, knn, None, &mut search);
Ok((self.names_of(&indices), distances))
}
pub fn search_by_query_data_reuse(
&self,
query: &[f32],
knn: usize,
scratch: &mut SearchScratch,
) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
if self.dim().unwrap_or(0) != query.len() {
return Err(anyhow::anyhow!("query's dim does not match"));
}
let query = VecPoint {
data: Arc::from(query),
};
let (indices, distances) = self.search_indices(&query, knn, None, &mut scratch.0);
Ok((self.names_of(&indices), distances))
}
pub fn match_by_query_name_against(
&self,
query_name: &K,
knn: usize,
against: &Self,
) -> anyhow::Result<(Vec<K>, Vec<f32>)> {
let self_idx = self.index_of(query_name)?;
let query = &self.data_vec[self_idx];
let mut search = Search::default();
let (indices, distances) = against.search_indices(query, knn, None, &mut search);
Ok((against.names_of(&indices), distances))
}
fn search_indices(
&self,
query: &VecPoint,
knn: usize,
exclude: Option<usize>,
search: &mut Search,
) -> (Vec<usize>, Vec<f32>) {
let fetch = if exclude.is_some() { knn + 1 } else { knn };
let (mut indices, mut distances) = match &self.backend {
Backend::Exact => exact::topk(&self.data_vec, &query.data, fetch),
Backend::Approx(map) => {
debug_assert!(
fetch <= EF_SEARCH,
"knn={knn} exceeds EF_SEARCH={EF_SEARCH}; approx results are truncated"
);
let mut indices = Vec::with_capacity(fetch);
let mut distances = Vec::with_capacity(fetch);
for item in map.search(query, search).take(fetch) {
indices.push(*item.value as usize);
distances.push(item.distance.sqrt());
}
(indices, distances)
}
};
if let Some(e) = exclude {
if let Some(pos) = indices.iter().position(|&i| i == e) {
indices.remove(pos);
distances.remove(pos);
}
indices.truncate(knn);
distances.truncate(knn);
}
(indices, distances)
}
#[inline]
fn names_of(&self, indices: &[usize]) -> Vec<K> {
indices.iter().map(|&i| self.names[i].clone()).collect()
}
#[inline]
fn index_of(&self, name: &K) -> anyhow::Result<usize> {
self.name2index
.get(name)
.copied()
.ok_or_else(|| anyhow::anyhow!("name {} not found", name))
}
}
#[derive(Default)]
pub struct SearchScratch(Search);
pub trait ColumnDictOps<K, V> {
fn empty() -> Self;
fn from_column_views(data: Vec<V>, names: Vec<K>) -> Self;
}
impl<T, V> ColumnDictOps<T, V> for ColumnDict<T>
where
T: Clone + Eq + std::hash::Hash + Debug + Display,
V: Sync + MakeVecPoint,
{
fn empty() -> Self {
Self {
backend: Backend::Exact,
data_vec: vec![],
name2index: Default::default(),
names: vec![],
}
}
fn from_column_views(data: Vec<V>, names: Vec<T>) -> Self {
build_column_dict(data, names, EXACT_THRESHOLD)
}
}
#[derive(Clone, Debug)]
pub struct VecPoint {
data: Arc<[f32]>,
}
impl VecPoint {
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
pub fn as_slice(&self) -> &[f32] {
&self.data
}
}
pub trait MakeVecPoint {
fn to_vp(&self) -> VecPoint;
}
impl MakeVecPoint for Vec<f32> {
fn to_vp(&self) -> VecPoint {
VecPoint {
data: self.iter().copied().collect(),
}
}
}
impl MakeVecPoint for nalgebra::DVectorView<'_, f32> {
fn to_vp(&self) -> VecPoint {
VecPoint {
data: self.iter().copied().collect(),
}
}
}
impl MakeVecPoint for ndarray::ArrayView1<'_, f32> {
fn to_vp(&self) -> VecPoint {
VecPoint {
data: self.iter().copied().collect(),
}
}
}