Skip to main content

recern_vector/
metric.rs

1use std::fmt;
2use std::str::FromStr;
3
4/// Distance function of a collection. For every metric, smaller is closer.
5#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
6pub enum Metric {
7    /// `1 - cos(a, b)`. Vectors are normalized when they are stored.
8    Cosine,
9    /// Squared Euclidean distance.
10    L2,
11    /// Negative inner product.
12    Dot,
13}
14
15impl Metric {
16    pub fn as_str(self) -> &'static str {
17        match self {
18            Metric::Cosine => "cosine",
19            Metric::L2 => "l2",
20            Metric::Dot => "dot",
21        }
22    }
23
24    #[inline]
25    pub fn distance(self, a: &[f32], b: &[f32]) -> f32 {
26        match self {
27            Metric::Cosine => 1.0 - dot(a, b),
28            Metric::L2 => l2_squared(a, b),
29            Metric::Dot => -dot(a, b),
30        }
31    }
32
33    pub(crate) fn code(self) -> u8 {
34        match self {
35            Metric::Cosine => 0,
36            Metric::L2 => 1,
37            Metric::Dot => 2,
38        }
39    }
40
41    pub(crate) fn from_code(code: u8) -> Option<Self> {
42        match code {
43            0 => Some(Metric::Cosine),
44            1 => Some(Metric::L2),
45            2 => Some(Metric::Dot),
46            _ => None,
47        }
48    }
49}
50
51impl fmt::Display for Metric {
52    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53        f.write_str(self.as_str())
54    }
55}
56
57impl FromStr for Metric {
58    type Err = String;
59
60    fn from_str(s: &str) -> Result<Self, Self::Err> {
61        match s.to_ascii_lowercase().as_str() {
62            "cosine" => Ok(Metric::Cosine),
63            "l2" | "euclidean" => Ok(Metric::L2),
64            "dot" | "ip" | "inner_product" => Ok(Metric::Dot),
65            other => Err(format!(
66                "unknown metric '{other}' (expected cosine, l2 or dot)"
67            )),
68        }
69    }
70}
71
72// Independent accumulators let the compiler vectorize the loops without
73// relaxing floating-point ordering.
74const LANES: usize = 8;
75
76#[inline]
77pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
78    debug_assert_eq!(a.len(), b.len());
79    let (chunks_a, chunks_b) = (a.chunks_exact(LANES), b.chunks_exact(LANES));
80    let (rest_a, rest_b) = (chunks_a.remainder(), chunks_b.remainder());
81    let mut acc = [0.0f32; LANES];
82    for (x, y) in chunks_a.zip(chunks_b) {
83        for i in 0..LANES {
84            acc[i] += x[i] * y[i];
85        }
86    }
87    let mut sum: f32 = acc.iter().sum();
88    for (x, y) in rest_a.iter().zip(rest_b) {
89        sum += x * y;
90    }
91    sum
92}
93
94#[inline]
95pub(crate) fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
96    debug_assert_eq!(a.len(), b.len());
97    let (chunks_a, chunks_b) = (a.chunks_exact(LANES), b.chunks_exact(LANES));
98    let (rest_a, rest_b) = (chunks_a.remainder(), chunks_b.remainder());
99    let mut acc = [0.0f32; LANES];
100    for (x, y) in chunks_a.zip(chunks_b) {
101        for i in 0..LANES {
102            let d = x[i] - y[i];
103            acc[i] += d * d;
104        }
105    }
106    let mut sum: f32 = acc.iter().sum();
107    for (x, y) in rest_a.iter().zip(rest_b) {
108        let d = x - y;
109        sum += d * d;
110    }
111    sum
112}
113
114/// Scales `v` to unit length. Returns `false` for a zero vector.
115pub(crate) fn normalize(v: &mut [f32]) -> bool {
116    let norm = dot(v, v).sqrt();
117    if norm == 0.0 || !norm.is_finite() {
118        return false;
119    }
120    for x in v {
121        *x /= norm;
122    }
123    true
124}
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129
130    #[test]
131    fn distances_match_naive_implementation() {
132        let a: Vec<f32> = (0..19).map(|i| i as f32 * 0.5 - 3.0).collect();
133        let b: Vec<f32> = (0..19).map(|i| (i as f32).sin()).collect();
134        let naive_dot: f32 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
135        let naive_l2: f32 = a.iter().zip(&b).map(|(x, y)| (x - y) * (x - y)).sum();
136        assert!((dot(&a, &b) - naive_dot).abs() < 1e-4);
137        assert!((l2_squared(&a, &b) - naive_l2).abs() < 1e-3);
138    }
139
140    #[test]
141    fn normalize_rejects_zero_vector() {
142        let mut v = [3.0, 4.0];
143        assert!(normalize(&mut v));
144        assert!((v[0] - 0.6).abs() < 1e-6 && (v[1] - 0.8).abs() < 1e-6);
145        assert!(!normalize(&mut [0.0, 0.0]));
146    }
147
148    #[test]
149    fn parses_metric_names() {
150        assert_eq!("Cosine".parse::<Metric>(), Ok(Metric::Cosine));
151        assert_eq!("euclidean".parse::<Metric>(), Ok(Metric::L2));
152        assert_eq!("ip".parse::<Metric>(), Ok(Metric::Dot));
153        assert!("manhattan".parse::<Metric>().is_err());
154    }
155}