use super::{ColumnDict, MakeVecPoint, VecPoint};
use instant_distance::{Builder, HnswMap};
use rustc_hash::FxHashMap as HashMap;
use std::fmt::{Debug, Display};
pub(crate) const KNN_SEED: u64 = 20_240_517;
pub(super) const EF_CONSTRUCTION: usize = 200;
pub(super) const EF_SEARCH: usize = 128;
pub(crate) const EXACT_THRESHOLD: usize = 8_192;
pub(super) enum Backend {
Exact,
Approx(HnswMap<VecPoint, u32>),
}
pub(super) fn build_column_dict<T, V>(
data: Vec<V>,
names: Vec<T>,
exact_threshold: usize,
) -> ColumnDict<T>
where
T: Clone + Eq + std::hash::Hash + Debug + Display,
V: Sync + MakeVecPoint,
{
let nn = data.len();
debug_assert!(
nn == names.len(),
"Data and names must have the same length"
);
let data_vec: Vec<VecPoint> = data.iter().map(|v| v.to_vp()).collect();
let mut name2index: HashMap<T, usize> = Default::default();
names.iter().enumerate().for_each(|(j, x)| {
name2index.insert(x.clone(), j);
});
let backend = if nn <= exact_threshold {
Backend::Exact
} else {
let values: Vec<u32> = (0..nn as u32).collect();
let map = Builder::default()
.seed(KNN_SEED)
.ef_construction(EF_CONSTRUCTION)
.ef_search(EF_SEARCH)
.build(data_vec.clone(), values);
Backend::Approx(map)
};
let ret = ColumnDict {
backend,
data_vec,
name2index,
names,
};
#[cfg(debug_assertions)]
{
for (i, x) in ret.names.iter().enumerate() {
if let Some(&j) = ret.name2index.get(x) {
debug_assert_eq!(i, j);
}
}
}
ret
}