use super::graph_helpers::from_node;
use super::{Graph, SearchScratch};
use crate::vector::{Candidate, VectorStore};
use std::cmp::Reverse;
impl Graph {
pub(super) fn greedy_index(
&self,
vectors: &VectorStore,
target: usize,
entry: usize,
level: usize,
) -> usize {
self.greedy(entry, level, |candidate| {
vectors.distance_indices(target, candidate)
})
}
fn greedy_query(
&self,
vectors: &VectorStore,
query: &[f32],
inverse_norm: f32,
entry: usize,
level: usize,
) -> usize {
self.greedy(entry, level, |candidate| {
vectors.distance_query(candidate, query, inverse_norm)
})
}
fn greedy<F>(&self, mut best: usize, level: usize, mut distance: F) -> usize
where
F: FnMut(usize) -> f32,
{
let mut best_candidate = Candidate::new(distance(best), best);
loop {
let mut improved = false;
for neighbor in &self.nodes[best].layers[level] {
let neighbor = from_node(*neighbor);
let candidate = Candidate::new(distance(neighbor), neighbor);
if candidate < best_candidate {
best = neighbor;
best_candidate = candidate;
improved = true;
}
}
if !improved {
return best;
}
}
}
#[allow(clippy::too_many_arguments)]
pub(super) fn search_into<F>(
&self,
vectors: &VectorStore,
query: &[f32],
inverse_norm: f32,
expansion: usize,
scratch: &mut SearchScratch,
output: &mut Vec<Candidate>,
accepts: &F,
) where
F: Fn(u64) -> bool,
{
let Some(mut entry) = self.entry else {
return;
};
for level in (1..=self.max_level).rev() {
entry = self.greedy_query(vectors, query, inverse_norm, entry, level);
}
self.explore_layer(
entry,
expansion,
0,
scratch,
|candidate| vectors.distance_query(candidate, query, inverse_norm),
|candidate| accepts(vectors.key(candidate)),
);
output.extend(scratch.results.iter().copied());
}
pub(super) fn search_layer<F>(
&self,
_vectors: &VectorStore,
entry: usize,
expansion: usize,
level: usize,
scratch: &mut SearchScratch,
mut distance: F,
) -> Vec<Candidate>
where
F: FnMut(usize) -> f32,
{
self.explore_layer(entry, expansion, level, scratch, &mut distance, |_| true);
let mut results = scratch.results.clone().into_vec();
results.sort_unstable();
results
}
fn explore_layer<F, A>(
&self,
entry: usize,
expansion: usize,
level: usize,
scratch: &mut SearchScratch,
mut distance: F,
accepts: A,
) where
F: FnMut(usize) -> f32,
A: Fn(usize) -> bool,
{
scratch.begin();
let first = Candidate::new(distance(entry), entry);
scratch.mark(entry);
scratch.candidates.push(Reverse(first));
if accepts(entry) {
scratch.results.push(first);
}
while let Some(Reverse(current)) = scratch.candidates.pop() {
if scratch.results.len() >= expansion
&& scratch.results.peek().is_some_and(|worst| current > *worst)
{
break;
}
for neighbor in &self.nodes[current.index()].layers[level] {
let neighbor = from_node(*neighbor);
if !scratch.mark(neighbor) {
continue;
}
let candidate = Candidate::new(distance(neighbor), neighbor);
if scratch.results.len() < expansion
|| scratch
.results
.peek()
.is_some_and(|worst| candidate < *worst)
{
scratch.candidates.push(Reverse(candidate));
if accepts(neighbor) {
scratch.results.push(candidate);
if scratch.results.len() > expansion {
scratch.results.pop();
}
}
}
}
}
}
}