weavatrix-search-vector 0.3.1

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
use super::{MAX_ROUTING_PROBES, ROUTING_BITS, ROUTING_PROBES, RoutingCandidate, VectorStore};
use crate::error::SearchError;
use std::cmp::Reverse;
use std::collections::BinaryHeap;

impl VectorStore {
    pub(crate) fn routing_code(&self, vector: &[f32]) -> u16 {
        let sums = self.routing_sums(vector);
        sums.iter().enumerate().fold(0_u16, |code, (plane, sum)| {
            if *sum >= 0.0 {
                code | (1_u16 << plane)
            } else {
                code
            }
        })
    }

    pub(crate) fn routing_probes(&self, vector: &[f32]) -> [u16; ROUTING_PROBES] {
        let sums = self.routing_sums(vector);
        let code = sums.iter().enumerate().fold(0_u16, |code, (plane, sum)| {
            if *sum >= 0.0 {
                code | (1_u16 << plane)
            } else {
                code
            }
        });
        let mut planes = std::array::from_fn::<_, ROUTING_BITS, _>(|plane| plane);
        planes.sort_unstable_by(|left, right| {
            sums[*left]
                .abs()
                .total_cmp(&sums[*right].abs())
                .then_with(|| left.cmp(right))
        });
        let mut probes = [code; ROUTING_PROBES];
        for (probe, plane) in probes[1..].iter_mut().zip(planes) {
            *probe = code ^ (1_u16 << plane);
        }
        probes
    }

    pub(crate) fn routing_probes_into(
        &self,
        vector: &[f32],
        limit: usize,
        output: &mut Vec<u16>,
        heap: &mut BinaryHeap<Reverse<RoutingCandidate>>,
    ) -> Result<(), SearchError> {
        if limit <= ROUTING_PROBES {
            output.clear();
            output
                .try_reserve(limit)
                .map_err(|_| SearchError::AllocationFailed)?;
            output.extend(self.routing_probes(vector).into_iter().take(limit));
            return Ok(());
        }
        routing_probes_from_signs(
            vector,
            self.routing_signs.iter().copied(),
            limit,
            output,
            heap,
        )
    }

    fn routing_sums(&self, vector: &[f32]) -> [f32; ROUTING_BITS] {
        debug_assert_eq!(vector.len(), self.routing_signs.len());
        let mut sums = [0.0_f32; ROUTING_BITS];
        for (value, signs) in vector.iter().copied().zip(&self.routing_signs) {
            for (plane, sum) in sums.iter_mut().enumerate() {
                if signs & (1_u16 << plane) == 0 {
                    *sum += value;
                } else {
                    *sum -= value;
                }
            }
        }
        sums
    }

    pub(crate) fn stored_routing_code(&self, index: usize) -> u16 {
        self.routing_code(self.vector(index))
    }
}

pub(crate) fn routing_probes_from_signs(
    vector: &[f32],
    signs: impl IntoIterator<Item = u16>,
    limit: usize,
    output: &mut Vec<u16>,
    heap: &mut BinaryHeap<Reverse<RoutingCandidate>>,
) -> Result<(), SearchError> {
    output.clear();
    let mut sums = [0.0_f32; ROUTING_BITS];
    for (value, signs) in vector.iter().copied().zip(signs) {
        for (plane, sum) in sums.iter_mut().enumerate() {
            if signs & (1_u16 << plane) == 0 {
                *sum += value;
            } else {
                *sum -= value;
            }
        }
    }
    let base = sums.iter().enumerate().fold(0_u16, |code, (plane, sum)| {
        if *sum >= 0.0 {
            code | (1_u16 << plane)
        } else {
            code
        }
    });
    if limit > 64 {
        return routing_probes_exhaustive(&sums, base, limit, output, heap);
    }
    let mut planes = std::array::from_fn::<_, ROUTING_BITS, _>(|plane| plane);
    planes.sort_unstable_by(|left, right| {
        sums[*left]
            .abs()
            .total_cmp(&sums[*right].abs())
            .then_with(|| left.cmp(right))
    });
    output.clear();
    output
        .try_reserve(limit)
        .map_err(|_| SearchError::AllocationFailed)?;
    output.push(base);
    if limit == 1 {
        return Ok(());
    }
    heap.clear();
    heap.try_reserve(limit.saturating_mul(2))
        .map_err(|_| SearchError::AllocationFailed)?;
    let first_plane = planes[0];
    heap.push(Reverse(RoutingCandidate {
        score: sums[first_plane].abs(),
        mask: 1_u16 << first_plane,
        last: 0,
        bits: 1,
    }));
    extend_routing_probes(&sums, &planes, base, limit, output, heap);
    Ok(())
}

