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]);
}
}
}