use super::Divergence;
const RRF_SCORE_EPSILON: f32 = 1e-6;
#[derive(Debug, Clone, PartialEq)]
pub struct FusionVector {
pub inputs: Vec<Vec<(u64, f32)>>,
pub k: u32,
pub expected: Vec<(u64, f32)>,
}
#[must_use]
pub fn rrf_reference_vectors() -> Vec<FusionVector> {
vec![
FusionVector {
inputs: vec![vec![(10, 0.9), (20, 0.8), (30, 0.7)]],
k: 60,
expected: vec![(10, 0.016_393_4), (20, 0.016_129_0), (30, 0.015_873_0)],
},
FusionVector {
inputs: vec![vec![(1, 0.9), (2, 0.8)], vec![(2, 0.7), (3, 0.6)]],
k: 60,
expected: vec![(2, 0.032_522_5), (1, 0.016_393_4), (3, 0.016_129_0)],
},
FusionVector {
inputs: vec![vec![(100, 0.5), (200, 0.4)]],
k: 10,
expected: vec![(100, 0.090_909_1), (200, 0.083_333_3)],
},
]
}
fn fused_matches(expected: &[(u64, f32)], actual: &[(u64, f32)]) -> bool {
expected.len() == actual.len()
&& expected
.iter()
.zip(actual)
.all(|(e, a)| e.0 == a.0 && (e.1 - a.1).abs() <= RRF_SCORE_EPSILON)
}
#[must_use]
pub fn check_rrf(
fuse_fn: impl Fn(Vec<Vec<(u64, f32)>>, u32) -> Vec<(u64, f32)>,
) -> Vec<Divergence> {
let mut divergences = Vec::new();
for vector in rrf_reference_vectors() {
let actual = fuse_fn(vector.inputs.clone(), vector.k);
if !fused_matches(&vector.expected, &actual) {
divergences.push(Divergence {
case: format!("rrf(k={}, inputs={:?})", vector.k, vector.inputs),
expected: format!("{:?}", vector.expected),
actual: format!("{actual:?}"),
});
}
}
divergences
}