use std::cell::RefCell;
use std::collections::{BinaryHeap, HashSet};
pub(crate) use crate::prefetch::prefetch_read_data;
#[cfg(feature = "benchmark")]
thread_local! {
static SEARCH_COUNTERS: RefCell<HnswSearchCounters> = const {
RefCell::new(HnswSearchCounters::new())
};
}
#[cfg(feature = "benchmark")]
#[derive(Clone, Copy, Debug, Default)]
pub struct HnswSearchCounters {
pub distance_evals: u64,
pub result_insertions: u64,
pub result_replacements: u64,
pub result_rejections: u64,
pub candidate_pushes: u64,
pub candidate_pops: u64,
pub frontier_retain_calls: u64,
pub frontier_pruned_candidates: u64,
pub max_frontier_len: usize,
}
#[cfg(feature = "benchmark")]
impl HnswSearchCounters {
const fn new() -> Self {
Self {
distance_evals: 0,
result_insertions: 0,
result_replacements: 0,
result_rejections: 0,
candidate_pushes: 0,
candidate_pops: 0,
frontier_retain_calls: 0,
frontier_pruned_candidates: 0,
max_frontier_len: 0,
}
}
}
#[cfg(feature = "benchmark")]
pub fn reset_search_counters() {
SEARCH_COUNTERS.with(|cell| {
*cell.borrow_mut() = HnswSearchCounters::new();
});
}
#[cfg(feature = "benchmark")]
pub fn take_search_counters() -> HnswSearchCounters {
SEARCH_COUNTERS.with(|cell| {
let mut counters = cell.borrow_mut();
let out = *counters;
*counters = HnswSearchCounters::new();
out
})
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_distance_evals(count: usize) {
SEARCH_COUNTERS.with(|cell| {
cell.borrow_mut().distance_evals += count as u64;
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_result_insertion() {
SEARCH_COUNTERS.with(|cell| {
cell.borrow_mut().result_insertions += 1;
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_result_replacement() {
SEARCH_COUNTERS.with(|cell| {
cell.borrow_mut().result_replacements += 1;
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_result_rejection() {
SEARCH_COUNTERS.with(|cell| {
cell.borrow_mut().result_rejections += 1;
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_candidate_push(candidates_len: usize) {
SEARCH_COUNTERS.with(|cell| {
let mut counters = cell.borrow_mut();
counters.candidate_pushes += 1;
counters.max_frontier_len = counters.max_frontier_len.max(candidates_len);
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_candidate_pop() {
SEARCH_COUNTERS.with(|cell| {
cell.borrow_mut().candidate_pops += 1;
});
}
#[cfg(feature = "benchmark")]
#[inline]
fn record_frontier_retain(before: usize, after: usize) {
SEARCH_COUNTERS.with(|cell| {
let mut counters = cell.borrow_mut();
counters.frontier_retain_calls += 1;
counters.frontier_pruned_candidates += before.saturating_sub(after) as u64;
});
}
const DENSE_VISITED_THRESHOLD: usize = 4_000_000;
const DISTANCE_BATCH_SIZE: usize = 8;
const CACHED_WORST_MIN_EF: usize = 64;
const FRONTIER_PRUNE_MIN_EF: usize = 64;
const FRONTIER_PRUNE_INTERVAL: usize = 64;
pub(crate) enum VisitedSet {
Dense { marks: Vec<u8>, generation: u8 },
Sparse(HashSet<u32>),
}
impl VisitedSet {
fn new(num_nodes: usize, capacity_hint: usize) -> Self {
if num_nodes <= DENSE_VISITED_THRESHOLD {
Self::dense(num_nodes)
} else {
VisitedSet::Sparse(HashSet::with_capacity(capacity_hint))
}
}
pub(crate) fn dense(num_nodes: usize) -> Self {
VisitedSet::Dense {
marks: vec![0u8; num_nodes],
generation: 1,
}
}
fn clear(&mut self) {
match self {
VisitedSet::Dense { marks, generation } => {
if let Some(next) = generation.checked_add(1) {
*generation = next;
} else {
marks.fill(0);
*generation = 1;
}
}
VisitedSet::Sparse(s) => s.clear(),
}
}
#[cfg(test)]
#[inline]
fn contains(&self, id: u32) -> bool {
match self {
VisitedSet::Dense { marks, generation } => {
let idx = id as usize;
idx < marks.len() && marks[idx] == *generation
}
VisitedSet::Sparse(s) => s.contains(&id),
}
}
#[inline]
pub(crate) fn insert(&mut self, id: u32) -> bool {
match self {
VisitedSet::Dense { marks, generation } => {
let idx = id as usize;
debug_assert!(
idx < marks.len(),
"VisitedSet::insert: id {} out of bounds (capacity {})",
id,
marks.len()
);
if idx < marks.len() {
if marks[idx] != *generation {
marks[idx] = *generation;
true
} else {
false
}
} else {
true
}
}
VisitedSet::Sparse(s) => s.insert(id),
}
}
fn prepare(&mut self, num_nodes: usize, capacity_hint: usize) {
match self {
VisitedSet::Dense { marks, .. } if num_nodes <= DENSE_VISITED_THRESHOLD => {
if marks.len() < num_nodes {
marks.resize(num_nodes, 0);
}
self.clear();
}
VisitedSet::Sparse(s) if num_nodes > DENSE_VISITED_THRESHOLD => {
s.clear();
}
_ => {
*self = VisitedSet::new(num_nodes, capacity_hint);
}
}
}
}
thread_local! {
static THREAD_VISITED: RefCell<VisitedSet> = const { RefCell::new(
VisitedSet::Dense { marks: Vec::new(), generation: 1 }
) };
static THREAD_SEARCH_SCRATCH: RefCell<SearchScratch> = const { RefCell::new(SearchScratch::new()) };
}
struct SearchScratch {
candidates: BinaryHeap<MinCandidate>,
results: BinaryHeap<MaxResult>,
}
impl SearchScratch {
const fn new() -> Self {
Self {
candidates: BinaryHeap::new(),
results: BinaryHeap::new(),
}
}
fn prepare(&mut self, ef: usize) {
self.candidates.clear();
self.results.clear();
self.candidates.reserve(ef.saturating_mul(2));
self.results.reserve(ef.saturating_add(1));
}
}
pub(crate) fn with_visited_set<F, R>(num_nodes: usize, capacity_hint: usize, f: F) -> R
where
F: FnOnce(&mut VisitedSet) -> R,
{
THREAD_VISITED.with(|cell| {
let mut visited = cell.borrow_mut();
visited.prepare(num_nodes, capacity_hint);
f(&mut visited)
})
}
#[derive(PartialEq)]
pub(crate) struct MinCandidate {
pub(crate) id: u32,
pub(crate) distance: f32,
}
impl Eq for MinCandidate {}
impl Ord for MinCandidate {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.distance.total_cmp(&self.distance)
}
}
impl PartialOrd for MinCandidate {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
#[derive(PartialEq)]
pub(crate) struct MaxResult {
pub(crate) id: u32,
pub(crate) distance: f32,
}
impl Eq for MaxResult {}
impl Ord for MaxResult {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.distance.total_cmp(&other.distance)
}
}
impl PartialOrd for MaxResult {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
#[inline]
fn insert_result_if_accepted(
results: &mut BinaryHeap<MaxResult>,
ef: usize,
id: u32,
distance: f32,
) -> bool {
if results.len() < ef {
results.push(MaxResult { id, distance });
#[cfg(feature = "benchmark")]
record_result_insertion();
return true;
}
let Some(mut worst) = results.peek_mut() else {
#[cfg(feature = "benchmark")]
record_result_rejection();
return false;
};
if distance >= worst.distance {
#[cfg(feature = "benchmark")]
record_result_rejection();
return false;
}
*worst = MaxResult { id, distance };
#[cfg(feature = "benchmark")]
record_result_replacement();
true
}
#[inline]
fn flush_batch(
query: &[f32],
batch_ids: &[u32; DISTANCE_BATCH_SIZE],
count: usize,
vectors: &[f32],
dimension: usize,
dist_fn: fn(&[f32], &[f32]) -> f32,
candidates: &mut std::collections::BinaryHeap<MinCandidate>,
results: &mut std::collections::BinaryHeap<MaxResult>,
ef: usize,
) {
let mut dists = [0.0f32; DISTANCE_BATCH_SIZE];
for i in 0..count {
let vec = get_vector(vectors, dimension, batch_ids[i] as usize);
dists[i] = dist_fn(query, vec);
}
#[cfg(feature = "benchmark")]
record_distance_evals(count);
if ef < CACHED_WORST_MIN_EF {
for i in 0..count {
if insert_result_if_accepted(results, ef, batch_ids[i], dists[i]) {
candidates.push(MinCandidate {
id: batch_ids[i],
distance: dists[i],
});
#[cfg(feature = "benchmark")]
record_candidate_push(candidates.len());
}
}
return;
}
let mut worst_dist = if results.len() < ef {
f32::INFINITY
} else {
results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY)
};
for i in 0..count {
if results.len() < ef || dists[i] < worst_dist {
if insert_result_if_accepted(results, ef, batch_ids[i], dists[i]) {
candidates.push(MinCandidate {
id: batch_ids[i],
distance: dists[i],
});
#[cfg(feature = "benchmark")]
record_candidate_push(candidates.len());
}
worst_dist = if results.len() < ef {
f32::INFINITY
} else {
results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY)
};
} else {
#[cfg(feature = "benchmark")]
record_result_rejection();
}
}
}
#[inline]
fn flush_batch_custom<F: Fn(&[f32], u32) -> f32>(
query: &[f32],
batch_ids: &[u32; DISTANCE_BATCH_SIZE],
count: usize,
dist_fn: &F,
candidates: &mut std::collections::BinaryHeap<MinCandidate>,
results: &mut std::collections::BinaryHeap<MaxResult>,
ef: usize,
) {
let mut dists = [0.0f32; DISTANCE_BATCH_SIZE];
for i in 0..count {
dists[i] = dist_fn(query, batch_ids[i]);
}
#[cfg(feature = "benchmark")]
record_distance_evals(count);
if ef < CACHED_WORST_MIN_EF {
for i in 0..count {
if insert_result_if_accepted(results, ef, batch_ids[i], dists[i]) {
candidates.push(MinCandidate {
id: batch_ids[i],
distance: dists[i],
});
#[cfg(feature = "benchmark")]
record_candidate_push(candidates.len());
}
}
return;
}
let mut worst_dist = if results.len() < ef {
f32::INFINITY
} else {
results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY)
};
for i in 0..count {
if results.len() < ef || dists[i] < worst_dist {
if insert_result_if_accepted(results, ef, batch_ids[i], dists[i]) {
candidates.push(MinCandidate {
id: batch_ids[i],
distance: dists[i],
});
#[cfg(feature = "benchmark")]
record_candidate_push(candidates.len());
}
worst_dist = if results.len() < ef {
f32::INFINITY
} else {
results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY)
};
} else {
#[cfg(feature = "benchmark")]
record_result_rejection();
}
}
}
#[inline]
fn prune_unpromising_candidates(
candidates: &mut BinaryHeap<MinCandidate>,
results: &BinaryHeap<MaxResult>,
ef: usize,
) {
if ef < FRONTIER_PRUNE_MIN_EF || results.len() < ef || candidates.len() < ef {
return;
}
let Some(worst) = results.peek() else {
return;
};
let worst_dist = worst.distance;
#[cfg(feature = "benchmark")]
let before = candidates.len();
candidates.retain(|candidate| candidate.distance <= worst_dist);
#[cfg(feature = "benchmark")]
record_frontier_retain(before, candidates.len());
}
#[cfg(feature = "hnsw")]
pub fn greedy_search_layer(
query: &[f32],
entry_point: u32,
layer: &crate::hnsw::graph::Layer,
vectors: &[f32],
dimension: usize,
ef: usize,
dist_fn: fn(&[f32], &[f32]) -> f32,
) -> Vec<(u32, f32)> {
let num_vectors = vectors.len() / dimension;
THREAD_SEARCH_SCRATCH.with(|scratch_cell| {
let mut scratch = scratch_cell.borrow_mut();
scratch.prepare(ef);
let SearchScratch {
candidates,
results,
} = &mut *scratch;
with_visited_set(num_vectors, ef * 2, |visited| {
let entry_vector = get_vector(vectors, dimension, entry_point as usize);
let entry_distance = dist_fn(query, entry_vector);
#[cfg(feature = "benchmark")]
record_distance_evals(1);
candidates.push(MinCandidate {
id: entry_point,
distance: entry_distance,
});
#[cfg(feature = "benchmark")]
record_candidate_push(candidates.len());
results.push(MaxResult {
id: entry_point,
distance: entry_distance,
});
#[cfg(feature = "benchmark")]
record_result_insertion();
visited.insert(entry_point);
let should_prune_frontier = ef >= FRONTIER_PRUNE_MIN_EF;
let mut pops_since_prune = 0usize;
while let Some(candidate) = candidates.pop() {
#[cfg(feature = "benchmark")]
record_candidate_pop();
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
if should_prune_frontier {
pops_since_prune += 1;
}
let neighbors = layer.get_neighbors(candidate.id);
let mut batch_ids: [u32; DISTANCE_BATCH_SIZE] = [0; DISTANCE_BATCH_SIZE];
let mut batch_count = 0usize;
for &neighbor_id in neighbors.iter() {
if visited.insert(neighbor_id) {
batch_ids[batch_count] = neighbor_id;
batch_count += 1;
if (neighbor_id as usize) < num_vectors {
let ptr = vectors
.as_ptr()
.wrapping_add(neighbor_id as usize * dimension);
prefetch_read_data(ptr);
if dimension > 16 {
prefetch_read_data(ptr.wrapping_add(16));
}
}
if batch_count == DISTANCE_BATCH_SIZE {
flush_batch(
query,
&batch_ids,
batch_count,
vectors,
dimension,
dist_fn,
candidates,
results,
ef,
);
batch_count = 0;
}
}
}
if batch_count > 0 {
flush_batch(
query,
&batch_ids,
batch_count,
vectors,
dimension,
dist_fn,
candidates,
results,
ef,
);
}
if should_prune_frontier && pops_since_prune >= FRONTIER_PRUNE_INTERVAL {
prune_unpromising_candidates(candidates, results, ef);
pops_since_prune = 0;
}
}
let mut output: Vec<(u32, f32)> = results.drain().map(|r| (r.id, r.distance)).collect();
candidates.clear();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
output
})
})
}
#[cfg(feature = "hnsw")]
pub fn greedy_search_layer_multi_entry(
query: &[f32],
entry_points: &[u32],
layer: &crate::hnsw::graph::Layer,
vectors: &[f32],
dimension: usize,
ef: usize,
dist_fn: fn(&[f32], &[f32]) -> f32,
) -> Vec<(u32, f32)> {
if entry_points.is_empty() {
return Vec::new();
}
let num_vectors = vectors.len() / dimension;
with_visited_set(num_vectors, ef * 2, |visited| {
let mut candidates: BinaryHeap<MinCandidate> = BinaryHeap::with_capacity(ef * 2);
let mut results: BinaryHeap<MaxResult> = BinaryHeap::with_capacity(ef + 1);
for &ep in entry_points {
if visited.insert(ep) {
let ep_vec = get_vector(vectors, dimension, ep as usize);
let ep_dist = dist_fn(query, ep_vec);
candidates.push(MinCandidate {
id: ep,
distance: ep_dist,
});
let _ = insert_result_if_accepted(&mut results, ef, ep, ep_dist);
}
}
while let Some(candidate) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
let neighbors = layer.get_neighbors(candidate.id);
for (i, &neighbor_id) in neighbors.iter().enumerate() {
if i + 1 < neighbors.len() {
let next_id = neighbors[i + 1] as usize;
if next_id < num_vectors {
let ptr = vectors.as_ptr().wrapping_add(next_id * dimension);
prefetch_read_data(ptr);
prefetch_read_data(ptr.wrapping_add(16));
}
}
if i + 4 < neighbors.len() {
let far_id = neighbors[i + 4] as usize;
if far_id < num_vectors {
prefetch_read_data(vectors.as_ptr().wrapping_add(far_id * dimension));
}
}
if visited.insert(neighbor_id) {
let neighbor_vector = get_vector(vectors, dimension, neighbor_id as usize);
let neighbor_distance = dist_fn(query, neighbor_vector);
if insert_result_if_accepted(&mut results, ef, neighbor_id, neighbor_distance) {
candidates.push(MinCandidate {
id: neighbor_id,
distance: neighbor_distance,
});
}
}
}
}
let mut output: Vec<(u32, f32)> = results.into_iter().map(|r| (r.id, r.distance)).collect();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
output
})
}
#[cfg(feature = "hnsw")]
pub fn greedy_search_layer_custom<F: Fn(&[f32], u32) -> f32>(
query: &[f32],
entry_point: u32,
layer: &crate::hnsw::graph::Layer,
vectors: &[f32],
dimension: usize,
ef: usize,
dist_fn: &F,
) -> Vec<(u32, f32)> {
let num_vectors = vectors.len() / dimension;
with_visited_set(num_vectors, ef * 2, |visited| {
let mut candidates: BinaryHeap<MinCandidate> = BinaryHeap::with_capacity(ef * 2);
let mut results: BinaryHeap<MaxResult> = BinaryHeap::with_capacity(ef + 1);
let entry_distance = dist_fn(query, entry_point);
candidates.push(MinCandidate {
id: entry_point,
distance: entry_distance,
});
results.push(MaxResult {
id: entry_point,
distance: entry_distance,
});
visited.insert(entry_point);
while let Some(candidate) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
let neighbors = layer.get_neighbors(candidate.id);
let mut batch_ids: [u32; DISTANCE_BATCH_SIZE] = [0; DISTANCE_BATCH_SIZE];
let mut batch_count = 0usize;
for &neighbor_id in neighbors.iter() {
if visited.insert(neighbor_id) {
batch_ids[batch_count] = neighbor_id;
batch_count += 1;
if (neighbor_id as usize) < num_vectors {
let ptr = vectors
.as_ptr()
.wrapping_add(neighbor_id as usize * dimension);
prefetch_read_data(ptr);
if dimension > 16 {
prefetch_read_data(ptr.wrapping_add(16));
}
}
if batch_count == DISTANCE_BATCH_SIZE {
flush_batch_custom(
query,
&batch_ids,
batch_count,
dist_fn,
&mut candidates,
&mut results,
ef,
);
batch_count = 0;
}
}
}
if batch_count > 0 {
flush_batch_custom(
query,
&batch_ids,
batch_count,
dist_fn,
&mut candidates,
&mut results,
ef,
);
}
}
let mut output: Vec<(u32, f32)> = results.into_iter().map(|r| (r.id, r.distance)).collect();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
output
})
}
#[cfg(feature = "ivf_rabitq")]
pub fn greedy_search_layer_edge_aware<F: Fn(u32, u32, usize) -> f32>(
entry_point: u32,
entry_dist: f32,
layer: &crate::hnsw::graph::Layer,
num_vectors: usize,
ef: usize,
dist_fn: &F,
) -> Vec<(u32, f32)> {
with_visited_set(num_vectors, ef * 2, |visited| {
let mut candidates: BinaryHeap<MinCandidate> = BinaryHeap::with_capacity(ef * 2);
let mut results: BinaryHeap<MaxResult> = BinaryHeap::with_capacity(ef + 1);
candidates.push(MinCandidate {
id: entry_point,
distance: entry_dist,
});
results.push(MaxResult {
id: entry_point,
distance: entry_dist,
});
visited.insert(entry_point);
while let Some(candidate) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
let neighbors = layer.get_neighbors(candidate.id);
for (slot, &neighbor_id) in neighbors.iter().enumerate() {
if visited.insert(neighbor_id) {
let neighbor_distance = dist_fn(candidate.id, neighbor_id, slot);
if insert_result_if_accepted(&mut results, ef, neighbor_id, neighbor_distance) {
candidates.push(MinCandidate {
id: neighbor_id,
distance: neighbor_distance,
});
}
}
}
}
let mut output: Vec<(u32, f32)> = results.into_iter().map(|r| (r.id, r.distance)).collect();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
output
})
}
#[cfg(feature = "hnsw")]
pub fn greedy_search_layer_adaptive(
query: &[f32],
entry_point: u32,
layer: &crate::hnsw::graph::Layer,
vectors: &[f32],
dimension: usize,
ef: usize,
k: usize,
config: &crate::adaptive::AdaptiveConfig,
dist_fn: fn(&[f32], &[f32]) -> f32,
) -> (Vec<(u32, f32)>, usize) {
use crate::adaptive::EarlyTerminationOracle;
let num_vectors = vectors.len() / dimension;
with_visited_set(num_vectors, ef * 2, |visited| {
let mut candidates: BinaryHeap<MinCandidate> = BinaryHeap::with_capacity(ef * 2);
let mut results: BinaryHeap<MaxResult> = BinaryHeap::with_capacity(ef + 1);
let mut oracle = EarlyTerminationOracle::new(k, config.clone());
let entry_vector = get_vector(vectors, dimension, entry_point as usize);
let entry_distance = dist_fn(query, entry_vector);
oracle.observe(entry_distance);
candidates.push(MinCandidate {
id: entry_point,
distance: entry_distance,
});
results.push(MaxResult {
id: entry_point,
distance: entry_distance,
});
visited.insert(entry_point);
while let Some(candidate) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
let neighbors = layer.get_neighbors(candidate.id);
for (i, &neighbor_id) in neighbors.iter().enumerate() {
if i + 1 < neighbors.len() {
let next_id = neighbors[i + 1] as usize;
if next_id < num_vectors {
let ptr = vectors.as_ptr().wrapping_add(next_id * dimension);
prefetch_read_data(ptr);
prefetch_read_data(ptr.wrapping_add(16));
}
}
if i + 4 < neighbors.len() {
let far_id = neighbors[i + 4] as usize;
if far_id < num_vectors {
prefetch_read_data(vectors.as_ptr().wrapping_add(far_id * dimension));
}
}
if visited.insert(neighbor_id) {
let neighbor_vector = get_vector(vectors, dimension, neighbor_id as usize);
let neighbor_distance = dist_fn(query, neighbor_vector);
oracle.observe(neighbor_distance);
if insert_result_if_accepted(&mut results, ef, neighbor_id, neighbor_distance) {
candidates.push(MinCandidate {
id: neighbor_id,
distance: neighbor_distance,
});
}
}
}
if oracle.should_terminate() && results.len() >= k {
break;
}
}
let num_evaluated = oracle.num_evaluated();
let mut output: Vec<(u32, f32)> = results.into_iter().map(|r| (r.id, r.distance)).collect();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
(output, num_evaluated)
})
}
#[inline]
fn get_vector(vectors: &[f32], dimension: usize, idx: usize) -> &[f32] {
let start = idx * dimension;
let end = start + dimension;
&vectors[start..end]
}
#[cfg(feature = "hnsw")]
pub fn greedy_search_layer_prt(
query: &[f32],
entry_point: u32,
layer: &crate::hnsw::graph::Layer,
vectors: &[f32],
dimension: usize,
ef: usize,
dist_fn: fn(&[f32], &[f32]) -> f32,
prt: &crate::prt::ProbabilisticRoutingTest,
query_proj: &[f32],
tfb: &mut crate::prt::TestFeedbackBuffer,
) -> (Vec<(u32, f32)>, usize) {
let num_vectors = vectors.len() / dimension;
let mut full_dist_count: usize = 0;
let results = with_visited_set(num_vectors, ef * 2, |visited| {
let mut candidates: BinaryHeap<MinCandidate> = BinaryHeap::with_capacity(ef * 2);
let mut results: BinaryHeap<MaxResult> = BinaryHeap::with_capacity(ef + 1);
let entry_vector = get_vector(vectors, dimension, entry_point as usize);
let entry_distance = dist_fn(query, entry_vector);
full_dist_count += 1;
candidates.push(MinCandidate {
id: entry_point,
distance: entry_distance,
});
results.push(MaxResult {
id: entry_point,
distance: entry_distance,
});
visited.insert(entry_point);
while let Some(candidate) = candidates.pop() {
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if candidate.distance > worst_dist && results.len() >= ef {
break;
}
let neighbors = layer.get_neighbors(candidate.id);
for (i, &neighbor_id) in neighbors.iter().enumerate() {
if i + 1 < neighbors.len() {
let next_id = neighbors[i + 1] as usize;
if next_id < num_vectors {
let ptr = vectors.as_ptr().wrapping_add(next_id * dimension);
prefetch_read_data(ptr);
}
}
if !visited.insert(neighbor_id) {
continue;
}
let worst_dist = results.peek().map(|r| r.distance).unwrap_or(f32::INFINITY);
if results.len() >= ef
&& !prt.should_compute_full_distance(query_proj, neighbor_id, worst_dist, tfb)
{
continue;
}
let neighbor_vector = get_vector(vectors, dimension, neighbor_id as usize);
let neighbor_distance = dist_fn(query, neighbor_vector);
full_dist_count += 1;
if insert_result_if_accepted(&mut results, ef, neighbor_id, neighbor_distance) {
tfb.record_true_positive();
candidates.push(MinCandidate {
id: neighbor_id,
distance: neighbor_distance,
});
} else {
tfb.record_false_positive();
}
}
}
let mut output: Vec<(u32, f32)> = results.into_iter().map(|r| (r.id, r.distance)).collect();
output.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
output
});
(results, full_dist_count)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn test_candidate_ordering() {
let mut heap = BinaryHeap::new();
heap.push(MinCandidate {
id: 0,
distance: 0.5,
});
heap.push(MinCandidate {
id: 1,
distance: 0.1,
});
heap.push(MinCandidate {
id: 2,
distance: 0.3,
});
assert_eq!(heap.pop().unwrap().distance, 0.1);
assert_eq!(heap.pop().unwrap().distance, 0.3);
assert_eq!(heap.pop().unwrap().distance, 0.5);
}
#[test]
fn test_visited_set_dense() {
let mut v = VisitedSet::new(100, 10);
assert!(!v.contains(5));
assert!(v.insert(5));
assert!(v.contains(5));
assert!(!v.insert(5)); }
#[test]
fn test_visited_set_dense_clear() {
let mut v = VisitedSet::new(100, 10);
assert!(v.insert(5));
assert!(v.contains(5));
v.clear();
assert!(!v.contains(5));
assert!(v.insert(5));
}
#[test]
fn test_visited_set_dense_generation_overflow() {
let mut v = VisitedSet::new(100, 10);
if let VisitedSet::Dense {
ref mut generation, ..
} = v
{
*generation = u8::MAX;
}
assert!(v.insert(5));
assert!(v.contains(5));
v.clear();
assert!(!v.contains(5));
assert!(v.insert(5));
assert!(v.contains(5));
}
#[test]
fn test_visited_set_sparse() {
let mut v = VisitedSet::new(DENSE_VISITED_THRESHOLD + 1, 10);
assert!(!v.contains(42));
assert!(v.insert(42));
assert!(v.contains(42));
assert!(!v.insert(42));
}
#[test]
fn test_visited_set_sparse_clear() {
let mut v = VisitedSet::new(DENSE_VISITED_THRESHOLD + 1, 10);
assert!(v.insert(42));
v.clear();
assert!(!v.contains(42));
assert!(v.insert(42));
}
}