1use std::fmt;
2use std::str::FromStr;
3
4#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
6pub enum Metric {
7 Cosine,
9 L2,
11 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
72const 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
114pub(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}