Skip to main content

rig_core/embeddings/
distance.rs

1//! Distance and similarity helpers for embedding vectors.
2//!
3//! [`Embedding`](crate::embeddings::Embedding) reductions use fixed chunks and
4//! left-to-right summation for reproducible ordering.
5//!
6//! ```
7//! use rig_core::embeddings::{Embedding, distance::VectorDistance};
8//!
9//! let vector = Embedding { document: String::new(), vec: vec![1.0, 0.0] };
10//! assert_eq!(vector.dot_product(&vector), 1.0);
11//! ```
12
13/// Distance and similarity metrics for embedding vectors.
14/// Supply equal-length vectors; the embedding implementation pairs only their
15/// shared prefix. Unnormalized cosine requires nonzero magnitudes.
16pub trait VectorDistance {
17    /// Get dot product of two embedding vectors
18    fn dot_product(&self, other: &Self) -> f64;
19
20    /// Get cosine similarity of two embedding vectors.
21    /// If `normalized` is true, the dot product is returned.
22    fn cosine_similarity(&self, other: &Self, normalized: bool) -> f64;
23
24    /// Get angular distance of two embedding vectors.
25    fn angular_distance(&self, other: &Self, normalized: bool) -> f64;
26
27    /// Get euclidean distance of two embedding vectors.
28    fn euclidean_distance(&self, other: &Self) -> f64;
29
30    /// Get manhattan distance of two embedding vectors.
31    fn manhattan_distance(&self, other: &Self) -> f64;
32
33    /// Get chebyshev distance of two embedding vectors.
34    fn chebyshev_distance(&self, other: &Self) -> f64;
35}
36
37/// Reduction chunk size. Preserve left-to-right summation within and across
38/// chunks to keep floating-point results reproducible.
39const CHUNK: usize = 256;
40
41/// Generates the [`VectorDistance`] method bodies for [`Embedding`](crate::embeddings::Embedding)
42/// from one pairwise sum, one unary sum and one max-reduction.
43macro_rules! impl_vector_distance {
44    ($pair_sum:ident, $unary_sum:ident, $pair_max:ident) => {
45        fn dot_product(&self, other: &Self) -> f64 {
46            $pair_sum(&self.vec, &other.vec, |x, y| x * y)
47        }
48
49        fn cosine_similarity(&self, other: &Self, normalized: bool) -> f64 {
50            let dot_product = self.dot_product(other);
51
52            if normalized {
53                dot_product
54            } else {
55                let magnitude1: f64 = $unary_sum(&self.vec, |x| x.powi(2)).sqrt();
56                let magnitude2: f64 = $unary_sum(&other.vec, |x| x.powi(2)).sqrt();
57
58                dot_product / (magnitude1 * magnitude2)
59            }
60        }
61
62        fn angular_distance(&self, other: &Self, normalized: bool) -> f64 {
63            let cosine_sim = self.cosine_similarity(other, normalized);
64            // Roundoff can push a valid cosine beyond the domain of acos.
65            cosine_sim.clamp(-1.0, 1.0).acos() / std::f64::consts::PI
66        }
67
68        fn euclidean_distance(&self, other: &Self) -> f64 {
69            $pair_sum(&self.vec, &other.vec, |x, y| (x - y).powi(2)).sqrt()
70        }
71
72        fn manhattan_distance(&self, other: &Self) -> f64 {
73            $pair_sum(&self.vec, &other.vec, |x, y| (x - y).abs())
74        }
75
76        fn chebyshev_distance(&self, other: &Self) -> f64 {
77            $pair_max(&self.vec, &other.vec, |x, y| (x - y).abs())
78        }
79    };
80}
81
82mod sequential {
83    use super::{CHUNK, VectorDistance};
84    use crate::embeddings::Embedding;
85
86    /// Fixed chunks, two left-to-right sums, one thread.
87    fn pair_sum(a: &[f64], b: &[f64], term: impl Fn(f64, f64) -> f64) -> f64 {
88        a.chunks(CHUNK)
89            .zip(b.chunks(CHUNK))
90            .map(|(a, b)| a.iter().zip(b).map(|(x, y)| term(*x, *y)).sum::<f64>())
91            .sum()
92    }
93
94    fn unary_sum(a: &[f64], term: impl Fn(f64) -> f64) -> f64 {
95        a.chunks(CHUNK)
96            .map(|a| a.iter().map(|x| term(*x)).sum::<f64>())
97            .sum()
98    }
99
100    fn pair_max(a: &[f64], b: &[f64], term: impl Fn(f64, f64) -> f64) -> f64 {
101        a.iter()
102            .zip(b)
103            .map(|(x, y)| term(*x, *y))
104            .fold(0.0, f64::max)
105    }
106
107    impl VectorDistance for Embedding {
108        impl_vector_distance!(pair_sum, unary_sum, pair_max);
109    }
110}
111
112#[cfg(test)]
113mod tests;