use std::collections::hash_map::Entry;
use ahash::AHashMap;
use itertools::Either;
use ordered_float::OrderedFloat;
use crate::segment::common::operation_error::{OperationError, OperationResult};
use crate::segment::types::{ExtendedPointId, ScoredPoint};
pub const DEFAULT_RRF_K: usize = 2;
fn position_score(position: usize, k: usize, weight: f32) -> f32 {
if weight <= 0.0 {
return 0.0;
}
1.0 / ((position + 1) as f32 / weight + k as f32 - 1.0)
}
pub fn rrf_scoring(
responses: Vec<Vec<ScoredPoint>>,
k: usize,
weights: Option<&[f32]>,
) -> OperationResult<Vec<ScoredPoint>> {
let mut points_by_id: AHashMap<ExtendedPointId, ScoredPoint> = AHashMap::new();
let weights = if let Some(weights) = weights {
if weights.len() != responses.len() {
return Err(OperationError::validation_error(format!(
"Number of weights in RRF should match number of pre-fetches: got {}, expected {}",
weights.len(),
responses.len()
)));
}
Either::Left(weights.iter().copied())
} else {
Either::Right(std::iter::repeat(1.0f32))
};
for (response, weight) in responses.into_iter().zip(weights) {
for (pos, mut point) in response.into_iter().enumerate() {
let rrf_score = position_score(pos, k, weight);
match points_by_id.entry(point.id) {
Entry::Occupied(mut entry) => {
entry.get_mut().score += rrf_score;
}
Entry::Vacant(entry) => {
point.score = rrf_score;
entry.insert(point);
}
}
}
}
let mut scores: Vec<_> = points_by_id.into_values().collect();
scores.sort_unstable_by(|a, b| {
OrderedFloat(b.score).cmp(&OrderedFloat(a.score))
});
Ok(scores)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::segment::types::ScoredPoint;
fn make_scored_point(id: u64, score: f32) -> ScoredPoint {
ScoredPoint {
id: id.into(),
version: 0,
score,
payload: None,
vector: None,
shard_key: None,
order_value: None,
}
}
#[test]
fn test_rrf_scoring_empty() {
let responses = vec![];
let scored_points = rrf_scoring(responses, DEFAULT_RRF_K, None).unwrap();
assert_eq!(scored_points.len(), 0);
}
#[test]
fn test_rrf_scoring_one() {
let responses = vec![vec![make_scored_point(1, 0.9)]];
let scored_points = rrf_scoring(responses, DEFAULT_RRF_K, None).unwrap();
assert_eq!(scored_points.len(), 1);
assert_eq!(scored_points[0].id, 1.into());
assert_eq!(scored_points[0].score, 0.5); }
#[test]
fn test_rrf_scoring() {
let responses = vec![
vec![make_scored_point(2, 0.9), make_scored_point(1, 0.8)],
vec![
make_scored_point(1, 0.7),
make_scored_point(2, 0.6),
make_scored_point(3, 0.5),
],
vec![
make_scored_point(5, 0.9),
make_scored_point(3, 0.5),
make_scored_point(1, 0.4),
],
];
let scored_points = rrf_scoring(responses, DEFAULT_RRF_K, None).unwrap();
assert_eq!(scored_points.len(), 4);
assert!(
scored_points
.array_windows()
.all(|[a, b]| a.score >= b.score),
);
assert_eq!(scored_points.len(), 4);
assert_eq!(scored_points[0].id, 1.into());
assert_eq!(scored_points[0].score, 1.0833334);
assert_eq!(scored_points[1].id, 2.into());
assert_eq!(scored_points[1].score, 0.8333334);
assert_eq!(scored_points[2].id, 3.into());
assert_eq!(scored_points[2].score, 0.5833334);
assert_eq!(scored_points[3].id, 5.into());
assert_eq!(scored_points[3].score, 0.5);
}
#[test]
fn test_rrf_scoring_weighted() {
let responses = vec![
vec![make_scored_point(1, 0.9), make_scored_point(2, 0.8)],
vec![make_scored_point(2, 0.9), make_scored_point(1, 0.8)],
];
let scored_points = rrf_scoring(responses.clone(), DEFAULT_RRF_K, None).unwrap();
assert_eq!(scored_points[0].score, scored_points[1].score);
let weights = [3.0, 1.0];
let scored_points = rrf_scoring(responses, DEFAULT_RRF_K, Some(&weights)).unwrap();
assert!(scored_points[0].id == 2.into());
assert!(scored_points[0].score > scored_points[1].score);
}
#[test]
fn test_rrf_scoring_weighted_ratio() {
let k = 60;
let responses = vec![
vec![
make_scored_point(11, 0.0),
make_scored_point(12, 0.0),
make_scored_point(13, 0.0),
make_scored_point(14, 0.0),
make_scored_point(15, 0.0),
make_scored_point(16, 0.0),
make_scored_point(17, 0.0),
make_scored_point(18, 0.0),
],
vec![
make_scored_point(21, 0.0),
make_scored_point(22, 0.0),
make_scored_point(23, 0.0),
make_scored_point(24, 0.0),
make_scored_point(25, 0.0),
make_scored_point(26, 0.0),
make_scored_point(27, 0.0),
make_scored_point(28, 0.0),
],
];
let weights = [3.0, 1.0];
let scored_points = rrf_scoring(responses, k, Some(&weights)).unwrap();
let top_10 = &scored_points[..10];
let count_source_1 = top_10
.iter()
.filter(|p| p.id.as_u64() >= 10 && p.id.as_u64() < 20)
.count();
let count_source_2 = top_10
.iter()
.filter(|p| p.id.as_u64() >= 20 && p.id.as_u64() < 30)
.count();
assert!(count_source_1 >= 2 * count_source_2); }
#[test]
fn test_rrf_scoring_weights_length_mismatch() {
let responses = vec![
vec![make_scored_point(1, 0.9)],
vec![make_scored_point(2, 0.9)],
];
let weights = [1.0, 2.0, 3.0];
let result = rrf_scoring(responses.clone(), DEFAULT_RRF_K, Some(&weights));
assert!(result.is_err());
let weights = [1.0];
let result = rrf_scoring(responses, DEFAULT_RRF_K, Some(&weights));
assert!(result.is_err());
}
#[test]
fn test_rrf_scoring_zero_weight() {
let responses = vec![
vec![make_scored_point(1, 0.9)],
vec![make_scored_point(2, 0.9)],
];
let weights = [1.0, 0.0];
let scored_points = rrf_scoring(responses, DEFAULT_RRF_K, Some(&weights)).unwrap();
let p1 = scored_points.iter().find(|p| p.id == 1.into()).unwrap();
let p2 = scored_points.iter().find(|p| p.id == 2.into()).unwrap();
assert_eq!(p1.score, 0.5); assert_eq!(p2.score, 0.0); }
}