use embeddenator_vsa::SparseVec;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SimilarityMetric {
#[default]
Cosine,
Hamming,
Jaccard,
DotProduct,
}
pub fn compute_similarity(a: &SparseVec, b: &SparseVec, metric: SimilarityMetric) -> f64 {
match metric {
SimilarityMetric::Cosine => a.cosine(b),
SimilarityMetric::Hamming => hamming_distance(a, b),
SimilarityMetric::Jaccard => jaccard_similarity(a, b),
SimilarityMetric::DotProduct => dot_product(a, b) as f64,
}
}
pub fn hamming_distance(a: &SparseVec, b: &SparseVec) -> f64 {
use std::collections::HashSet;
let mut differing: HashSet<usize> = HashSet::new();
for &dim in &a.pos {
if !b.pos.contains(&dim) {
differing.insert(dim);
}
}
for &dim in &a.neg {
if !b.neg.contains(&dim) {
differing.insert(dim);
}
}
for &dim in &b.pos {
if !a.pos.contains(&dim) {
differing.insert(dim);
}
}
for &dim in &b.neg {
if !a.neg.contains(&dim) {
differing.insert(dim);
}
}
differing.len() as f64
}
pub fn jaccard_similarity(a: &SparseVec, b: &SparseVec) -> f64 {
let mut intersection = 0usize;
for &dim in &a.pos {
if b.pos.contains(&dim) {
intersection += 1;
}
}
for &dim in &a.neg {
if b.neg.contains(&dim) {
intersection += 1;
}
}
let union = a.pos.len() + a.neg.len() + b.pos.len() + b.neg.len() - intersection;
if union == 0 {
return 1.0; }
intersection as f64 / union as f64
}
pub fn dot_product(a: &SparseVec, b: &SparseVec) -> i32 {
let mut score = 0i32;
for &dim in &a.pos {
if b.pos.contains(&dim) {
score += 1;
} else if b.neg.contains(&dim) {
score -= 1;
}
}
for &dim in &a.neg {
if b.neg.contains(&dim) {
score += 1;
} else if b.pos.contains(&dim) {
score -= 1;
}
}
score
}
#[cfg(test)]
mod tests {
use super::*;
use embeddenator_vsa::ReversibleVSAConfig;
fn make_vec(data: &[u8]) -> SparseVec {
let config = ReversibleVSAConfig::default();
SparseVec::encode_data(data, &config, None)
}
#[test]
fn test_cosine_identical() {
let a = make_vec(b"test");
let b = make_vec(b"test");
let sim = compute_similarity(&a, &b, SimilarityMetric::Cosine);
assert!(
sim > 0.99,
"Identical vectors should have ~1.0 cosine similarity"
);
}
#[test]
fn test_cosine_different() {
let a = make_vec(b"hello");
let b = make_vec(b"world");
let sim = compute_similarity(&a, &b, SimilarityMetric::Cosine);
assert!(sim < 0.5, "Different vectors should have low similarity");
}
#[test]
fn test_hamming_identical() {
let a = make_vec(b"test");
let b = make_vec(b"test");
let dist = hamming_distance(&a, &b);
assert_eq!(
dist, 0.0,
"Identical vectors should have 0 Hamming distance"
);
}
#[test]
fn test_hamming_different() {
let a = make_vec(b"hello");
let b = make_vec(b"world");
let dist = hamming_distance(&a, &b);
assert!(
dist > 0.0,
"Different vectors should have positive Hamming distance"
);
}
#[test]
fn test_jaccard_identical() {
let a = make_vec(b"test");
let b = make_vec(b"test");
let sim = jaccard_similarity(&a, &b);
assert!(
(sim - 1.0).abs() < 0.01,
"Identical vectors should have ~1.0 Jaccard similarity"
);
}
#[test]
fn test_jaccard_disjoint() {
let a = make_vec(b"aaa");
let b = make_vec(b"zzz");
let sim = jaccard_similarity(&a, &b);
assert!(
sim < 0.5,
"Different vectors should have low Jaccard similarity"
);
}
#[test]
fn test_dot_product_identical() {
let a = make_vec(b"test");
let b = make_vec(b"test");
let dot = dot_product(&a, &b);
assert!(
dot > 0,
"Identical vectors should have positive dot product"
);
}
}