use crate::{cosine, dot, hamming_distance, l1_distance, l2_distance, slot::jaccard_distance};
pub trait Distance<T> {
fn eval(&self, a: &[T], b: &[T]) -> f32;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistCosine;
impl Distance<f32> for DistCosine {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
1.0 - cosine(a, b)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistDot;
impl Distance<f32> for DistDot {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
-dot(a, b)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistL2;
impl Distance<f32> for DistL2 {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
l2_distance(a, b)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistL1;
impl Distance<f32> for DistL1 {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
l1_distance(a, b)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistHamming;
impl Distance<u8> for DistHamming {
#[inline]
fn eval(&self, a: &[u8], b: &[u8]) -> f32 {
hamming_distance(a, b) as f32
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DistSlotU32;
impl Distance<u32> for DistSlotU32 {
#[inline]
fn eval(&self, a: &[u32], b: &[u32]) -> f32 {
jaccard_distance(a, b)
}
}
#[cfg(feature = "anndists")]
mod anndists_impls {
use super::*;
impl anndists::dist::Distance<f32> for DistCosine {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
Distance::eval(self, a, b)
}
}
impl anndists::dist::Distance<f32> for DistDot {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
Distance::eval(self, a, b)
}
}
impl anndists::dist::Distance<f32> for DistL2 {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
Distance::eval(self, a, b)
}
}
impl anndists::dist::Distance<f32> for DistL1 {
#[inline]
fn eval(&self, a: &[f32], b: &[f32]) -> f32 {
Distance::eval(self, a, b)
}
}
impl anndists::dist::Distance<u8> for DistHamming {
#[inline]
fn eval(&self, a: &[u8], b: &[u8]) -> f32 {
Distance::eval(self, a, b)
}
}
impl anndists::dist::Distance<u32> for DistSlotU32 {
#[inline]
fn eval(&self, a: &[u32], b: &[u32]) -> f32 {
Distance::eval(self, a, b)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cosine_distance_zero_for_parallel() {
let d = DistCosine.eval(&[1.0, 2.0, 3.0], &[2.0, 4.0, 6.0]);
assert!(
d.abs() < 1e-6,
"parallel vectors should have cosine distance 0, got {d}"
);
}
#[test]
fn dot_distance_orders_by_inner_product() {
let near = DistDot.eval(&[1.0, 0.0], &[1.0, 0.0]);
let far = DistDot.eval(&[1.0, 0.0], &[0.0, 1.0]);
assert!(near < far);
}
#[test]
fn l2_matches_free_function() {
let a = [1.0f32, 2.0, 3.0];
let b = [4.0f32, 0.0, 3.0];
assert_eq!(DistL2.eval(&a, &b), l2_distance(&a, &b));
}
#[test]
fn l1_matches_free_function() {
let a = [1.0f32, 2.0, 3.0];
let b = [4.0f32, 0.0, 3.0];
assert_eq!(DistL1.eval(&a, &b), l1_distance(&a, &b));
}
#[test]
fn hamming_matches_free_function() {
let a = [0b1111_0000u8, 0xFF];
let b = [0b1010_1010u8, 0x00];
assert_eq!(DistHamming.eval(&a, &b), hamming_distance(&a, &b) as f32);
}
#[test]
fn slot_distance_is_normalized_differing_fraction() {
let a = [1u32, 2, 3, 4];
let b = [1u32, 0, 3, 9];
assert_eq!(DistSlotU32.eval(&a, &b), 0.5);
}
fn closest<T, D: Distance<T>>(metric: &D, q: &[T], corpus: &[Vec<T>]) -> usize {
corpus
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| metric.eval(q, a).total_cmp(&metric.eval(q, b)))
.map(|(i, _)| i)
.unwrap()
}
#[test]
fn generic_index_over_metric() {
let corpus = vec![vec![1.0f32, 0.0], vec![0.0, 1.0], vec![1.0, 1.0]];
assert_eq!(closest(&DistCosine, &[1.0, 0.05], &corpus), 0);
assert_eq!(closest(&DistL2, &[0.95, 0.95], &corpus), 2);
let sketches = vec![vec![1u32, 2, 3, 4], vec![1, 2, 3, 9], vec![9, 9, 9, 9]];
assert_eq!(closest(&DistSlotU32, &[1, 2, 3, 4], &sketches), 0);
}
}