pub(crate) fn cosine(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let mut dot = 0.0f32;
let mut na = 0.0f32;
let mut nb = 0.0f32;
for i in 0..a.len() {
dot += a[i] * b[i];
na += a[i] * a[i];
nb += b[i] * b[i];
}
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na.sqrt() * nb.sqrt())
}
pub(crate) fn mmr_select(rel: &[f32], emb: &[Option<&[f32]>], lambda: f32, k: usize) -> Vec<usize> {
let n = rel.len();
debug_assert_eq!(n, emb.len());
let k = k.min(n);
if k == 0 {
return Vec::new();
}
let lambda = lambda.clamp(0.0, 1.0);
let mut lo = f32::INFINITY;
let mut hi = f32::NEG_INFINITY;
for &r in rel {
lo = lo.min(r);
hi = hi.max(r);
}
let span = hi - lo;
let reln = |i: usize| {
if span > 0.0 {
(rel[i] - lo) / span
} else {
1.0
}
};
let mut selected: Vec<usize> = Vec::with_capacity(k);
let mut chosen = vec![false; n];
while selected.len() < k {
let mut best: Option<usize> = None;
let mut best_score = f32::NEG_INFINITY;
for i in 0..n {
if chosen[i] {
continue;
}
let max_sim = match emb[i] {
Some(ei) => selected
.iter()
.filter_map(|&s| emb[s].map(|es| cosine(ei, es)))
.fold(0.0f32, f32::max),
None => 0.0,
};
let mmr = lambda * reln(i) - (1.0 - lambda) * max_sim;
if mmr > best_score {
best_score = mmr;
best = Some(i);
}
}
match best {
Some(i) => {
chosen[i] = true;
selected.push(i);
}
None => break,
}
}
selected
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lambda_one_is_pure_relevance_order() {
let rel = [0.2, 0.9, 0.5];
let a = [1.0f32, 0.0];
let emb = [Some(&a[..]), Some(&a[..]), Some(&a[..])];
let order = mmr_select(&rel, &emb, 1.0, 3);
assert_eq!(order, vec![1, 2, 0]);
}
#[test]
fn diversity_demotes_a_near_duplicate() {
let rel = [1.0, 0.9, 0.8];
let v0 = [1.0f32, 0.0, 0.0];
let v1 = [1.0f32, 0.0, 0.0]; let v2 = [0.0f32, 1.0, 0.0]; let emb = [Some(&v0[..]), Some(&v1[..]), Some(&v2[..])];
let order = mmr_select(&rel, &emb, 0.5, 3);
assert_eq!(order[0], 0, "most relevant is picked first");
assert_eq!(order[1], 2, "diverse candidate beats the near-duplicate");
assert_eq!(order[2], 1);
}
#[test]
fn missing_embeddings_fall_back_to_relevance() {
let rel = [0.3, 0.7];
let emb: [Option<&[f32]>; 2] = [None, None];
let order = mmr_select(&rel, &emb, 0.5, 2);
assert_eq!(order, vec![1, 0]);
}
#[test]
fn k_is_clamped_and_zero_is_empty() {
let rel = [0.5, 0.4];
let emb: [Option<&[f32]>; 2] = [None, None];
assert_eq!(mmr_select(&rel, &emb, 0.5, 0), Vec::<usize>::new());
assert_eq!(mmr_select(&rel, &emb, 0.5, 9).len(), 2);
}
}