weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
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();
                        }
                    }
                }
            }
        }
    }
}