use std::collections::HashMap;
use crate::common::counter::hardware_accumulator::HwMeasurementAcc;
use ordered_float::OrderedFloat;
use rstest::rstest;
use crate::segment::data_types::vectors::{
MultiDenseVectorInternal, VectorInternal, VectorStructInternal,
};
use crate::segment::types::{
Distance, MultiVectorComparator, MultiVectorConfig, PointIdType, ScoredPoint, VectorNameBuf,
};
use crate::sparse::common::sparse_vector::SparseVector;
use strum::IntoEnumIterator;
use super::mmr_from_points_with_vector;
use crate::shard::query::MmrInternal;
fn create_scored_point_with_vector(
id: PointIdType,
vector: Vec<f32>,
vector_name: Option<&str>,
) -> ScoredPoint {
let vector_internal = VectorInternal::Dense(vector);
let mut vectors = HashMap::new();
let name = vector_name.unwrap_or("");
vectors.insert(name.to_string(), vector_internal);
ScoredPoint {
id,
version: 0,
score: 0.0,
payload: None,
vector: Some(VectorStructInternal::Named(vectors)),
shard_key: None,
order_value: None,
}
}
fn create_scored_point_without_vector(id: PointIdType) -> ScoredPoint {
ScoredPoint {
id,
version: 0,
score: 0.0,
payload: None,
vector: None,
shard_key: None,
order_value: None,
}
}
fn create_scored_point_with_sparse_vector(
id: PointIdType,
indices: Vec<u32>,
values: Vec<f32>,
vector_name: Option<&str>,
) -> ScoredPoint {
let sparse_vector = SparseVector::new(indices, values).expect("Valid sparse vector");
let vector_internal = VectorInternal::Sparse(sparse_vector);
let mut vectors = HashMap::new();
let name = vector_name.unwrap_or("");
vectors.insert(name.to_string(), vector_internal);
ScoredPoint {
id,
version: 0,
score: 0.0,
payload: None,
vector: Some(VectorStructInternal::Named(vectors)),
shard_key: None,
order_value: None,
}
}
fn create_scored_point_with_multi_vector(
id: PointIdType,
vectors: Vec<Vec<f32>>,
vector_name: Option<&str>,
) -> ScoredPoint {
let multi_vector = MultiDenseVectorInternal::new_unchecked(vectors);
let vector_internal = VectorInternal::MultiDense(multi_vector);
let mut vector_map = HashMap::new();
let name = vector_name.unwrap_or("");
vector_map.insert(name.to_string(), vector_internal);
ScoredPoint {
id,
version: 0,
score: 0.0,
payload: None,
vector: Some(VectorStructInternal::Named(vector_map)),
shard_key: None,
order_value: None,
}
}
#[rstest]
#[case::full_relevance(1.0, &[1, 2, 3])]
#[case::balanced(0.5, &[1, 3, 2])]
#[case::more_diversity(0.01, &[1, 5, 4])]
fn test_mmr_lambda(#[case] lambda: f32, #[case] expected_order: &[u64]) {
let distance = Distance::Euclid;
let points = vec![
create_scored_point_with_vector(1.into(), vec![1.0, 0.05], None),
create_scored_point_with_vector(2.into(), vec![0.95, 0.15], None),
create_scored_point_with_vector(3.into(), vec![0.8, 0.0], None),
create_scored_point_with_vector(4.into(), vec![0.85, 0.25], None),
create_scored_point_with_vector(5.into(), vec![1.0, 0.5], None),
];
let mmr = MmrInternal {
vector: vec![1.0, 0.0].into(),
using: VectorNameBuf::from(""),
lambda: OrderedFloat(lambda),
candidates_limit: 100,
};
let result = mmr_from_points_with_vector(
points.clone(),
mmr,
distance,
None,
3,
HwMeasurementAcc::new(),
);
let scored_points = result.unwrap();
assert_eq!(scored_points.len(), 3);
let selected_ids: Vec<_> = scored_points.iter().map(|p| p.id).collect();
let expected_ids = expected_order
.iter()
.map(|&id| id.into())
.collect::<Vec<_>>();
assert_eq!(selected_ids, expected_ids);
assert_ne!(scored_points[0].score, 0.9); }
#[test]
fn test_mmr_less_than_two_points() {
let distance = Distance::Cosine;
let mmr = MmrInternal {
vector: vec![1.0, 0.0].into(), using: VectorNameBuf::from(""),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
let empty_points = vec![];
let result = mmr_from_points_with_vector(
empty_points,
mmr.clone(),
distance,
None,
5,
HwMeasurementAcc::new(),
);
assert!(result.is_ok());
assert_eq!(result.unwrap().len(), 0);
let single_point = vec![create_scored_point_with_vector(
1.into(),
vec![1.0, 0.0, 0.0],
None,
)];
let result = mmr_from_points_with_vector(
single_point,
mmr,
distance,
None,
5,
HwMeasurementAcc::new(),
);
assert!(result.is_ok());
let scored_points = result.unwrap();
assert_eq!(scored_points.len(), 1);
assert_eq!(scored_points[0].id, 1.into());
}
#[test]
fn test_mmr_points_without_required_vector() {
let distance = Distance::Cosine;
let points = vec![
create_scored_point_with_vector(1.into(), vec![1.0, 0.0, 0.0], Some("custom")),
create_scored_point_without_vector(2.into()), create_scored_point_with_vector(3.into(), vec![0.0, 1.0, 0.0], Some("other")), create_scored_point_with_vector(4.into(), vec![0.0, 0.0, 1.0], Some("custom")),
];
let mmr = MmrInternal {
vector: vec![1.0, 0.0, 0.0].into(),
using: VectorNameBuf::from("custom"),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
let result =
mmr_from_points_with_vector(points, mmr, distance, None, 5, HwMeasurementAcc::new());
assert!(result.is_ok());
let scored_points = result.unwrap();
assert_eq!(scored_points.len(), 2);
let selected_ids: Vec<_> = scored_points.iter().map(|p| p.id).collect();
assert!(selected_ids.contains(&(1.into())));
assert!(selected_ids.contains(&(4.into())));
}
#[test]
fn test_mmr_duplicate_points() {
let distance = Distance::Cosine;
let points = vec![
create_scored_point_with_vector(1.into(), vec![1.0, 0.0, 0.0], None),
create_scored_point_with_vector(1.into(), vec![0.5, 0.5, 0.0], None), create_scored_point_with_vector(2.into(), vec![0.0, 1.0, 0.0], None),
];
let mmr = MmrInternal {
vector: vec![1.0, 0.0, 0.0].into(),
using: VectorNameBuf::from(""),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
let result =
mmr_from_points_with_vector(points, mmr, distance, None, 5, HwMeasurementAcc::new());
assert!(result.is_ok());
let scored_points = result.unwrap();
assert_eq!(scored_points.len(), 2);
let unique_ids: std::collections::HashSet<_> = scored_points.iter().map(|p| p.id).collect();
assert_eq!(unique_ids.len(), 2);
}
#[test]
fn test_mmr_dense_vectors() {
let dense_points = vec![
create_scored_point_with_vector(1.into(), vec![1.0, 0.0, 0.0], None),
create_scored_point_with_vector(2.into(), vec![0.0, 1.0, 0.0], None),
create_scored_point_with_vector(3.into(), vec![0.0, 0.0, 1.0], None),
];
let mmr = MmrInternal {
vector: vec![1.0, 0.0, 0.0].into(),
using: VectorNameBuf::from(""),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
for distance in Distance::iter() {
let result = mmr_from_points_with_vector(
dense_points.clone(),
mmr.clone(),
distance,
None,
3,
HwMeasurementAcc::new(),
);
assert!(
result.is_ok(),
"Dense vectors failed for distance metric: {distance:?}"
);
let scored_points = result.unwrap();
assert_eq!(scored_points.len(), 3);
}
}
#[test]
fn test_mmr_sparse_vectors() {
let distance = Distance::Dot;
let sparse_vector_name = "sparse";
let sparse_points = vec![
create_scored_point_with_sparse_vector(
4.into(),
vec![0, 2, 5],
vec![1.0, 0.5, 0.3],
Some(sparse_vector_name),
),
create_scored_point_with_sparse_vector(
5.into(),
vec![1, 3, 4],
vec![0.8, 0.6, 0.4],
Some(sparse_vector_name),
),
create_scored_point_with_sparse_vector(
6.into(),
vec![0, 1, 6],
vec![0.7, 0.9, 0.2],
Some(sparse_vector_name),
),
];
let sparse_mmr = MmrInternal {
vector: SparseVector::new(vec![0, 2, 5], vec![1.0, 0.5, 0.3])
.unwrap()
.into(),
using: VectorNameBuf::from(sparse_vector_name),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
let sparse_result = mmr_from_points_with_vector(
sparse_points,
sparse_mmr,
distance,
None,
3,
HwMeasurementAcc::new(),
)
.unwrap();
assert_eq!(sparse_result.len(), 3);
}
#[test]
fn test_mmr_multi_vector() {
let multi_vector_config = MultiVectorConfig {
comparator: MultiVectorComparator::MaxSim,
};
let multi_vector_name = "multi";
let multi_points = vec![
create_scored_point_with_multi_vector(
7.into(),
vec![vec![1.0, 0.0], vec![0.0, 1.0]],
Some(multi_vector_name),
),
create_scored_point_with_multi_vector(
8.into(),
vec![vec![0.0, 1.0], vec![1.0, 0.0]],
Some(multi_vector_name),
),
create_scored_point_with_multi_vector(
9.into(),
vec![vec![1.0, 1.0], vec![0.0, 0.0]],
Some(multi_vector_name),
),
];
let multi_mmr = MmrInternal {
vector: MultiDenseVectorInternal::new(vec![1.0, 0.0, 0.0, 1.0], 2).into(),
using: VectorNameBuf::from(multi_vector_name),
lambda: OrderedFloat(0.5),
candidates_limit: 100,
};
for distance in Distance::iter() {
let multi_result = mmr_from_points_with_vector(
multi_points.clone(),
multi_mmr.clone(),
distance,
Some(multi_vector_config),
3,
HwMeasurementAcc::new(),
);
assert!(
multi_result.is_ok(),
"Multi-vectors failed for distance metric: {distance:?}"
);
let multi_scored_points = multi_result.unwrap();
assert_eq!(multi_scored_points.len(), 3);
}
}