fn cosine_similarity(a: &[f32], b: &[f32]) -> f64 {
let dot: f64 = a
.iter()
.zip(b)
.map(|(x, y)| (*x as f64) * (*y as f64))
.sum();
let na: f64 = a
.iter()
.map(|x| (*x as f64) * (*x as f64))
.sum::<f64>()
.sqrt();
let nb: f64 = b
.iter()
.map(|x| (*x as f64) * (*x as f64))
.sum::<f64>()
.sqrt();
if na < f64::EPSILON || nb < f64::EPSILON {
0.0
} else {
dot / (na * nb)
}
}
fn mmr_gain<V: AsRef<[f32]>>(
candidates: &[(String, f64, V)],
selected: &[usize],
idx: usize,
lambda: f32,
) -> f64 {
let relevance = candidates[idx].1;
let max_sim = selected
.iter()
.map(|&s| cosine_similarity(candidates[idx].2.as_ref(), candidates[s].2.as_ref()))
.fold(0.0f64, f64::max);
lambda as f64 * relevance - (1.0 - lambda) as f64 * max_sim
}
pub fn mmr<V: AsRef<[f32]>>(candidates: &[(String, f64, V)], lambda: f32, k: usize) -> Vec<String> {
if candidates.is_empty() || k == 0 {
return Vec::new();
}
let lambda = lambda.clamp(0.0, 1.0);
let k = k.min(candidates.len());
let mut unselected: Vec<usize> = (0..candidates.len()).collect();
let mut selected: Vec<usize> = Vec::with_capacity(k);
while selected.len() < k {
let next = unselected
.iter()
.copied()
.max_by(|&a, &b| {
mmr_gain(candidates, &selected, a, lambda)
.partial_cmp(&mmr_gain(candidates, &selected, b, lambda))
.unwrap_or(std::cmp::Ordering::Equal)
})
.unwrap_or_default();
unselected.retain(|&i| i != next);
selected.push(next);
}
selected
.into_iter()
.map(|i| candidates[i].0.clone())
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mmr_trades_diversity_against_relevance_exactly() {
let v_a = vec![1.0f32];
let v_b = vec![-1.0f32];
let v_c = v_a.clone();
let candidates: Vec<(String, f64, Vec<f32>)> = vec![
("A".into(), 0.9, v_a),
("B".into(), 0.5, v_b),
("C".into(), 0.8, v_c),
];
assert_eq!(mmr(&candidates, 0.5, 3), vec!["A", "B", "C"]);
assert_eq!(mmr(&candidates, 1.0, 3), vec!["A", "C", "B"]);
}
#[test]
fn mmr_truncates_to_k() {
let candidates = vec![
("A".to_string(), 0.9, vec![1.0f32]),
("B".to_string(), 0.5, vec![-1.0f32]),
];
assert_eq!(mmr(&candidates, 0.5, 1), vec!["A"]);
}
#[test]
fn mmr_handles_empty_and_zero() {
assert!(mmr::<Vec<f32>>(&[], 0.5, 3).is_empty());
let candidates = vec![("A".to_string(), 0.9, vec![1.0f32])];
assert!(mmr(&candidates, 0.5, 0).is_empty());
}
#[test]
fn mmr_first_pick_is_pure_relevance() {
let candidates = vec![
("A".to_string(), 0.9, vec![1.0f32]),
("C".to_string(), 0.8, vec![1.0f32]), ];
assert_eq!(mmr(&candidates, 0.5, 2), vec!["A", "C"]);
}
#[test]
fn cosine_zero_and_unit_vectors() {
assert_eq!(cosine_similarity(&[0.0, 0.0], &[0.5, 0.5]), 0.0);
assert!((cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]) - 1.0).abs() < 1e-9);
assert!((cosine_similarity(&[1.0, 0.0], &[-1.0, 0.0]) + 1.0).abs() < 1e-9);
}
}