use super::super::distance::{batch_distance_with_prefetch, DistanceEngine};
use super::super::layer::{Layer, NodeId};
use super::super::ordered_float::OrderedFloat;
use super::search_pools::should_prefetch;
use super::search_state::{gather_unvisited_neighbors, process_batch_results, SearchState};
use super::{NativeHnsw, NO_ENTRY_POINT};
use crate::perf_optimizations::ContiguousVectors;
use smallvec::SmallVec;
use std::borrow::Cow;
use std::cell::{Cell, RefCell};
use std::cmp::Reverse;
use std::sync::atomic::{AtomicU64, Ordering};
thread_local! {
static QUERY_BUF: RefCell<Vec<f32>> = RefCell::new(Vec::with_capacity(1536));
}
static PROBE_RNG_SEED_COUNTER: AtomicU64 = AtomicU64::new(0x5DEE_CE66_D1A4_B5B5);
thread_local! {
static PROBE_RNG: Cell<u64> = const { Cell::new(0) };
}
pub struct ResumableSearch {
state: SearchState,
}
impl<D: DistanceEngine> NativeHnsw<D> {
#[inline]
#[must_use]
pub fn search(&self, query: &[f32], k: usize, ef_search: usize) -> Vec<(NodeId, f32)> {
let prepared_query = self.prepare_query(query);
let results = self.search_prepared(&prepared_query, k, ef_search);
Self::recycle_cow(prepared_query);
results
}
#[inline]
fn search_prepared(&self, query: &[f32], k: usize, ef_search: usize) -> Vec<(NodeId, f32)> {
let ep = self.entry_point.load(Ordering::Acquire);
if ep == NO_ENTRY_POINT {
return Vec::new();
}
let max_layer = self.max_layer.load(Ordering::Relaxed);
let mut current_ep = ep;
for layer_idx in (1..=max_layer).rev() {
current_ep = self.search_layer_single(query, current_ep, layer_idx);
}
let count = self.count.load(Ordering::Relaxed);
let probes = self.adaptive_num_probes(count, ef_search, k);
if probes > 1 {
self.search_multi_entry_prepared(query, k, ef_search, probes)
} else {
self.search_layer(
query,
&[current_ep],
ef_search,
0,
self.stagnation_limit,
Some(k),
)
}
}
#[inline]
#[allow(clippy::unused_self)] fn adaptive_num_probes(&self, count: usize, ef_search: usize, k: usize) -> usize {
if count <= 10_000 || ef_search <= (k * 4).max(64) {
return 1;
}
if ef_search >= 1024 {
4
} else if ef_search >= 512 {
3
} else {
2
}
}
#[must_use]
pub fn search_multi_entry(
&self,
query: &[f32],
k: usize,
ef_search: usize,
num_probes: usize,
) -> Vec<(NodeId, f32)> {
let prepared_query = self.prepare_query(query);
let result = self.search_multi_entry_prepared(&prepared_query, k, ef_search, num_probes);
Self::recycle_cow(prepared_query);
result
}
#[must_use]
fn search_multi_entry_prepared(
&self,
query: &[f32],
k: usize,
ef_search: usize,
num_probes: usize,
) -> Vec<(NodeId, f32)> {
let ep = self.entry_point.load(Ordering::Acquire);
if ep == NO_ENTRY_POINT {
return Vec::new();
}
let count = self.count.load(Ordering::Relaxed);
if count == 0 {
return Vec::new();
}
let max_layer = self.max_layer.load(Ordering::Relaxed);
let mut current_ep = ep;
for layer_idx in (1..=max_layer).rev() {
current_ep = self.search_layer_single(query, current_ep, layer_idx);
}
let entry_points = Self::gather_multi_entry_points(current_ep, count, num_probes);
self.search_layer(
query,
&entry_points,
ef_search,
0,
self.stagnation_limit,
Some(k),
)
}
#[inline]
fn gather_multi_entry_points(
primary_ep: NodeId,
count: usize,
num_probes: usize,
) -> Vec<NodeId> {
let mut entry_points = vec![primary_ep];
if num_probes > 1 && count > 10 {
for _ in 1..num_probes.min(4) {
let random_id = (Self::next_probe_rng() as usize) % count;
if !entry_points.contains(&random_id) {
entry_points.push(random_id);
}
}
}
entry_points
}
#[inline]
fn next_probe_rng() -> u64 {
PROBE_RNG.with(|cell| {
let mut s = cell.get();
if s == 0 {
s = PROBE_RNG_SEED_COUNTER.fetch_add(0x9e37_79b9_7f4a_7c15, Ordering::Relaxed);
if s == 0 {
s = 1; }
}
let next = super::xorshift64(s);
cell.set(next);
next
})
}
#[inline]
fn recycle_cow(cow: Cow<'_, [f32]>) {
if let Cow::Owned(buf) = cow {
Self::return_query_buf(buf);
}
}
#[must_use]
pub(in crate::index::hnsw) fn search_resumable(
&self,
query: &[f32],
k: usize,
ef_search: usize,
) -> (Vec<(NodeId, f32)>, Option<ResumableSearch>) {
let prepared_query = self.prepare_query(query);
let outcome = self.search_resumable_prepared(&prepared_query, k, ef_search);
Self::recycle_cow(prepared_query);
outcome
}
fn search_resumable_prepared(
&self,
query: &[f32],
k: usize,
ef_search: usize,
) -> (Vec<(NodeId, f32)>, Option<ResumableSearch>) {
let ep = self.entry_point.load(Ordering::Acquire);
if ep == NO_ENTRY_POINT {
return (Vec::new(), None);
}
let max_layer = self.max_layer.load(Ordering::Relaxed);
let mut current_ep = ep;
for layer_idx in (1..=max_layer).rev() {
current_ep = self.search_layer_single(query, current_ep, layer_idx);
}
let count = self.count.load(Ordering::Relaxed);
let probes = self.adaptive_num_probes(count, ef_search, k);
let mut state = SearchState::new(count);
if probes > 1 {
let entry_points = Self::gather_multi_entry_points(current_ep, count, probes);
self.run_layer_search(
query,
&entry_points,
ef_search,
0,
self.stagnation_limit,
&mut state,
);
} else {
self.run_layer_search(
query,
&[current_ep],
ef_search,
0,
self.stagnation_limit,
&mut state,
);
}
let results = state.peek_sorted_results(Some(k));
(results, Some(ResumableSearch { state }))
}
#[must_use]
pub(in crate::index::hnsw) fn resume_search(
&self,
resume: ResumableSearch,
query: &[f32],
k: usize,
ef_search: usize,
) -> Vec<(NodeId, f32)> {
let ResumableSearch { mut state } = resume;
state.stagnation_count = 0;
let prepared_query = self.prepare_query(query);
self.continue_layer_search(
&prepared_query,
&mut state,
ef_search,
0,
self.stagnation_limit,
);
Self::recycle_cow(prepared_query);
state.into_sorted_results(Some(k))
}
#[inline]
pub(in crate::index::hnsw::native::graph) fn search_layer_single(
&self,
query: &[f32],
entry: NodeId,
layer: usize,
) -> NodeId {
self.with_vectors_and_layers_read(|vectors, layers| {
let dimension = vectors.dimension();
let prefetch_dist = crate::simd_native::calculate_prefetch_distance(dimension);
let mut best = entry;
debug_assert!(
entry < vectors.len(),
"entry {entry} out of bounds (len {})",
vectors.len()
);
let entry_vec = unsafe { vectors.get_unchecked(entry) };
let mut best_dist = self.distance.distance(query, entry_vec);
loop {
let improved = layers[layer]
.with_neighbors(best, |neighbors| {
self.greedy_scan_with_prefetch(
query,
neighbors,
vectors,
dimension,
prefetch_dist,
&mut best,
&mut best_dist,
)
})
.unwrap_or(false);
if !improved {
break;
}
}
best
})
}
#[inline]
fn prefetch_neighbors(
neighbors: &[NodeId],
vectors: &crate::perf_optimizations::ContiguousVectors,
start: usize,
count: usize,
) {
for &neighbor_id in neighbors.iter().skip(start).take(count) {
if neighbor_id < vectors.len() {
vectors.prefetch(neighbor_id);
}
}
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn greedy_scan_with_prefetch(
&self,
query: &[f32],
neighbors: &[NodeId],
vectors: &crate::perf_optimizations::ContiguousVectors,
dimension: usize,
prefetch_dist: usize,
best: &mut NodeId,
best_dist: &mut f32,
) -> bool {
let use_prefetch = should_prefetch(dimension);
if use_prefetch && neighbors.len() > prefetch_dist {
Self::prefetch_neighbors(neighbors, vectors, 0, prefetch_dist);
}
let mut improved = false;
for (i, &neighbor) in neighbors.iter().enumerate() {
if use_prefetch && i + prefetch_dist < neighbors.len() {
Self::prefetch_neighbors(neighbors, vectors, i + prefetch_dist, 1);
}
debug_assert!(
neighbor < vectors.len(),
"neighbor {neighbor} out of bounds (len {})",
vectors.len()
);
let neighbor_vec = unsafe { vectors.get_unchecked(neighbor) };
let dist = self.distance.distance(query, neighbor_vec);
if dist < *best_dist {
*best = neighbor;
*best_dist = dist;
improved = true;
}
}
improved
}
#[inline]
pub(in crate::index::hnsw::native::graph) fn search_layer(
&self,
query: &[f32],
entry_points: &[NodeId],
ef: usize,
layer: usize,
stagnation_limit: usize,
result_limit: Option<usize>,
) -> Vec<(NodeId, f32)> {
let capacity_hint = self.count.load(Ordering::Relaxed);
let mut state = SearchState::new(capacity_hint);
self.run_layer_search(query, entry_points, ef, layer, stagnation_limit, &mut state);
state.into_sorted_results(result_limit)
}
fn run_layer_search(
&self,
query: &[f32],
entry_points: &[NodeId],
ef: usize,
layer: usize,
stagnation_limit: usize,
state: &mut SearchState,
) {
self.with_vectors_and_layers_read(|vectors, layers| {
let use_prefetch = should_prefetch(vectors.dimension());
for &ep in entry_points {
debug_assert!(
ep < vectors.len(),
"ep {ep} out of bounds (len {})",
vectors.len()
);
let ep_vec = unsafe { vectors.get_unchecked(ep) };
let dist = self.distance.distance(query, ep_vec);
state.push_candidate(ep, dist);
}
Self::dispatch_layer_search(
&self.distance,
query,
vectors,
layers,
state,
ef,
layer,
stagnation_limit,
use_prefetch,
);
});
}
fn continue_layer_search(
&self,
query: &[f32],
state: &mut SearchState,
ef: usize,
layer: usize,
stagnation_limit: usize,
) {
self.with_vectors_and_layers_read(|vectors, layers| {
let use_prefetch = should_prefetch(vectors.dimension());
Self::dispatch_layer_search(
&self.distance,
query,
vectors,
layers,
state,
ef,
layer,
stagnation_limit,
use_prefetch,
);
});
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn dispatch_layer_search(
distance: &D,
query: &[f32],
vectors: &ContiguousVectors,
layers: &[Layer],
state: &mut SearchState,
ef: usize,
layer: usize,
stagnation_limit: usize,
use_prefetch: bool,
) {
let use_pipeline = use_prefetch && vectors.len() >= 10_000;
if use_pipeline {
super::search_pipeline::search_layer_pipelined(
distance,
query,
vectors,
layers,
state,
ef,
layer,
stagnation_limit,
use_prefetch,
);
} else {
Self::search_loop_sequential(
distance,
query,
vectors,
layers,
state,
ef,
layer,
stagnation_limit,
use_prefetch,
);
}
}
#[inline]
#[allow(clippy::too_many_arguments)]
fn search_loop_sequential(
distance: &D,
query: &[f32],
vectors: &ContiguousVectors,
layers: &[Layer],
state: &mut SearchState,
ef: usize,
layer: usize,
stagnation_limit: usize,
use_prefetch: bool,
) {
while let Some(Reverse((OrderedFloat(c_dist), c_node))) = state.candidates.pop() {
if state.should_terminate(c_dist, ef, stagnation_limit) {
break;
}
let improved = layers[layer]
.with_neighbors(c_node, |neighbors| {
let batch = gather_unvisited_neighbors(
neighbors,
&mut state.visited,
vectors,
use_prefetch,
);
if batch.is_empty() {
return false;
}
let vecs: SmallVec<[&[f32]; 64]> = batch.iter().map(|(_, v)| *v).collect();
let distances = batch_distance_with_prefetch(distance, query, &vecs);
process_batch_results(&batch, &distances, ef, state)
})
.unwrap_or(false);
state.update_stagnation(improved);
}
}
#[inline]
pub(in crate::index::hnsw::native) fn prepare_query<'a>(
&self,
query: &'a [f32],
) -> Cow<'a, [f32]> {
if self.distance.is_pre_normalized()
&& self.distance.metric() == crate::DistanceMetric::Cosine
{
let mut buf = QUERY_BUF.with(|cell| {
let mut borrow = cell.borrow_mut();
if borrow.capacity() == 0 {
Vec::with_capacity(query.len())
} else {
std::mem::take(&mut *borrow)
}
});
buf.clear();
buf.extend_from_slice(query);
crate::simd_native::normalize_inplace_native(&mut buf);
Cow::Owned(buf)
} else {
Cow::Borrowed(query)
}
}
#[inline]
pub(in crate::index::hnsw::native) fn with_prepared_query<R>(
&self,
query: &[f32],
f: impl FnOnce(&[f32]) -> R,
) -> R {
let prepared = self.prepare_query(query);
let result = f(&prepared);
Self::recycle_cow(prepared);
result
}
#[inline]
fn return_query_buf(buf: Vec<f32>) {
QUERY_BUF.with(|cell| {
let mut borrow = cell.borrow_mut();
if borrow.is_empty() {
*borrow = buf;
}
});
}
}
#[cfg(test)]
#[path = "probe_tests.rs"]
mod probe_tests;