use thiserror::Error;
use crate::{
ANNError, error,
graph::{AdjacencyList, config::PruneKind, internal::SortedNeighbors},
neighbor::Neighbor,
utils::{IntoUsize, VectorId},
};
#[derive(Debug, Clone, Copy)]
pub(crate) struct Options {
pub(in crate::graph) force_saturate: bool,
}
#[derive(Debug)]
pub(crate) struct Scratch<I>
where
I: VectorId,
{
pub(in crate::graph) pool: Vec<Neighbor<I>>,
pub(in crate::graph) states: Vec<State>,
pub(in crate::graph) neighbors: AdjacencyList<I>,
}
impl<I> Scratch<I>
where
I: VectorId,
{
pub(in crate::graph) fn new() -> Self {
Self {
pool: Vec::new(),
states: Vec::new(),
neighbors: AdjacencyList::new(),
}
}
pub(in crate::graph) fn as_context(&mut self, max_candidates: usize) -> Context<'_, I> {
Context {
pool: SortedNeighbors::new(&mut self.pool, max_candidates),
states: &mut self.states,
neighbors: &mut self.neighbors,
}
}
}
#[derive(Debug)]
pub(crate) struct Context<'ctx, I>
where
I: VectorId,
{
pub(in crate::graph) pool: SortedNeighbors<'ctx, I>,
pub(in crate::graph) states: &'ctx mut Vec<State>,
pub(in crate::graph) neighbors: &'ctx mut AdjacencyList<I>,
}
#[derive(Debug, Clone, Copy, Default)]
pub(crate) struct State {
pub(in crate::graph) occlude_factor: f32,
pub(in crate::graph) last_checked: u16,
pub(in crate::graph) neighbor: u16,
}
pub(in crate::graph) fn robust_prune<V, D>(
candidates: SortedNeighbors<'_, Option<V>>,
states: &mut [State],
degree: usize,
alpha: f32,
prune_kind: PruneKind,
mut compute_distance: D,
) -> usize
where
D: FnMut(&V, &V) -> f32,
{
let mut current_alpha = 1.0f32;
let increment_factor = alpha.min(1.2);
let mut found = 0;
while found < degree {
for (i, neighbor) in candidates.iter().enumerate() {
if found >= degree {
break;
}
let State {
mut occlude_factor,
mut last_checked,
..
} = states[i];
if occlude_factor > current_alpha {
continue;
}
let neighbor_distance = neighbor.distance();
let neighbor = match neighbor.id() {
Some(n) => n,
None => {
debug_assert!(states.get(i).is_some(), "index {i} is out of bounds");
unsafe { states.get_unchecked_mut(i) }.occlude_factor = f32::MAX;
continue;
}
};
while last_checked as usize != found {
let result_position = states[last_checked as usize].neighbor.into_usize();
last_checked += 1;
if result_position >= i {
debug_assert!(states.get(i).is_some(), "index {i} is out of bounds");
unsafe { states.get_unchecked_mut(i) }.last_checked = last_checked;
continue;
}
let distance = match candidates[result_position].id() {
Some(v) => compute_distance(neighbor, v),
None => f32::MAX,
};
occlude_factor = prune_kind.update_occlude_factor(
*neighbor_distance,
distance,
occlude_factor,
current_alpha,
);
if occlude_factor > current_alpha {
break;
}
}
debug_assert!(states.get(i).is_some(), "index {i} is out of bounds");
let state = unsafe { states.get_unchecked_mut(i) };
state.last_checked = last_checked;
if occlude_factor > current_alpha {
state.occlude_factor = occlude_factor;
continue;
}
state.occlude_factor = f32::MAX;
states[found].neighbor = i as u16;
found += 1;
}
if current_alpha == alpha {
break;
}
current_alpha = (current_alpha * increment_factor).min(alpha);
}
found
}
#[derive(Debug, Clone, Copy, Error)]
#[error("retrieval of main vector id {} failed during prune aggregation", self.0)]
pub(crate) struct FailedVectorRetrieval<I>(I)
where
I: VectorId;
impl<I> error::TransientError<ANNError> for FailedVectorRetrieval<I>
where
I: VectorId,
{
fn acknowledge<D>(self, _why: D)
where
D: std::fmt::Display,
{
}
#[track_caller]
#[inline(never)]
fn escalate<D>(self, why: D) -> ANNError
where
D: std::fmt::Display,
{
ANNError::new(self).context(why.to_string())
}
}
#[derive(Debug)]
pub(crate) enum ListError<I>
where
I: VectorId,
{
FailedVectorRetrieval(FailedVectorRetrieval<I>),
Other(ANNError),
}
impl<I> ListError<I>
where
I: VectorId,
{
pub(in crate::graph) fn failed_retrieval(id: I) -> Self {
Self::FailedVectorRetrieval(FailedVectorRetrieval(id))
}
}
impl<I> From<ANNError> for ListError<I>
where
I: VectorId,
{
fn from(err: ANNError) -> Self {
Self::Other(err)
}
}
impl<I> error::ToRanked for ListError<I>
where
I: VectorId,
{
type Transient = FailedVectorRetrieval<I>;
type Error = ANNError;
fn to_ranked(self) -> error::RankedError<Self::Transient, Self::Error> {
match self {
Self::FailedVectorRetrieval(err) => error::RankedError::Transient(err),
Self::Other(err) => error::RankedError::Error(err),
}
}
fn from_transient(transient: Self::Transient) -> Self {
Self::FailedVectorRetrieval(transient)
}
fn from_error(error: Self::Error) -> Self {
Self::Other(error)
}
}