Skip to main content

rskit_ai/
vector.rs

1//! Vector math helpers shared across AI crates.
2
3/// Compute the cosine similarity between two vectors.
4#[must_use]
5pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
6    if a.len() != b.len() {
7        return 0.0;
8    }
9    let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
10    let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
11    let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
12    if norm_a == 0.0 || norm_b == 0.0 {
13        return 0.0;
14    }
15    dot / (norm_a * norm_b)
16}
17
18/// Compute the Euclidean (L2) distance between two vectors.
19#[must_use]
20pub fn euclidean_distance(a: &[f32], b: &[f32]) -> f32 {
21    if a.len() != b.len() {
22        return f32::NAN;
23    }
24    a.iter()
25        .zip(b.iter())
26        .map(|(x, y)| (x - y) * (x - y))
27        .sum::<f32>()
28        .sqrt()
29}
30
31/// Compute the dot product of two vectors.
32#[must_use]
33pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
34    if a.len() != b.len() {
35        return 0.0;
36    }
37    a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
38}
39
40/// Compute the element-wise mean of a collection of vectors.
41#[must_use]
42pub fn mean_pooling(vectors: &[Vec<f32>]) -> Option<Vec<f32>> {
43    let first = vectors.first()?;
44    let dims = first.len();
45    if vectors.iter().any(|vector| vector.len() != dims) {
46        return None;
47    }
48    let count = vectors.len() as f32;
49    let mut result = vec![0.0_f32; dims];
50    for vector in vectors {
51        for (index, value) in vector.iter().enumerate() {
52            result[index] += value;
53        }
54    }
55    for value in &mut result {
56        *value /= count;
57    }
58    Some(result)
59}
60
61/// Compute the element-wise maximum of a collection of vectors.
62#[must_use]
63pub fn max_pooling(vectors: &[Vec<f32>]) -> Option<Vec<f32>> {
64    let first = vectors.first()?;
65    let dims = first.len();
66    if vectors.iter().any(|vector| vector.len() != dims) {
67        return None;
68    }
69    let mut result = vec![f32::NEG_INFINITY; dims];
70    for vector in vectors {
71        for (index, value) in vector.iter().enumerate() {
72            if *value > result[index] {
73                result[index] = *value;
74            }
75        }
76    }
77    Some(result)
78}
79
80/// Normalize a vector to unit length.
81#[must_use]
82pub fn normalize(vector: &[f32]) -> Option<Vec<f32>> {
83    let norm = vector.iter().map(|value| value * value).sum::<f32>().sqrt();
84    if norm == 0.0 {
85        return None;
86    }
87    Some(vector.iter().map(|value| value / norm).collect())
88}