mod lazy_matrix;
#[cfg(test)]
mod tests;
use crate::common::counter::hardware_accumulator::HwMeasurementAcc;
use crate::common::counter::hardware_counter::HardwareCounterCell;
use crate::common::types::ScoreType;
use indexmap::IndexSet;
use itertools::Itertools as _;
use ordered_float::OrderedFloat;
use crate::segment::common::operation_error::{OperationError, OperationResult};
use crate::segment::data_types::vectors::{QueryVector, VectorInternal, VectorRef};
use crate::segment::types::{Distance, MultiVectorConfig, ScoredPoint};
use crate::segment::vector_storage::dense::volatile_dense_vector_storage::new_volatile_dense_vector_storage;
use crate::segment::vector_storage::multi_dense::volatile_multi_dense_vector_storage::new_volatile_multi_dense_vector_storage;
use crate::segment::vector_storage::sparse::volatile_sparse_vector_storage::new_volatile_sparse_vector_storage;
use crate::segment::vector_storage::{
VectorStorage as _, VectorStorageEnum, VectorStorageRead as _, new_raw_scorer,
};
use self::lazy_matrix::LazyMatrix;
use super::MmrInternal;
pub fn mmr_from_points_with_vector(
points_with_vector: impl IntoIterator<Item = ScoredPoint>,
mmr: MmrInternal,
distance: Distance,
multivector_config: Option<MultiVectorConfig>,
limit: usize,
hw_measurement_acc: HwMeasurementAcc,
) -> OperationResult<Vec<ScoredPoint>> {
let (vectors, candidates): (Vec<_>, Vec<_>) = points_with_vector
.into_iter()
.unique_by(|p| p.id)
.filter_map(|p| {
let vector = p
.vector
.as_ref()
.and_then(|v| v.get(&mmr.using))
.map(|v| v.to_owned())?;
Some((vector, p))
})
.unzip();
debug_assert_eq!(vectors.len(), candidates.len());
if candidates.is_empty() {
return Ok(candidates);
}
let volatile_storage = create_volatile_storage(
&vectors,
distance,
multivector_config,
hw_measurement_acc.get_counter_cell(),
)?;
if candidates.len() < 2 {
return Ok(candidates);
}
let query_similarities = relevance_similarities(
&volatile_storage,
mmr.vector,
hw_measurement_acc.get_counter_cell(),
)?;
let similarity_matrix = similarity_matrix(&volatile_storage, vectors, hw_measurement_acc)?;
Ok(maximal_marginal_relevance(
candidates,
query_similarities,
similarity_matrix,
mmr.lambda.0,
limit,
))
}
fn create_volatile_storage(
vectors: &[VectorInternal],
distance: Distance,
multivector_config: Option<MultiVectorConfig>,
hw_counter: HardwareCounterCell,
) -> OperationResult<VectorStorageEnum> {
let mut volatile_storage = {
match &vectors[0] {
VectorInternal::Dense(vector) => {
new_volatile_dense_vector_storage(vector.len(), distance)
}
VectorInternal::MultiDense(typed_multi_dense_vector) => {
let multivector_config = multivector_config.ok_or_else(|| {
OperationError::service_error(
"multivectors are present, but no multivector config provided",
)
})?;
new_volatile_multi_dense_vector_storage(
typed_multi_dense_vector.dim,
distance,
multivector_config,
)
}
VectorInternal::Sparse(_) => new_volatile_sparse_vector_storage(),
}
};
for (key, vector) in (0..).zip(vectors) {
volatile_storage.insert_vector(key, VectorRef::from(vector), &hw_counter)?;
}
Ok(volatile_storage)
}
fn relevance_similarities(
volatile_storage: &VectorStorageEnum,
query_vector: VectorInternal,
hw_counter: HardwareCounterCell,
) -> OperationResult<Vec<ScoreType>> {
let query = QueryVector::Nearest(query_vector);
let query_scorer = new_raw_scorer(query, volatile_storage, hw_counter)?;
let ids: Vec<_> = (0..volatile_storage.total_vector_count() as u32).collect();
let mut similarities = vec![0.0; ids.len()];
query_scorer.score_points(&ids, &mut similarities);
Ok(similarities)
}
fn similarity_matrix(
volatile_storage: &VectorStorageEnum,
vectors: Vec<VectorInternal>,
hw_measurement_acc: HwMeasurementAcc,
) -> OperationResult<LazyMatrix<'_>> {
let num_vectors = vectors.len();
debug_assert!(
num_vectors >= 2,
"There should be at least two vectors to calculate similarity matrix"
);
if num_vectors < 2 {
return Err(OperationError::service_error(
"There should be at least two vectors to calculate similarity matrix",
));
}
LazyMatrix::new(vectors, volatile_storage, hw_measurement_acc)
}
fn maximal_marginal_relevance(
candidates: Vec<ScoredPoint>,
query_similarities: Vec<ScoreType>,
mut similarity_matrix: LazyMatrix,
lambda: f32,
limit: usize,
) -> Vec<ScoredPoint> {
let num_candidates = candidates.len();
if num_candidates == 0 || limit == 0 {
return Vec::new();
}
let mut selected_indices = Vec::with_capacity(limit);
let mut remaining_indices: IndexSet<usize, ahash::RandomState> = (0..num_candidates).collect();
if let Some(best_idx) = remaining_indices
.iter()
.max_by_key(|&candidate_idx| OrderedFloat(query_similarities[*candidate_idx]))
.copied()
{
selected_indices.push(best_idx);
remaining_indices.swap_remove(&best_idx);
}
while selected_indices.len() < limit && !remaining_indices.is_empty() {
let best_candidate = remaining_indices
.iter()
.map(|&candidate_idx| {
let relevance_score = query_similarities[candidate_idx];
debug_assert!(
selected_indices
.iter()
.all(|&selected_idx| selected_idx != candidate_idx)
);
let max_similarity_to_selected = selected_indices
.iter()
.map(|selected_idx| {
similarity_matrix.get_similarity(candidate_idx, *selected_idx)
})
.max_by_key(|&sim| OrderedFloat(sim))
.unwrap_or(0.0);
let mmr_score =
lambda * relevance_score - (1.0 - lambda) * max_similarity_to_selected;
(candidate_idx, mmr_score)
})
.max_by_key(|(_candidate_idx, mmr_score)| OrderedFloat(*mmr_score));
if let Some((selected_idx, _mmr_score)) = best_candidate {
remaining_indices.swap_remove(&selected_idx);
selected_indices.push(selected_idx);
} else {
break;
}
}
selected_indices
.into_iter()
.map(|idx| {
candidates[idx].clone()
})
.collect()
}