weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::mapped_search::MappedScratch;
use super::{MappedGraph, MappedVectorIndex};
use crate::mmap::Mapping;
use crate::vector::Candidate;
use std::cmp::Reverse;

impl MappedGraph {
    pub(super) fn node_layers<'a>(&self, mapping: &'a Mapping) -> &'a [u64] {
        mapping
            .u64_slice(self.node_layers_offset, self.node_count + 1)
            .expect("validated graph node-layer section")
    }

    pub(super) fn layer_neighbors<'a>(&self, mapping: &'a Mapping) -> &'a [u64] {
        mapping
            .u64_slice(self.layer_neighbors_offset, self.total_layers + 1)
            .expect("validated graph layer-neighbor section")
    }

    pub(super) fn neighbors<'a>(&self, mapping: &'a Mapping) -> &'a [u32] {
        mapping
            .u32_slice(self.neighbors_offset, self.neighbor_count)
            .expect("validated graph neighbor section")
    }

    fn node_neighbors<'a>(&self, mapping: &'a Mapping, node: usize, level: usize) -> &'a [u32] {
        let node_layers = self.node_layers(mapping);
        let layer = usize::try_from(node_layers[node]).expect("validated layer offset") + level;
        let layer_neighbors = self.layer_neighbors(mapping);
        let start = usize::try_from(layer_neighbors[layer]).expect("validated neighbor offset");
        let end = usize::try_from(layer_neighbors[layer + 1]).expect("validated neighbor offset");
        &self.neighbors(mapping)[start..end]
    }

    #[allow(clippy::too_many_arguments)]
    pub(super) fn search_into<F>(
        &self,
        index: &MappedVectorIndex,
        query: &[f32],
        inverse_norm: f32,
        expansion: usize,
        scratch: &mut MappedScratch,
        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(index, query, inverse_norm, entry, level);
        }
        self.explore_layer(
            index,
            query,
            inverse_norm,
            entry,
            expansion,
            0,
            scratch,
            accepts,
        );
        output.extend(scratch.results.iter().copied());
    }

    fn greedy(
        &self,
        index: &MappedVectorIndex,
        query: &[f32],
        inverse_norm: f32,
        mut best: usize,
        level: usize,
    ) -> usize {
        let mut best_candidate =
            Candidate::new(index.distance_query(best, query, inverse_norm), best);
        loop {
            let mut improved = false;
            for neighbor in self.node_neighbors(&index.mapping, best, level) {
                let neighbor = *neighbor as usize;
                let candidate = Candidate::new(
                    index.distance_query(neighbor, query, inverse_norm),
                    neighbor,
                );
                if candidate < best_candidate {
                    best = neighbor;
                    best_candidate = candidate;
                    improved = true;
                }
            }
            if !improved {
                return best;
            }
        }
    }

    #[allow(clippy::too_many_arguments)]
    fn explore_layer<F>(
        &self,
        index: &MappedVectorIndex,
        query: &[f32],
        inverse_norm: f32,
        entry: usize,
        expansion: usize,
        level: usize,
        scratch: &mut MappedScratch,
        accepts: &F,
    ) where
        F: Fn(u64) -> bool,
    {
        scratch.begin();
        let first = Candidate::new(index.distance_query(entry, query, inverse_norm), entry);
        scratch.mark(entry);
        scratch.candidates.push(Reverse(first));
        if accepts(index.key_slice()[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.node_neighbors(&index.mapping, current.index(), level) {
                let neighbor = *neighbor as usize;
                if !scratch.mark(neighbor) {
                    continue;
                }
                let candidate = Candidate::new(
                    index.distance_query(neighbor, query, inverse_norm),
                    neighbor,
                );
                if scratch.results.len() < expansion
                    || scratch
                        .results
                        .peek()
                        .is_some_and(|worst| candidate < *worst)
                {
                    scratch.candidates.push(Reverse(candidate));
                    if accepts(index.key_slice()[neighbor]) {
                        scratch.results.push(candidate);
                        if scratch.results.len() > expansion {
                            scratch.results.pop();
                        }
                    }
                }
            }
        }
    }
}