pub trait VectorDistance {
fn dot_product(&self, other: &Self) -> f64;
fn cosine_similarity(&self, other: &Self, normalized: bool) -> f64;
fn angular_distance(&self, other: &Self, normalized: bool) -> f64;
fn euclidean_distance(&self, other: &Self) -> f64;
fn manhattan_distance(&self, other: &Self) -> f64;
fn chebyshev_distance(&self, other: &Self) -> f64;
}
const CHUNK: usize = 256;
macro_rules! impl_vector_distance {
($pair_sum:ident, $unary_sum:ident, $pair_max:ident) => {
fn dot_product(&self, other: &Self) -> f64 {
$pair_sum(&self.vec, &other.vec, |x, y| x * y)
}
fn cosine_similarity(&self, other: &Self, normalized: bool) -> f64 {
let dot_product = self.dot_product(other);
if normalized {
dot_product
} else {
let magnitude1: f64 = $unary_sum(&self.vec, |x| x.powi(2)).sqrt();
let magnitude2: f64 = $unary_sum(&other.vec, |x| x.powi(2)).sqrt();
dot_product / (magnitude1 * magnitude2)
}
}
fn angular_distance(&self, other: &Self, normalized: bool) -> f64 {
let cosine_sim = self.cosine_similarity(other, normalized);
cosine_sim.clamp(-1.0, 1.0).acos() / std::f64::consts::PI
}
fn euclidean_distance(&self, other: &Self) -> f64 {
$pair_sum(&self.vec, &other.vec, |x, y| (x - y).powi(2)).sqrt()
}
fn manhattan_distance(&self, other: &Self) -> f64 {
$pair_sum(&self.vec, &other.vec, |x, y| (x - y).abs())
}
fn chebyshev_distance(&self, other: &Self) -> f64 {
$pair_max(&self.vec, &other.vec, |x, y| (x - y).abs())
}
};
}
mod sequential {
use super::{CHUNK, VectorDistance};
use crate::embeddings::Embedding;
fn pair_sum(a: &[f64], b: &[f64], term: impl Fn(f64, f64) -> f64) -> f64 {
a.chunks(CHUNK)
.zip(b.chunks(CHUNK))
.map(|(a, b)| a.iter().zip(b).map(|(x, y)| term(*x, *y)).sum::<f64>())
.sum()
}
fn unary_sum(a: &[f64], term: impl Fn(f64) -> f64) -> f64 {
a.chunks(CHUNK)
.map(|a| a.iter().map(|x| term(*x)).sum::<f64>())
.sum()
}
fn pair_max(a: &[f64], b: &[f64], term: impl Fn(f64, f64) -> f64) -> f64 {
a.iter()
.zip(b)
.map(|(x, y)| term(*x, *y))
.fold(0.0, f64::max)
}
impl VectorDistance for Embedding {
impl_vector_distance!(pair_sum, unary_sum, pair_max);
}
}
#[cfg(test)]
mod tests;