fn extend_routing_probes(
    sums: &[f32; ROUTING_BITS],
    planes: &[usize; ROUTING_BITS],
    base: u16,
    limit: usize,
    output: &mut Vec<u16>,
    heap: &mut BinaryHeap<Reverse<RoutingCandidate>>,
) {
    while output.len() < limit {
        let Some(Reverse(candidate)) = heap.pop() else {
            break;
        };
        output.push(base ^ candidate.mask);
        let next = usize::from(candidate.last) + 1;
        if next == ROUTING_BITS {
            continue;
        }
        let previous_plane = planes[usize::from(candidate.last)];
        let next_plane = planes[next];
        let next_bit = 1_u16 << next_plane;
        let next_position = u8::try_from(next).expect("routing bit index fits in u8");
        if candidate.bits < 3 {
            let mask = candidate.mask | next_bit;
            heap.push(Reverse(RoutingCandidate {
                score: routing_mask_score(sums, mask),
                mask,
                last: next_position,
                bits: candidate.bits + 1,
            }));
        }
        let mask = (candidate.mask ^ (1_u16 << previous_plane)) | next_bit;
        heap.push(Reverse(RoutingCandidate {
            score: routing_mask_score(sums, mask),
            mask,
            last: next_position,
            bits: candidate.bits,
        }));
    }
    heap.clear();
}

fn routing_probes_exhaustive(
    sums: &[f32; ROUTING_BITS],
    base: u16,
    limit: usize,
    output: &mut Vec<u16>,
    heap: &mut BinaryHeap<Reverse<RoutingCandidate>>,
) -> Result<(), SearchError> {
    heap.clear();
    heap.try_reserve(MAX_ROUTING_PROBES)
        .map_err(|_| SearchError::AllocationFailed)?;
    heap.push(Reverse(RoutingCandidate::new(0.0, 0, 0)));
    for first in 0..ROUTING_BITS {
        let first_mask = 1_u16 << first;
        heap.push(Reverse(RoutingCandidate::new(
            routing_mask_score(sums, first_mask),
            first_mask,
            1,
        )));
        for second in first + 1..ROUTING_BITS {
            let second_mask = first_mask | (1_u16 << second);
            heap.push(Reverse(RoutingCandidate::new(
                routing_mask_score(sums, second_mask),
                second_mask,
                2,
            )));
            for third in second + 1..ROUTING_BITS {
                let mask = second_mask | (1_u16 << third);
                heap.push(Reverse(RoutingCandidate::new(
                    routing_mask_score(sums, mask),
                    mask,
                    3,
                )));
            }
        }
    }
    output
        .try_reserve(limit)
        .map_err(|_| SearchError::AllocationFailed)?;
    output.extend(
        std::iter::from_fn(|| heap.pop())
            .take(limit)
            .map(|Reverse(candidate)| base ^ candidate.mask),
    );
    heap.clear();
    Ok(())
}

fn routing_mask_score(sums: &[f32; ROUTING_BITS], mask: u16) -> f32 {
    sums.iter()
        .enumerate()
        .filter(|(plane, _)| mask & (1_u16 << plane) != 0)
        .map(|(_, sum)| sum.abs())
        .sum()
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::config::DistanceMetric;

    #[test]
    fn bounded_probe_heap_matches_exhaustive_three_flip_order() {
        let stored = [0.5_f32, -0.25, 0.75, 1.0, -0.6, 0.4, 0.2, -0.9];
        let store =
            VectorStore::build(stored.len(), DistanceMetric::Cosine, &[(7, &stored)]).unwrap();
        let query = [0.9_f32, -0.4, 0.1, 0.7, -0.8, 0.3, 0.6, -0.2];
        let sums = store.routing_sums(&query);
        let base = sums.iter().enumerate().fold(0_u16, |code, (plane, sum)| {
            if *sum >= 0.0 {
                code | (1_u16 << plane)
            } else {
                code
            }
        });
        let mut exhaustive = vec![(0.0_f32, 0_u16)];
        for first in 0..ROUTING_BITS {
            exhaustive.push((sums[first].abs(), 1_u16 << first));
            for second in first + 1..ROUTING_BITS {
                exhaustive.push((
                    sums[first].abs() + sums[second].abs(),
                    (1_u16 << first) | (1_u16 << second),
                ));
                for third in second + 1..ROUTING_BITS {
                    exhaustive.push((
                        sums[first].abs() + sums[second].abs() + sums[third].abs(),
                        (1_u16 << first) | (1_u16 << second) | (1_u16 << third),
                    ));
                }
            }
        }
        exhaustive.sort_unstable_by(|left, right| {
            left.0
                .total_cmp(&right.0)
                .then_with(|| left.1.cmp(&right.1))
        });
        let expected = exhaustive
            .into_iter()
            .map(|(_, mask)| base ^ mask)
            .collect::<Vec<_>>();
        let mut actual = Vec::new();
        let mut heap = BinaryHeap::new();
        for limit in [6, 12, 64, MAX_ROUTING_PROBES] {
            store
                .routing_probes_into(&query, limit, &mut actual, &mut heap)
                .unwrap();
            assert_eq!(actual, expected[..limit]);
        }
    }
}