leaves12138 commented on code in PR #62: URL: https://github.com/apache/paimon-vector-index/pull/62#discussion_r3651305455
########## core/src/diskann_search.rs: ########## @@ -0,0 +1,6979 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use crate::diskann::DiskAnnRawVectorEncoding; +use crate::diskann_io::{ + decode_adjacency_list, CacheLockMetrics, OffsetLru, SectionRange, SharedWindowCacheLookup, + DISKANN_PAGE_SIZE, +}; +use crate::distance::{ + fvec_distance, fvec_l2sqr, pq_distance_four_codes, pq_distance_from_table, preprocess_vectors, + MetricType, +}; +use crate::index_io_util::decode_roaring_filter; +use crate::io::{ReadRequest, SeekRead}; +use crate::read_options::ReadPlan; +use crate::sparse_table::{estimated_memory_bytes as sparse_table_memory_bytes, SparseTable}; +use half::prelude::{HalfBitsSliceExt, HalfFloatSliceExt}; +use rayon::prelude::*; +use roaring::{RoaringBitmap, RoaringTreemap}; +use std::borrow::Cow; +use std::cmp::{Ordering, Reverse}; +use std::collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet}; +use std::io; +use std::sync::Arc; + +const QUERY_WINDOW_BUFFER_LIMIT_BYTES: usize = 8 * 1024 * 1024; +const QUERY_ADJACENCY_WINDOW_LIMIT_BYTES: usize = 8 * 1024 * 1024; +const BATCH_WINDOW_BUFFER_LIMIT_BYTES: usize = 64 * 1024 * 1024; +const FILTERED_BATCH_RERANK_MAX_BYTES: usize = 64 * 1024 * 1024; +const FILTERED_BATCH_RERANK_MAX_RANGES: usize = 1024; +const BATCH_QUERY_CHUNK_SIZE: usize = 1024; +const FILTERED_PQ_MAX_QUERY_TILE_SIZE: usize = 4; +const FILTERED_PQ_TILE_TABLE_LIMIT_BYTES: usize = 2 * 1024 * 1024; +const FILTERED_SINGLE_PQ_NODE_CHUNK_SIZE: usize = 1024; +const PARALLEL_EXACT_RERANK_MIN_COMPONENTS: usize = 16 * 1024; +const PARALLEL_SESSION_MAX_QUERIES_PER_WORKER: usize = 4; +const SPARSE_VISITED_MIN_MEMORY_SAVINGS: usize = 2; + +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub struct DiskAnnSearchStats { + pub query_count: usize, + pub query_chunks: usize, + pub max_queries_per_chunk: usize, + pub filtered_exhaustive_queries: usize, + pub filtered_graph_queries: usize, + pub filtered_graph_fallbacks: usize, + pub pq_distance_evaluations: usize, + pub pq_code_loads: usize, + pub adjacency_cache_hits: usize, + pub adjacency_cache_misses: usize, + pub adjacency_cache_waits: usize, + pub adjacency_cache_evictions: usize, + pub adjacency_cache_lock_acquisitions: usize, + pub adjacency_cache_lock_wait_nanos: u64, + pub query_adjacency_cache_peak_bytes: usize, + pub query_adjacency_cache_evictions: usize, + pub rerank_candidate_references: usize, + pub rerank_unique_windows: usize, + pub rerank_chunks: usize, + pub raw_vector_cache_hits: usize, + pub raw_vector_cache_misses: usize, + pub raw_vector_cache_evictions: usize, + pub parallel_exact_rerank_chunks: usize, + pub parallel_exact_rerank_references: usize, + pub parallel_session_queries: usize, +} + +impl DiskAnnSearchStats { + fn record_adjacency_cache_lock(&mut self, metrics: CacheLockMetrics) { + self.adjacency_cache_lock_acquisitions = self + .adjacency_cache_lock_acquisitions + .saturating_add(metrics.acquisitions); + self.adjacency_cache_lock_wait_nanos = self + .adjacency_cache_lock_wait_nanos + .saturating_add(metrics.wait_nanos); + } + + fn merge_candidate_generation(&mut self, worker: Self) { + self.filtered_exhaustive_queries = self + .filtered_exhaustive_queries + .saturating_add(worker.filtered_exhaustive_queries); + self.filtered_graph_queries = self + .filtered_graph_queries + .saturating_add(worker.filtered_graph_queries); + self.filtered_graph_fallbacks = self + .filtered_graph_fallbacks + .saturating_add(worker.filtered_graph_fallbacks); + self.pq_distance_evaluations = self + .pq_distance_evaluations + .saturating_add(worker.pq_distance_evaluations); + self.pq_code_loads = self.pq_code_loads.saturating_add(worker.pq_code_loads); + self.adjacency_cache_hits = self + .adjacency_cache_hits + .saturating_add(worker.adjacency_cache_hits); + self.adjacency_cache_misses = self + .adjacency_cache_misses + .saturating_add(worker.adjacency_cache_misses); + self.adjacency_cache_waits = self + .adjacency_cache_waits + .saturating_add(worker.adjacency_cache_waits); + self.adjacency_cache_evictions = self + .adjacency_cache_evictions + .saturating_add(worker.adjacency_cache_evictions); + self.adjacency_cache_lock_acquisitions = self + .adjacency_cache_lock_acquisitions + .saturating_add(worker.adjacency_cache_lock_acquisitions); + self.adjacency_cache_lock_wait_nanos = self + .adjacency_cache_lock_wait_nanos + .saturating_add(worker.adjacency_cache_lock_wait_nanos); + self.query_adjacency_cache_peak_bytes = self + .query_adjacency_cache_peak_bytes + .max(worker.query_adjacency_cache_peak_bytes); + self.query_adjacency_cache_evictions = self + .query_adjacency_cache_evictions + .saturating_add(worker.query_adjacency_cache_evictions); + } + + fn merge_complete_query(&mut self, worker: Self) { + let worker_query_count = worker.query_count; + self.merge_candidate_generation(worker); + self.rerank_candidate_references = self + .rerank_candidate_references + .saturating_add(worker.rerank_candidate_references); + self.rerank_unique_windows = self + .rerank_unique_windows + .saturating_add(worker.rerank_unique_windows); + self.rerank_chunks = self.rerank_chunks.saturating_add(worker.rerank_chunks); + self.raw_vector_cache_hits = self + .raw_vector_cache_hits + .saturating_add(worker.raw_vector_cache_hits); + self.raw_vector_cache_misses = self + .raw_vector_cache_misses + .saturating_add(worker.raw_vector_cache_misses); + self.raw_vector_cache_evictions = self + .raw_vector_cache_evictions + .saturating_add(worker.raw_vector_cache_evictions); + self.parallel_exact_rerank_chunks = self + .parallel_exact_rerank_chunks + .saturating_add(worker.parallel_exact_rerank_chunks); + self.parallel_exact_rerank_references = self + .parallel_exact_rerank_references + .saturating_add(worker.parallel_exact_rerank_references); + self.parallel_session_queries = self + .parallel_session_queries + .saturating_add(worker_query_count.max(1)); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct ReadWindow { + pub offset: u64, + pub length: usize, +} + +impl ReadWindow { + pub const fn new(offset: u64, length: usize) -> Self { + Self { offset, length } + } +} + +pub(crate) struct ReadWindowPlanner { + plan: ReadPlan, + section: SectionRange, +} + +impl ReadWindowPlanner { + pub const fn new(plan: ReadPlan, section: SectionRange) -> Self { + Self { plan, section } + } + + #[cfg(test)] + pub const fn beam_width(&self) -> usize { + self.plan.graph_beam_width + } + + pub fn plan_logical_pages( + &self, + logical_pages: impl IntoIterator<Item = usize>, + ) -> Vec<ReadWindow> { + let mut windows = BTreeMap::new(); + for logical_page in logical_pages { + if let Some(window) = self.window_for_logical_page(logical_page) { + windows.insert(window.offset, window.length); + } + } + windows + .into_iter() + .map(|(offset, length)| ReadWindow::new(offset, length)) + .collect() + } + + fn window_for_logical_page(&self, logical_page: usize) -> Option<ReadWindow> { + let window_size = self.plan.window_bytes as u64; + let relative_page = (logical_page as u64).checked_mul(DISKANN_PAGE_SIZE as u64)?; + if relative_page >= self.section.length { + return None; + } + let relative_window = relative_page / window_size * window_size; + let length = window_size.min(self.section.length - relative_window) as usize; + Some(ReadWindow::new( + self.section.offset + relative_window, + length, + )) + } +} + +pub(crate) struct VectorWindowPlanner { + section: SectionRange, + record_size: usize, + records_per_window: usize, +} + +impl VectorWindowPlanner { + fn new(plan: ReadPlan, section: SectionRange, record_size: usize) -> io::Result<Self> { + if record_size == 0 { + return Err(invalid_data( + "DiskANN raw-vector record size must be greater than zero", + )); + } + Ok(Self { + section, + record_size, + records_per_window: (plan.window_bytes / record_size).max(1), + }) + } + + fn window_for_node(&self, node: usize) -> Option<ReadWindow> { + let record_offset = node.checked_mul(self.record_size)?; + if u64::try_from(record_offset).ok()? >= self.section.length { + return None; + } + let first_node = node / self.records_per_window * self.records_per_window; + let relative_offset = first_node.checked_mul(self.record_size)?; + let maximum_length = self.records_per_window.checked_mul(self.record_size)?; + let remaining = usize::try_from(self.section.length - relative_offset as u64).ok()?; + Some(ReadWindow::new( + self.section.offset.checked_add(relative_offset as u64)?, + maximum_length.min(remaining), + )) + } + + fn plan_nodes(&self, nodes: impl IntoIterator<Item = usize>) -> Vec<ReadWindow> { + let mut windows = BTreeMap::new(); + for node in nodes { + if let Some(window) = self.window_for_node(node) { + windows.insert(window.offset, window.length); + } + } + windows + .into_iter() + .map(|(offset, length)| ReadWindow::new(offset, length)) + .collect() + } + + fn record<'a>( + &self, + window: ReadWindow, + payload: &'a [u8], + node: usize, + ) -> io::Result<&'a [u8]> { + let absolute_offset = self + .section + .offset + .checked_add( + u64::try_from( + node.checked_mul(self.record_size) + .ok_or_else(|| invalid_data("DiskANN raw-vector offset overflows"))?, + ) + .map_err(|_| invalid_data("DiskANN raw-vector offset exceeds u64"))?, + ) + .ok_or_else(|| invalid_data("DiskANN raw-vector offset overflows"))?; + let relative_offset = absolute_offset + .checked_sub(window.offset) + .and_then(|offset| usize::try_from(offset).ok()) + .ok_or_else(|| invalid_data("DiskANN raw-vector window starts after its record"))?; + let record_end = relative_offset + .checked_add(self.record_size) + .ok_or_else(|| invalid_data("DiskANN raw-vector record range overflows"))?; + payload + .get(relative_offset..record_end) + .ok_or_else(|| invalid_data("DiskANN raw-vector record is truncated")) + } +} + +#[derive(Debug, Clone, Copy)] +struct SearchCandidate { + node: usize, + distance: f32, +} + +#[derive(Debug, Clone, Copy)] +struct ExactSearchResult { + row_id: i64, + distance: f32, +} + +#[derive(Clone, Copy)] +struct ExactRerankReference<'a> { + query_index: usize, + row_id: i64, + record: &'a [u8], +} + +enum WindowPayload { + Owned(Vec<u8>), + Shared(Arc<Vec<u8>>), +} + +impl WindowPayload { + fn as_slice(&self) -> &[u8] { + match self { + Self::Owned(payload) => payload, + Self::Shared(payload) => payload, + } + } + + fn capacity(&self) -> usize { + match self { + Self::Owned(payload) => payload.capacity(), + Self::Shared(payload) => payload.capacity(), + } + } +} + +impl From<Vec<u8>> for WindowPayload { + fn from(payload: Vec<u8>) -> Self { + Self::Owned(payload) + } +} + +#[derive(Default)] +struct AdjacencyWindowCache { + entries: HashMap<u64, WindowPayload>, + recency: OffsetLru, + retained_capacity: usize, +} + +impl AdjacencyWindowCache { + fn contains_key(&self, offset: &u64) -> bool { + self.entries.contains_key(offset) + } + + fn get(&self, offset: &u64) -> Option<&WindowPayload> { + self.entries.get(offset) + } + + fn insert(&mut self, offset: u64, payload: WindowPayload) { + let payload_capacity = payload.capacity(); + if let Some(previous) = self.entries.insert(offset, payload) { + self.retained_capacity = self.retained_capacity.saturating_sub(previous.capacity()); + self.recency.remove(offset); + } + self.retained_capacity = self.retained_capacity.saturating_add(payload_capacity); + self.recency.touch(offset); + } + + fn touch_windows(&mut self, windows: &[ReadWindow]) { + for window in windows { + if self.entries.contains_key(&window.offset) { + self.recency.touch(window.offset); + } + } + } + + fn trim(&mut self, window_buffers: &mut WindowBufferPool, capacity_limit: usize) -> usize { + let mut evictions = 0usize; + while self.retained_capacity > capacity_limit { + let Some(offset) = self.recency.pop_oldest() else { + break; + }; + if let Some(payload) = self.entries.remove(&offset) { + self.retained_capacity = self.retained_capacity.saturating_sub(payload.capacity()); + if let WindowPayload::Owned(payload) = payload { + window_buffers.recycle(payload); + } + evictions = evictions.saturating_add(1); + } + } + evictions + } + + fn recycle(&mut self, window_buffers: &mut WindowBufferPool) { + for (_, payload) in self.entries.drain() { + if let WindowPayload::Owned(payload) = payload { + window_buffers.recycle(payload); + } + } + self.recency.clear(); + self.retained_capacity = 0; + } + + #[cfg(test)] + fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + fn retained_capacity(&self) -> usize { + self.retained_capacity + } + + #[cfg(test)] + fn reserve(&mut self, additional: usize) { + self.entries.reserve(additional); + } + + #[cfg(test)] + fn capacity(&self) -> usize { + self.entries.capacity() + } +} + +fn prepare_adjacency_window_cache( + required_windows: &[ReadWindow], + incoming_bytes: usize, + cache: &mut AdjacencyWindowCache, + window_buffers: &mut WindowBufferPool, +) -> usize { + // A window that was available when the read plan was assembled may still be + // needed by this round even when its individual page does not need loading. + // Mark all round inputs as most-recent before making room for missing + // windows, otherwise trimming for a different page can evict one that + // decode is about to consume. + cache.touch_windows(required_windows); + cache.trim( + window_buffers, + QUERY_ADJACENCY_WINDOW_LIMIT_BYTES.saturating_sub(incoming_bytes), + ) +} + +fn share_window_payload(payload: Vec<u8>) -> Arc<Vec<u8>> { + Arc::new(payload) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum FilteredCandidateStrategy { + Exhaustive { + target_candidates: usize, + }, + Graph { + target_candidates: usize, + search_list_size: usize, + }, +} + +#[derive(Default)] +pub(crate) struct DiskAnnQueryScratch { + visited: Vec<bool>, + sparse_visited: SparseTable<()>, + uses_sparse_visited: bool, + touched_nodes: Vec<usize>, + distance_table: Vec<f32>, + candidates: Vec<SearchCandidate>, + rerank_candidates: Vec<SearchCandidate>, + rerank_windows: HashSet<u64>, + retained_candidates: BinaryHeap<SearchCandidate>, + frontier: BinaryHeap<Reverse<SearchCandidate>>, + selected_nodes: Vec<usize>, + loaded_adjacency_pages: HashSet<usize>, + adjacency_windows: AdjacencyWindowCache, + vector_windows: VectorWindowCache, + window_buffers: WindowBufferPool, + neighbor_buffer: Vec<u32>, + scored_neighbors: Vec<SearchCandidate>, +} + +#[derive(Default)] +struct VectorWindowCache { + entries: HashMap<u64, WindowPayload>, + recency: OffsetLru, + retained_capacity: usize, +} + +#[derive(Debug, Default)] +struct VectorWindowLoadStats { + hits: usize, + misses: usize, + evictions: usize, +} + +impl VectorWindowCache { + fn contains_key(&self, offset: &u64) -> bool { + self.entries.contains_key(offset) + } + + fn get(&self, offset: &u64) -> Option<&[u8]> { + self.entries.get(offset).map(WindowPayload::as_slice) + } + + fn insert(&mut self, offset: u64, payload: impl Into<WindowPayload>) { + let payload = payload.into(); + let payload_capacity = payload.capacity(); + if let Some(previous) = self.entries.insert(offset, payload) { + self.retained_capacity = self.retained_capacity.saturating_sub(previous.capacity()); + self.recency.remove(offset); + } + self.retained_capacity = self.retained_capacity.saturating_add(payload_capacity); + self.recency.touch(offset); + } + + #[cfg(test)] + fn remove(&mut self, offset: u64) -> Option<WindowPayload> { + let payload = self.entries.remove(&offset)?; + self.retained_capacity = self.retained_capacity.saturating_sub(payload.capacity()); + self.recency.remove(offset); + Some(payload) + } + + #[cfg(test)] + fn touch(&mut self, offset: u64) { + if self.entries.contains_key(&offset) { + self.recency.touch(offset); + } + } + + fn touch_windows(&mut self, windows: &[ReadWindow]) { + for window in windows { + debug_assert!(self.entries.contains_key(&window.offset)); + self.recency.touch(window.offset); + } + } + + fn trim(&mut self, window_buffers: &mut WindowBufferPool, capacity_limit: usize) -> usize { + let mut evictions = 0usize; + while self.retained_capacity > capacity_limit { + let Some(offset) = self.recency.pop_oldest() else { + break; + }; + if let Some(payload) = self.entries.remove(&offset) { + self.retained_capacity = self.retained_capacity.saturating_sub(payload.capacity()); + if let WindowPayload::Owned(payload) = payload { + window_buffers.recycle(payload); + } + evictions = evictions.saturating_add(1); + } + } + evictions + } + + fn recycle(&mut self, window_buffers: &mut WindowBufferPool) { + for (_, payload) in self.entries.drain() { + if let WindowPayload::Owned(payload) = payload { + window_buffers.recycle(payload); + } + } + self.recency.clear(); + self.retained_capacity = 0; + } + + fn len(&self) -> usize { + self.entries.len() + } + + #[cfg(test)] + fn is_empty(&self) -> bool { + self.entries.is_empty() + } + + #[cfg(test)] + fn retained_capacity(&self) -> usize { + self.retained_capacity + } +} + +struct WindowBufferPool { + buffers: Vec<Vec<u8>>, + retained_capacity: usize, + retained_capacity_limit: usize, +} + +impl Default for WindowBufferPool { + fn default() -> Self { + Self { + buffers: Vec::new(), + retained_capacity: 0, + retained_capacity_limit: QUERY_WINDOW_BUFFER_LIMIT_BYTES, + } + } +} + +impl WindowBufferPool { + #[cfg(test)] + fn with_retained_capacity_limit(retained_capacity_limit: usize) -> Self { + Self { + retained_capacity_limit, + ..Self::default() + } + } + + fn recycle(&mut self, mut buffer: Vec<u8>) { + buffer.clear(); + let capacity = buffer.capacity(); + let Some(retained_capacity) = self.retained_capacity.checked_add(capacity) else { + return; + }; + if retained_capacity > self.retained_capacity_limit { + return; + } + self.retained_capacity = retained_capacity; + self.buffers.push(buffer); + } + + fn take(&mut self, len: usize) -> io::Result<Vec<u8>> { + let best_fit = self + .buffers + .last() + .is_some_and(|buffer| buffer.capacity() == len) + .then(|| self.buffers.len() - 1) + .or_else(|| { + self.buffers + .iter() + .enumerate() + .filter(|(_, buffer)| buffer.capacity() >= len) + .min_by_key(|(_, buffer)| buffer.capacity()) + .map(|(index, _)| index) + }) + .or_else(|| { + self.buffers + .iter() + .enumerate() + .max_by_key(|(_, buffer)| buffer.capacity()) + .map(|(index, _)| index) + }); + let mut buffer = if let Some(index) = best_fit { + let buffer = self.buffers.swap_remove(index); + self.retained_capacity -= buffer.capacity(); + buffer + } else { + Vec::new() + }; + let additional_capacity = len.saturating_sub(buffer.capacity()); + if additional_capacity != 0 && buffer.try_reserve_exact(additional_capacity).is_err() { + self.recycle(buffer); + return Err(invalid_input( + "DiskANN query window buffer allocation failed", + )); + } + buffer.resize(len, 0); + Ok(buffer) + } + + fn set_retained_capacity_limit(&mut self, retained_capacity_limit: usize) { + self.retained_capacity_limit = retained_capacity_limit; + while self.retained_capacity > self.retained_capacity_limit { + let Some(buffer) = self.buffers.pop() else { + self.retained_capacity = 0; + break; + }; + self.retained_capacity -= buffer.capacity(); + } + } +} + +impl DiskAnnQueryScratch { + #[cfg(test)] + fn with_window_buffer_limit(retained_capacity_limit: usize) -> Self { + Self { + window_buffers: WindowBufferPool::with_retained_capacity_limit(retained_capacity_limit), + ..Self::default() + } + } + + fn set_window_buffer_limit(&mut self, retained_capacity_limit: usize) { + self.window_buffers + .set_retained_capacity_limit(retained_capacity_limit); + } + + #[cfg(test)] + fn begin_search(&mut self, vector_count: usize) { + self.begin_graph_search(vector_count, vector_count, 1) + .expect("test-sized DiskANN visited allocation"); + } + + fn begin_graph_search( + &mut self, + vector_count: usize, + search_list_size: usize, + max_degree: usize, + ) -> io::Result<()> { + self.begin_rerank(); + let expected_visited = search_list_size + .saturating_mul(max_degree) + .saturating_add(1) + .min(vector_count); + let dense_bytes = vector_count.div_ceil(8); + let sparse_bytes = + sparse_table_memory_bytes(expected_visited, size_of::<()>()).unwrap_or(usize::MAX); + // Dense bitmap probes are substantially cheaper than open-addressed + // hashing. Prefer them unless sparse storage saves at least 2x memory. + self.uses_sparse_visited = sparse_bytes + .checked_mul(SPARSE_VISITED_MIN_MEMORY_SAVINGS) + .is_some_and(|threshold| threshold < dense_bytes); + if self.uses_sparse_visited { + if expected_visited > self.sparse_visited.entry_capacity() { + self.sparse_visited = SparseTable::try_with_capacity(expected_visited) + .map_err(|_| invalid_input("DiskANN sparse visited allocation failed"))?; + } + } else { + self.visited.resize(vector_count, false); + } + Ok(()) + } + + fn begin_rerank(&mut self) { + if self.uses_sparse_visited { + self.sparse_visited.clear(); + self.touched_nodes.clear(); + } else { + for node in self.touched_nodes.drain(..) { + self.visited[node] = false; + } + } + self.candidates.clear(); + self.rerank_candidates.clear(); + self.rerank_windows.clear(); + self.retained_candidates.clear(); + self.frontier.clear(); + self.selected_nodes.clear(); + self.loaded_adjacency_pages.clear(); + self.recycle_adjacency_windows(); + self.neighbor_buffer.clear(); + self.scored_neighbors.clear(); + } + + fn recycle_adjacency_windows(&mut self) { + self.adjacency_windows.recycle(&mut self.window_buffers); + } + + fn recycle_vector_windows(&mut self) { + self.vector_windows.recycle(&mut self.window_buffers); + } + + fn recycle_window_caches(&mut self) { + self.recycle_adjacency_windows(); + self.recycle_vector_windows(); + } + + fn prepare_distance_table(&mut self, len: usize) -> &mut [f32] { + self.distance_table.resize(len, 0.0); + &mut self.distance_table + } + + fn select_round(&mut self, limit: usize) { + self.selected_nodes.clear(); + while self.selected_nodes.len() < limit { + let Some(Reverse(candidate)) = self.frontier.pop() else { + break; + }; + if self + .retained_candidates + .peek() + .is_some_and(|worst| candidate > *worst) + { + self.frontier.clear(); + break; + } + self.selected_nodes.push(candidate.node); + } + } + + fn insert_graph_candidate( + &mut self, + candidate: SearchCandidate, + limit: usize, + ) -> io::Result<()> { + if limit == 0 { + return Ok(()); + } + let replacing_worst = self.retained_candidates.len() == limit; + if replacing_worst { + let Some(worst) = self.retained_candidates.peek().copied() else { + return Ok(()); + }; + if candidate >= worst { + return Ok(()); + } + } else { + self.retained_candidates + .try_reserve(1) + .map_err(|_| invalid_input("DiskANN graph candidate allocation failed"))?; + } + self.frontier + .try_reserve(1) + .map_err(|_| invalid_input("DiskANN graph frontier allocation failed"))?; + if replacing_worst { + self.retained_candidates.pop(); + } + self.retained_candidates.push(candidate); + self.frontier.push(Reverse(candidate)); + if self.frontier.len() > limit.saturating_mul(2) { + let worst = *self + .retained_candidates + .peek() + .expect("non-empty retained DiskANN candidates"); + self.frontier + .retain(|Reverse(candidate)| *candidate <= worst); + } + Ok(()) + } + + fn finish_graph_candidates(&mut self) { + self.candidates.extend(self.retained_candidates.drain()); + sort_candidates(&mut self.candidates); + } + + #[cfg(test)] + fn is_visited(&self, node: usize) -> bool { + if self.uses_sparse_visited { + self.sparse_visited.get(node as u32).is_some() + } else { + self.visited[node] + } + } + + fn mark_visited(&mut self, node: usize) -> bool { + if self.uses_sparse_visited { + return self.sparse_visited.insert(node as u32, ()).is_none(); + } + if self.visited[node] { + return false; + } + self.visited[node] = true; + self.touched_nodes.push(node); + true + } + + #[cfg(test)] + fn visited_capacity(&self) -> usize { + self.visited.capacity() + } + + #[cfg(test)] + fn uses_sparse_visited(&self) -> bool { + self.uses_sparse_visited + } + + #[cfg(test)] + fn retained_window_capacity(&self) -> usize { + self.window_buffers.retained_capacity + } +} + +fn window_buffer_limit_per_worker(worker_count: usize) -> usize { + QUERY_WINDOW_BUFFER_LIMIT_BYTES.min(BATCH_WINDOW_BUFFER_LIMIT_BYTES / worker_count.max(1)) +} + +fn prepare_vector_window_cache( + windows: &[ReadWindow], + cache: &mut VectorWindowCache, + window_buffers: &mut WindowBufferPool, + capacity_limit: usize, +) -> (bool, usize) { + let retain = windows + .iter() + .try_fold(0usize, |total, window| total.checked_add(window.length)) + .is_some_and(|total| total <= capacity_limit); + if !retain { + let evictions = cache.len(); + cache.recycle(window_buffers); + return (false, evictions); + } + let evictions = cache.trim(window_buffers, capacity_limit); + (true, evictions) +} + +type CandidatePartition = Vec<(usize, Vec<usize>)>; +type SessionQueryOutput = (usize, Vec<i64>, Vec<f32>, DiskAnnSearchStats); + +impl PartialEq for SearchCandidate { + fn eq(&self, other: &Self) -> bool { + self.node == other.node && self.distance.to_bits() == other.distance.to_bits() + } +} + +impl Eq for SearchCandidate {} + +impl PartialOrd for SearchCandidate { + fn partial_cmp(&self, other: &Self) -> Option<Ordering> { + Some(self.cmp(other)) + } +} + +impl Ord for SearchCandidate { + fn cmp(&self, other: &Self) -> Ordering { + self.distance + .total_cmp(&other.distance) + .then_with(|| self.node.cmp(&other.node)) + } +} + +impl PartialEq for ExactSearchResult { + fn eq(&self, other: &Self) -> bool { + self.row_id == other.row_id && self.distance.to_bits() == other.distance.to_bits() + } +} + +impl Eq for ExactSearchResult {} + +impl PartialOrd for ExactSearchResult { + fn partial_cmp(&self, other: &Self) -> Option<Ordering> { + Some(self.cmp(other)) + } +} + +impl Ord for ExactSearchResult { + fn cmp(&self, other: &Self) -> Ordering { + self.distance + .total_cmp(&other.distance) + .then_with(|| self.row_id.cmp(&other.row_id)) + } +} + +impl<R: SeekRead> crate::diskann_io::DiskAnnIndexReader<R> { + fn preprocess_queries<'a>(&self, queries: &'a [f32], query_count: usize) -> Cow<'a, [f32]> { + if self.header.metric_type() == MetricType::Cosine { + Cow::Owned(preprocess_vectors( + queries, + query_count, + self.header.dimension as usize, + MetricType::Cosine, + )) + } else { + Cow::Borrowed(queries) + } + } + + fn take_batch_workers( + &mut self, + worker_count: usize, + filtered: bool, + ) -> io::Result<Option<Vec<Self>>> { + if filtered { + self.ensure_resident()?; + } else { + self.optimize_for_search()?; + } + let mut workers = std::mem::take(&mut self.batch_workers); + workers.truncate(worker_count); + while workers.len() < worker_count { + let worker = if filtered { + self.try_clone_for_filtered_search() + } else { + self.try_clone_for_search() + }; + match worker { + Ok(Some(worker)) => workers.push(worker), + Ok(None) => { + self.batch_workers = workers; + return Ok(None); + } + Err(error) => { + self.batch_workers = workers; + return Err(error); + } + } + } + let window_buffer_limit = window_buffer_limit_per_worker(worker_count); + for worker in &mut workers { + worker.refresh_shared_state_from(self); + worker.limit_raw_vector_cache_bytes(window_buffer_limit); + worker + .query_scratch + .set_window_buffer_limit(window_buffer_limit); + } + Ok(Some(workers)) + } + + pub(crate) fn search_batch( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + self.last_search_stats = DiskAnnSearchStats { + query_count, + ..DiskAnnSearchStats::default() + }; + if top_k == 0 { + return Ok((Vec::new(), Vec::new())); + } + let processed_queries = self.preprocess_queries(queries, query_count); + let queries = processed_queries.as_ref(); + let worker_count = query_count.min(rayon::current_num_threads()); + if worker_count <= 1 { + self.batch_workers.clear(); + if self.header.is_interleaved() { + return self.search_batch_direct_serial(queries, top_k, l_search); + } + return self.search_batch_serial(queries, top_k, l_search); + } + + let Some(mut workers) = self.take_batch_workers(worker_count, false)? else { + if self.header.is_interleaved() { + return self.search_batch_direct_serial(queries, top_k, l_search); + } + return self.search_batch_serial(queries, top_k, l_search); + }; + if self.header.is_interleaved() + || query_count <= worker_count.saturating_mul(PARALLEL_SESSION_MAX_QUERIES_PER_WORKER) + { + let result = + self.search_batch_in_parallel_sessions(queries, top_k, l_search, &mut workers); + self.batch_workers = workers; + return result; + } + let result = (|| { + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query_chunk in queries.chunks(BATCH_QUERY_CHUNK_SIZE * dimension) { + let chunk_query_count = query_chunk.len() / dimension; + self.record_query_chunk(chunk_query_count); + let worker_outputs = workers + .par_iter_mut() + .enumerate() + .map(|(worker_index, worker)| { + worker.last_search_stats = DiskAnnSearchStats::default(); + let mut partition = Vec::new(); + for query_index in (worker_index..chunk_query_count).step_by(worker_count) { + let query = &query_chunk + [query_index * dimension..(query_index + 1) * dimension]; + let candidates = worker + .generate_unfiltered_candidate_nodes(query, top_k, l_search)?; + partition.push((query_index, candidates)); + } + Ok::<_, io::Error>((partition, worker.last_search_stats)) + }) + .collect::<io::Result<Vec<_>>>()?; + let mut partitions = Vec::with_capacity(worker_outputs.len()); + for (partition, worker_stats) in worker_outputs { + self.last_search_stats + .merge_candidate_generation(worker_stats); + partitions.push(partition); + } + let (chunk_ids, chunk_distances) = + self.rerank_candidate_batch_streaming(query_chunk, top_k, partitions)?; + ids.extend(chunk_ids); + distances.extend(chunk_distances); + } + Ok((ids, distances)) + })(); + self.batch_workers = workers; + result + } + + /// Warms resident metadata and the query-dependent adjacency/raw-vector + /// caches with representative queries without changing reported search + /// statistics. + pub fn warmup_queries(&mut self, queries: &[f32], l_search: usize) -> io::Result<()> { + let dimension = self.header.dimension as usize; + if !queries.len().is_multiple_of(dimension) { + return Err(invalid_input(format!( + "warmup query length {} is not divisible by dimension {}", + queries.len(), + dimension + ))); + } + if queries.iter().any(|value| !value.is_finite()) { + return Err(invalid_input("warmup query values must be finite")); + } + self.optimize_for_search()?; + if queries.is_empty() { + return Ok(()); + } + let saved_stats = self.last_search_stats; + // Replay on the parent Reader so both adjacency and raw-vector windows + // are useful to the subsequent single-query path. A batch warm-up may + // otherwise populate only retained worker-local raw-vector caches. + let result = queries + .chunks_exact(dimension) + .try_for_each(|query| self.search(query, 1, l_search).map(|_| ())); + self.last_search_stats = saved_stats; + result + } + + /// Calibrates the automatic search width from representative queries. + /// + /// This is a stability proxy, not a ground-truth recall guarantee: it + /// chooses the first width whose Top-K overlap with the next wider search + /// reaches 98% across the sample. + pub fn calibrate_l_search(&mut self, queries: &[f32], top_k: usize) -> io::Result<usize> { + let dimension = self.header.dimension as usize; + if queries.is_empty() || !queries.len().is_multiple_of(dimension) { + return Err(invalid_input( + "calibration queries must contain one or more complete vectors", + )); + } + if top_k == 0 { + return Err(invalid_input("calibration top_k must be greater than 0")); + } + if queries.iter().any(|value| !value.is_finite()) { + return Err(invalid_input("calibration query values must be finite")); + } + self.optimize_for_search()?; + let widths = [ + 100usize.max(top_k), + 200usize.max(top_k), + 400usize.max(top_k), + ]; + let mut results = Vec::with_capacity(widths.len()); + let saved_stats = self.last_search_stats; + for width in widths { + let (ids, _) = self.search_batch(queries, top_k, width)?; + results.push(ids); + } + self.last_search_stats = saved_stats; + let chosen = if topk_result_stability(&results[0], &results[1], top_k) >= 0.98 { + widths[0] + } else if topk_result_stability(&results[1], &results[2], top_k) >= 0.98 { + widths[1] + } else { + widths[2] + }; + self.calibrated_l_search = Some(chosen); + for worker in &mut self.batch_workers { + worker.calibrated_l_search = Some(chosen); + } + Ok(chosen) + } + + pub(crate) fn search_batch_with_roaring_filter( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + roaring_filter_bytes: &[u8], + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let filter = decode_roaring_filter(roaring_filter_bytes)?; + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + self.last_search_stats = DiskAnnSearchStats { + query_count, + ..DiskAnnSearchStats::default() + }; + if top_k == 0 { + return Ok((Vec::new(), Vec::new())); + } + let processed_queries = self.preprocess_queries(queries, query_count); + let queries = processed_queries.as_ref(); + if filter.is_empty() { + return Ok(( + vec![-1; query_count * top_k], + vec![f32::MAX; query_count * top_k], + )); + } + let matching_nodes = self.matching_nodes_for_filter(&filter)?; + if matching_nodes.is_empty() { + return Ok(( + vec![-1; query_count * top_k], + vec![f32::MAX; query_count * top_k], + )); + } + if self.header.is_interleaved() { + let worker_count = query_count.min(rayon::current_num_threads()); + if worker_count <= 1 { + self.batch_workers.clear(); + return self.search_batch_with_matching_nodes_direct_serial( + queries, + top_k, + l_search, + &matching_nodes, + ); + } + let Some(mut workers) = self.take_batch_workers(worker_count, true)? else { + return self.search_batch_with_matching_nodes_direct_serial( + queries, + top_k, + l_search, + &matching_nodes, + ); + }; + let result = self.search_batch_with_matching_nodes_in_parallel_sessions( + queries, + top_k, + l_search, + &matching_nodes, + &mut workers, + ); + self.batch_workers = workers; + return result; + } + let matching_count = usize::try_from(matching_nodes.len()).unwrap_or(usize::MAX); + if let FilteredCandidateStrategy::Exhaustive { target_candidates } = + select_filtered_candidate_strategy( + self.header.vector_count as usize, + matching_count, + top_k, + l_search, + self.header.max_degree as usize, + self.read_plan(), + self.adjacency_fully_preloaded(), + ) + { + return self.search_batch_filtered_exhaustive( + queries, + top_k, + &matching_nodes, + target_candidates, + ); + } + let worker_count = query_count.min(rayon::current_num_threads()); + if worker_count <= 1 { + self.batch_workers.clear(); + return self.search_batch_with_matching_nodes_serial( + queries, + top_k, + l_search, + &matching_nodes, + ); + } + + let Some(mut workers) = self.take_batch_workers(worker_count, true)? else { + return self.search_batch_with_matching_nodes_serial( + queries, + top_k, + l_search, + &matching_nodes, + ); + }; + if query_count <= worker_count.saturating_mul(PARALLEL_SESSION_MAX_QUERIES_PER_WORKER) { + let result = self.search_batch_with_matching_nodes_in_parallel_sessions( + queries, + top_k, + l_search, + &matching_nodes, + &mut workers, + ); + self.batch_workers = workers; + return result; + } + let result = (|| { + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query_chunk in queries.chunks(BATCH_QUERY_CHUNK_SIZE * dimension) { + let chunk_query_count = query_chunk.len() / dimension; + self.record_query_chunk(chunk_query_count); + let worker_outputs = workers + .par_iter_mut() + .enumerate() + .map(|(worker_index, worker)| { + worker.last_search_stats = DiskAnnSearchStats::default(); + let mut partition = Vec::new(); + for query_index in (worker_index..chunk_query_count).step_by(worker_count) { + let query = &query_chunk + [query_index * dimension..(query_index + 1) * dimension]; + let candidates = worker.generate_filtered_candidates( + query, + top_k, + l_search, + &matching_nodes, + )?; + partition.push(( + query_index, + candidates + .into_iter() + .map(|candidate| candidate.node) + .collect(), + )); + } + Ok::<_, io::Error>((partition, worker.last_search_stats)) + }) + .collect::<io::Result<Vec<_>>>()?; + let mut partitions = Vec::with_capacity(worker_outputs.len()); + for (partition, worker_stats) in worker_outputs { + self.last_search_stats + .merge_candidate_generation(worker_stats); + partitions.push(partition); + } + let (chunk_ids, chunk_distances) = + self.rerank_candidate_batch_streaming(query_chunk, top_k, partitions)?; + ids.extend(chunk_ids); + distances.extend(chunk_distances); + } + Ok((ids, distances)) + })(); + self.batch_workers = workers; + result + } + + fn search_batch_in_parallel_sessions( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + workers: &mut [Self], + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + self.record_query_chunk(query_count); + let worker_count = workers.len(); + let worker_outputs = workers + .par_iter_mut() + .enumerate() + .map(|(worker_index, worker)| { + let mut outputs = Vec::new(); + for query_index in (worker_index..query_count).step_by(worker_count) { + let query = &queries[query_index * dimension..(query_index + 1) * dimension]; + let (ids, distances) = worker.search_preprocessed(query, top_k, l_search)?; + outputs.push((query_index, ids, distances, worker.last_search_stats)); + } + Ok::<_, io::Error>(outputs) + }) + .collect::<io::Result<Vec<_>>>()?; + self.collect_parallel_session_outputs(worker_outputs, query_count, top_k) + } + + fn search_batch_with_matching_nodes_in_parallel_sessions( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + matching_nodes: &RoaringBitmap, + workers: &mut [Self], + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + self.record_query_chunk(query_count); + let worker_count = workers.len(); + let worker_outputs = workers + .par_iter_mut() + .enumerate() + .map(|(worker_index, worker)| { + let mut outputs = Vec::new(); + for query_index in (worker_index..query_count).step_by(worker_count) { + let query = &queries[query_index * dimension..(query_index + 1) * dimension]; + worker.last_search_stats = DiskAnnSearchStats { + query_count: 1, + ..DiskAnnSearchStats::default() + }; + let (ids, distances) = worker.search_with_matching_nodes( + query, + top_k, + l_search, + matching_nodes, + )?; + outputs.push((query_index, ids, distances, worker.last_search_stats)); + } + Ok::<_, io::Error>(outputs) + }) + .collect::<io::Result<Vec<_>>>()?; + self.collect_parallel_session_outputs(worker_outputs, query_count, top_k) + } + + fn collect_parallel_session_outputs( + &mut self, + worker_outputs: Vec<Vec<SessionQueryOutput>>, + query_count: usize, + top_k: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let mut ordered = (0..query_count) + .map(|_| None) + .collect::<Vec<Option<(Vec<i64>, Vec<f32>)>>>(); + for (query_index, ids, distances, stats) in worker_outputs.into_iter().flatten() { + let slot = ordered + .get_mut(query_index) + .ok_or_else(|| invalid_data("DiskANN parallel session query index is invalid"))?; + if slot.replace((ids, distances)).is_some() { + return Err(invalid_data( + "DiskANN parallel sessions returned a duplicate query", + )); + } + self.last_search_stats.merge_complete_query(stats); + } + if ordered.iter().any(Option::is_none) { + return Err(invalid_data( + "DiskANN parallel sessions did not return every query", + )); + } + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for result in ordered { + let (query_ids, query_distances) = + result.expect("validated DiskANN parallel session output"); + ids.extend(query_ids); + distances.extend(query_distances); + } + Ok((ids, distances)) + } + + fn search_batch_filtered_exhaustive( + &mut self, + queries: &[f32], + top_k: usize, + matching_nodes: &RoaringBitmap, + candidate_limit: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let matching_count = usize::try_from(matching_nodes.len()).unwrap_or(usize::MAX); + let pq_m = self.header.pq_m as usize; + let pq_ksub = 1usize << self.header.pq_bits; + let tile_size = filtered_pq_query_tile_size(pq_m, pq_ksub); + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query_chunk in queries.chunks(BATCH_QUERY_CHUNK_SIZE * dimension) { + let chunk_query_count = query_chunk.len() / dimension; + self.record_query_chunk(chunk_query_count); + let candidate_sets = self.exhaustive_filtered_candidate_nodes_batch( + query_chunk, + matching_nodes, + candidate_limit, + )?; + self.last_search_stats.filtered_exhaustive_queries = self + .last_search_stats + .filtered_exhaustive_queries + .saturating_add(chunk_query_count); + self.last_search_stats.pq_distance_evaluations = self + .last_search_stats + .pq_distance_evaluations + .saturating_add(chunk_query_count.saturating_mul(matching_count)); + self.last_search_stats.pq_code_loads = + self.last_search_stats.pq_code_loads.saturating_add( + chunk_query_count + .div_ceil(tile_size) + .saturating_mul(matching_count), + ); + let partition = candidate_sets.into_iter().enumerate().collect(); + let (chunk_ids, chunk_distances) = + self.rerank_candidate_batch_streaming(query_chunk, top_k, vec![partition])?; + ids.extend(chunk_ids); + distances.extend(chunk_distances); + } + Ok((ids, distances)) + } + + fn exhaustive_filtered_candidate_nodes_batch( + &self, + queries: &[f32], + matching_nodes: &RoaringBitmap, + candidate_limit: usize, + ) -> io::Result<Vec<Vec<usize>>> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let pq_m = self.header.pq_m as usize; + let pq = self.pq()?; + let pq_ksub = pq.ksub; + let pq_code_size = pq.code_size(); + let pq_codes = self.pq_codes()?; + let metric = self.header.metric_type(); + let tile_size = filtered_pq_query_tile_size(pq_m, pq_ksub); + let tile_count = query_count.div_ceil(tile_size); + let mut tiles = (0..tile_count) + .into_par_iter() + .map(|tile_index| { + let query_start = tile_index * tile_size; + let query_end = (query_start + tile_size).min(query_count); + let tile_query_count = query_end - query_start; + let mut distance_tables = vec![0.0f32; tile_query_count * pq_m * pq_ksub]; + for tile_query_index in 0..tile_query_count { + let query_index = query_start + tile_query_index; + let query = &queries[query_index * dimension..(query_index + 1) * dimension]; + let table_start = tile_query_index * pq_m * pq_ksub; + pq.compute_distance_table( + query, + metric, + &mut distance_tables[table_start..table_start + pq_m * pq_ksub], + ); + } + let mut heaps = (0..tile_query_count) + .map(|_| BinaryHeap::new()) + .collect::<Vec<BinaryHeap<SearchCandidate>>>(); + for heap in &mut heaps { + heap.try_reserve(candidate_limit.min(1024)).map_err(|_| { + invalid_input("DiskANN filtered candidate allocation failed") + })?; + } + for node in matching_nodes.iter() { + let node = node as usize; + let code_start = node + .checked_mul(pq_code_size) + .ok_or_else(|| invalid_data("DiskANN PQ code offset overflows"))?; + let codes = pq_codes + .get(code_start..code_start + pq_code_size) + .ok_or_else(|| invalid_data("DiskANN PQ codes are truncated"))?; + for tile_query_index in 0..tile_query_count { + let table_start = tile_query_index * pq_m * pq_ksub; + push_bounded_candidate( + &mut heaps[tile_query_index], + SearchCandidate { + node, + distance: pq.distance_from_table( + &distance_tables[table_start..table_start + pq_m * pq_ksub], + codes, + ), + }, + candidate_limit, + )?; + } + } + let candidates = heaps + .into_iter() + .map(|heap| { + let mut candidates = heap.into_vec(); + sort_candidates(&mut candidates); + candidates + .into_iter() + .map(|candidate| candidate.node) + .collect::<Vec<_>>() + }) + .collect::<Vec<_>>(); + Ok::<_, io::Error>((query_start, candidates)) + }) + .collect::<io::Result<Vec<_>>>()?; + tiles.sort_unstable_by_key(|(query_start, _)| *query_start); + Ok(tiles + .into_iter() + .flat_map(|(_, candidates)| candidates) + .collect()) + } + + fn search_batch_direct_serial( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let mut aggregate = DiskAnnSearchStats { + query_count, + query_chunks: usize::from(query_count != 0), + max_queries_per_chunk: query_count, + ..DiskAnnSearchStats::default() + }; + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query in queries.chunks_exact(dimension) { + let (query_ids, query_distances) = self.search_preprocessed(query, top_k, l_search)?; + aggregate.merge_complete_query(self.last_search_stats); + ids.extend(query_ids); + distances.extend(query_distances); + } + aggregate.parallel_session_queries = 0; + self.last_search_stats = aggregate; + Ok((ids, distances)) + } + + fn search_batch_with_matching_nodes_direct_serial( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + matching_nodes: &RoaringBitmap, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let mut aggregate = DiskAnnSearchStats { + query_count, + query_chunks: usize::from(query_count != 0), + max_queries_per_chunk: query_count, + ..DiskAnnSearchStats::default() + }; + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query in queries.chunks_exact(dimension) { + self.last_search_stats = DiskAnnSearchStats { + query_count: 1, + ..DiskAnnSearchStats::default() + }; + let (query_ids, query_distances) = + self.search_with_matching_nodes(query, top_k, l_search, matching_nodes)?; + aggregate.merge_complete_query(self.last_search_stats); + ids.extend(query_ids); + distances.extend(query_distances); + } + aggregate.parallel_session_queries = 0; + self.last_search_stats = aggregate; + Ok((ids, distances)) + } + + fn search_batch_serial( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query_chunk in queries.chunks(BATCH_QUERY_CHUNK_SIZE * dimension) { + let chunk_query_count = query_chunk.len() / dimension; + self.record_query_chunk(chunk_query_count); + let mut partition = Vec::with_capacity(chunk_query_count); + for (query_index, query) in query_chunk.chunks_exact(dimension).enumerate() { + let candidates = + self.generate_unfiltered_candidate_nodes(query, top_k, l_search)?; + partition.push((query_index, candidates)); + } + let (chunk_ids, chunk_distances) = + self.rerank_candidate_batch_streaming(query_chunk, top_k, vec![partition])?; + ids.extend(chunk_ids); + distances.extend(chunk_distances); + } + Ok((ids, distances)) + } + + fn search_batch_with_matching_nodes_serial( + &mut self, + queries: &[f32], + top_k: usize, + l_search: usize, + matching_nodes: &RoaringBitmap, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for query_chunk in queries.chunks(BATCH_QUERY_CHUNK_SIZE * dimension) { + let chunk_query_count = query_chunk.len() / dimension; + self.record_query_chunk(chunk_query_count); + let mut partition = Vec::with_capacity(chunk_query_count); + for (query_index, query) in query_chunk.chunks_exact(dimension).enumerate() { + let candidates = + self.generate_filtered_candidates(query, top_k, l_search, matching_nodes)?; + partition.push(( + query_index, + candidates + .into_iter() + .map(|candidate| candidate.node) + .collect(), + )); + } + let (chunk_ids, chunk_distances) = + self.rerank_candidate_batch_streaming(query_chunk, top_k, vec![partition])?; + ids.extend(chunk_ids); + distances.extend(chunk_distances); + } + Ok((ids, distances)) + } + + fn record_query_chunk(&mut self, query_count: usize) { + self.last_search_stats.query_chunks = self.last_search_stats.query_chunks.saturating_add(1); + self.last_search_stats.max_queries_per_chunk = self + .last_search_stats + .max_queries_per_chunk + .max(query_count); + } + + fn generate_unfiltered_candidate_nodes( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<Vec<usize>> { + self.ensure_resident()?; + let search_list_size = resolve_diskann_l_search(top_k, l_search); + let mut scratch = std::mem::take(&mut self.query_scratch); + let result = (|| { + self.generate_graph_candidates( + query, + search_list_size, + self.read_plan().graph_beam_width, + &mut scratch, + )?; + let rerank_count = search_list_size + .min(top_k.saturating_mul(4).max(64)) + .min(scratch.candidates.len()); + if rerank_count == scratch.candidates.len() { + return Ok(scratch + .candidates + .iter() + .map(|candidate| candidate.node) + .collect()); + } + if self.header.is_interleaved() { + let planner = + ReadWindowPlanner::new(self.read_plan(), self.header.sections.adjacency); + expand_rerank_candidates_within_seed_windows( + &scratch.candidates, + rerank_count, + |node| { + let page = self.adjacency_locator(node)?.page_index as usize; + planner + .window_for_logical_page(page) + .map(|window| window.offset) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range")) + }, + &mut scratch.rerank_windows, + &mut scratch.rerank_candidates, + )?; + return Ok(scratch + .rerank_candidates + .iter() + .map(|candidate| candidate.node) + .collect()); + } + let planner = VectorWindowPlanner::new( + self.read_plan(), + self.header.sections.vectors, + self.header.vector_record_size as usize, + )?; + expand_rerank_candidates_within_seed_windows( + &scratch.candidates, + rerank_count, + |node| { + planner + .window_for_node(node) + .map(|window| window.offset) + .ok_or_else(|| invalid_data("DiskANN raw-vector record is out of range")) + }, + &mut scratch.rerank_windows, + &mut scratch.rerank_candidates, + )?; + Ok(scratch + .rerank_candidates + .iter() + .map(|candidate| candidate.node) + .collect()) + })(); + scratch.recycle_adjacency_windows(); + self.query_scratch = scratch; + result + } + + pub fn search( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + self.last_search_stats = DiskAnnSearchStats { + query_count: 1, + ..DiskAnnSearchStats::default() + }; + let dimension = self.header.dimension as usize; + if query.len() != dimension { + return Err(invalid_input(format!( + "query dimension mismatch: expected {}, got {}", + dimension, + query.len() + ))); + } + if query.iter().any(|value| !value.is_finite()) { + return Err(invalid_input("query values must be finite")); + } + let processed_query = self.preprocess_queries(query, 1); + self.search_preprocessed(processed_query.as_ref(), top_k, l_search) + } + + fn search_preprocessed( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + self.last_search_stats = DiskAnnSearchStats { + query_count: 1, + ..DiskAnnSearchStats::default() + }; + if top_k == 0 { + return Ok((Vec::new(), Vec::new())); + } + self.ensure_resident()?; + let search_list_size = resolve_diskann_l_search(top_k, l_search); + let mut scratch = std::mem::take(&mut self.query_scratch); + let result = (|| { + self.generate_graph_candidates( + query, + search_list_size, + self.read_plan().graph_beam_width, + &mut scratch, + )?; + + let rerank_count = search_list_size + .min(top_k.saturating_mul(4).max(64)) + .min(scratch.candidates.len()); + if rerank_count == scratch.candidates.len() { + if self.header.is_interleaved() { + return self.rerank_interleaved( + query, + &scratch.candidates, + top_k, + &mut scratch.adjacency_windows, + &mut scratch.loaded_adjacency_pages, + &mut scratch.window_buffers, + ); + } + return self.rerank( + query, + &scratch.candidates, + top_k, + &mut scratch.vector_windows, + &mut scratch.window_buffers, + ); + } + if self.header.is_interleaved() { + let planner = + ReadWindowPlanner::new(self.read_plan(), self.header.sections.adjacency); + expand_rerank_candidates_within_seed_windows( + &scratch.candidates, + rerank_count, + |node| { + let page = self.adjacency_locator(node)?.page_index as usize; + planner + .window_for_logical_page(page) + .map(|window| window.offset) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range")) + }, + &mut scratch.rerank_windows, + &mut scratch.rerank_candidates, + )?; + return self.rerank_interleaved( + query, + &scratch.rerank_candidates, + top_k, + &mut scratch.adjacency_windows, + &mut scratch.loaded_adjacency_pages, + &mut scratch.window_buffers, + ); + } + let planner = VectorWindowPlanner::new( + self.read_plan(), + self.header.sections.vectors, + self.header.vector_record_size as usize, + )?; + expand_rerank_candidates_within_seed_windows( + &scratch.candidates, + rerank_count, + |node| { + planner + .window_for_node(node) + .map(|window| window.offset) + .ok_or_else(|| invalid_data("DiskANN raw-vector record is out of range")) + }, + &mut scratch.rerank_windows, + &mut scratch.rerank_candidates, + )?; + self.rerank( + query, + &scratch.rerank_candidates, + top_k, + &mut scratch.vector_windows, + &mut scratch.window_buffers, + ) + })(); + scratch.recycle_adjacency_windows(); + self.query_scratch = scratch; + result + } + + fn generate_graph_candidates( + &mut self, + query: &[f32], + search_list_size: usize, + beam_width: usize, + scratch: &mut DiskAnnQueryScratch, + ) -> io::Result<()> { + let vector_count = self.header.vector_count as usize; + let search_list_size = search_list_size.min(vector_count); + scratch.begin_graph_search( + vector_count, + search_list_size, + self.header.max_degree as usize, + )?; + let pq = self.pq()?; + let distance_table_len = pq.m * pq.ksub; + pq.compute_distance_table( + query, + self.header.metric_type(), + scratch.prepare_distance_table(distance_table_len), + ); + + let entry_node = self.header.entry_node as usize; + scratch.mark_visited(entry_node); + scratch.insert_graph_candidate( + SearchCandidate { + node: entry_node, + distance: self.pq_distance(entry_node, &scratch.distance_table)?, + }, + search_list_size, + )?; + let mut expanded_count = 0usize; + while expanded_count < search_list_size { + scratch.select_round(beam_width.min(search_list_size - expanded_count)); + if scratch.selected_nodes.is_empty() { + break; + } + self.load_adjacency_pages( + &scratch.selected_nodes, + &mut scratch.adjacency_windows, + &mut scratch.loaded_adjacency_pages, + &mut scratch.window_buffers, + )?; + expanded_count += scratch.selected_nodes.len(); + for selected_node_index in 0..scratch.selected_nodes.len() { + let node = scratch.selected_nodes[selected_node_index]; + self.decode_adjacency_neighbors( + node, + &scratch.adjacency_windows, + &mut scratch.neighbor_buffer, + )?; + let mut retained_neighbors = 0; + for neighbor_index in 0..scratch.neighbor_buffer.len() { + let neighbor = scratch.neighbor_buffer[neighbor_index] as usize; + if !scratch.mark_visited(neighbor) { + continue; + } + scratch.neighbor_buffer[retained_neighbors] = neighbor as u32; + retained_neighbors += 1; + } + scratch.neighbor_buffer.truncate(retained_neighbors); + score_pq_neighbors( + &scratch.distance_table, + self.pq_codes()?, + self.header.pq_m as usize, + self.header.pq_bits as usize, + &scratch.neighbor_buffer, + &mut scratch.scored_neighbors, + )?; + for scored_index in 0..scratch.scored_neighbors.len() { + let candidate = scratch.scored_neighbors[scored_index]; + scratch.insert_graph_candidate(candidate, search_list_size)?; + } + } + let evictions = scratch.adjacency_windows.trim( + &mut scratch.window_buffers, + QUERY_ADJACENCY_WINDOW_LIMIT_BYTES, + ); + self.last_search_stats.query_adjacency_cache_evictions = self + .last_search_stats + .query_adjacency_cache_evictions + .saturating_add(evictions); + } + scratch.finish_graph_candidates(); + Ok(()) + } + + pub fn search_with_roaring_filter( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + roaring_filter_bytes: &[u8], + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let dimension = self.header.dimension as usize; + if query.len() != dimension { + return Err(invalid_input(format!( + "query dimension mismatch: expected {}, got {}", + dimension, + query.len() + ))); + } + if query.iter().any(|value| !value.is_finite()) { + return Err(invalid_input("query values must be finite")); + } + let processed_query = self.preprocess_queries(query, 1); + let query = processed_query.as_ref(); + let filter = decode_roaring_filter(roaring_filter_bytes)?; + self.search_with_decoded_roaring_filter(query, top_k, l_search, &filter) + } + + fn search_with_decoded_roaring_filter( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + filter: &RoaringTreemap, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + self.last_search_stats = DiskAnnSearchStats { + query_count: 1, + ..DiskAnnSearchStats::default() + }; + if top_k == 0 { + return Ok((Vec::new(), Vec::new())); + } + if filter.is_empty() { + return Ok((vec![-1; top_k], vec![f32::MAX; top_k])); + } + let matching_nodes = self.matching_nodes_for_filter(filter)?; + if matching_nodes.is_empty() { + return Ok((vec![-1; top_k], vec![f32::MAX; top_k])); + } + self.search_with_matching_nodes(query, top_k, l_search, &matching_nodes) + } + + fn exhaustive_filtered_candidates( + &mut self, + query: &[f32], + matching_nodes: &RoaringBitmap, + candidate_limit: usize, + ) -> io::Result<Vec<SearchCandidate>> { + let matching_count = usize::try_from(matching_nodes.len()).unwrap_or(usize::MAX); + self.last_search_stats.filtered_exhaustive_queries = self + .last_search_stats + .filtered_exhaustive_queries + .saturating_add(1); + self.last_search_stats.pq_distance_evaluations = self + .last_search_stats + .pq_distance_evaluations + .saturating_add(matching_count); + self.last_search_stats.pq_code_loads = self + .last_search_stats + .pq_code_loads + .saturating_add(matching_count); + let mut scratch = std::mem::take(&mut self.query_scratch); + let result = (|| { + scratch.begin_rerank(); + let pq = self.pq()?; + let distance_table_len = pq.m * pq.ksub; + pq.compute_distance_table( + query, + self.header.metric_type(), + scratch.prepare_distance_table(distance_table_len), + ); + let pq_codes = self.pq_codes()?; + let pq_m = self.header.pq_m as usize; + let pq_bits = self.header.pq_bits as usize; + let mut candidates = BinaryHeap::new(); + candidates + .try_reserve(candidate_limit.min(1024)) + .map_err(|_| invalid_input("DiskANN filtered candidate allocation failed"))?; + for node in matching_nodes.iter() { + scratch.neighbor_buffer.push(node); + if scratch.neighbor_buffer.len() == FILTERED_SINGLE_PQ_NODE_CHUNK_SIZE { + score_filtered_candidate_chunk( + &scratch.distance_table, + pq_codes, + pq_m, + pq_bits, + &scratch.neighbor_buffer, + &mut scratch.scored_neighbors, + &mut candidates, + candidate_limit, + )?; + scratch.neighbor_buffer.clear(); + } + } + if !scratch.neighbor_buffer.is_empty() { + score_filtered_candidate_chunk( + &scratch.distance_table, + pq_codes, + pq_m, + pq_bits, + &scratch.neighbor_buffer, + &mut scratch.scored_neighbors, + &mut candidates, + candidate_limit, + )?; + scratch.neighbor_buffer.clear(); + } + let mut candidates = candidates.into_vec(); + sort_candidates(&mut candidates); + Ok(candidates) + })(); + self.query_scratch = scratch; + result + } + + fn matching_nodes_for_filter(&mut self, filter: &RoaringTreemap) -> io::Result<RoaringBitmap> { + self.ensure_resident()?; + if use_row_id_order(filter.len(), self.header.vector_count as usize) { + if let Some(row_id_order) = self.ensure_row_id_order()? { + let matching_ranges = + matching_ranges_from_row_id_order(&row_id_order, filter, |node| { + self.row_id(node) + })?; + let mut matching = RoaringBitmap::new(); + for range in matching_ranges { + for &node in &row_id_order[range] { + matching.insert(node); + } + } + return Ok(matching); + } + } + matching_nodes_from_sequential_row_ids(filter, |visitor| self.try_for_each_row_id(visitor)) + } + + fn search_with_matching_nodes( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + matching_nodes: &RoaringBitmap, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + if top_k == 0 { + return Ok((Vec::new(), Vec::new())); + } + let candidates = + self.generate_filtered_candidates(query, top_k, l_search, matching_nodes)?; + self.rerank_with_query_scratch(query, &candidates, top_k) + } + + fn generate_filtered_candidates( + &mut self, + query: &[f32], + top_k: usize, + l_search: usize, + matching_nodes: &RoaringBitmap, + ) -> io::Result<Vec<SearchCandidate>> { + let matching_count = usize::try_from(matching_nodes.len()).unwrap_or(usize::MAX); + let strategy = select_filtered_candidate_strategy( + self.header.vector_count as usize, + matching_count, + top_k, + l_search, + self.header.max_degree as usize, + self.read_plan(), + self.adjacency_fully_preloaded(), + ); + match strategy { + FilteredCandidateStrategy::Exhaustive { target_candidates } => { + self.exhaustive_filtered_candidates(query, matching_nodes, target_candidates) + } + FilteredCandidateStrategy::Graph { + target_candidates, + search_list_size, + } => { + self.last_search_stats.filtered_graph_queries = self + .last_search_stats + .filtered_graph_queries + .saturating_add(1); + let mut scratch = std::mem::take(&mut self.query_scratch); + let graph_candidates = (|| { + self.generate_graph_candidates( + query, + search_list_size, + self.options() + .storage_profile + .read_plan() + .filtered_graph_beam_width, + &mut scratch, + )?; + Ok::<_, io::Error>(post_filter_graph_candidates( + &scratch.candidates, + matching_nodes, + target_candidates, + )) + })(); + scratch.recycle_adjacency_windows(); + self.query_scratch = scratch; + if let Some(candidates) = graph_candidates? { + Ok(candidates) + } else { + self.last_search_stats.filtered_graph_fallbacks = self + .last_search_stats + .filtered_graph_fallbacks + .saturating_add(1); + self.exhaustive_filtered_candidates(query, matching_nodes, target_candidates) + } + } + } + } + + fn pq_distance(&self, node: usize, distance_table: &[f32]) -> io::Result<f32> { + let pq = self.pq()?; + let code_size = pq.code_size(); + let start = node + .checked_mul(code_size) + .ok_or_else(|| invalid_data("DiskANN PQ code offset overflows"))?; + let end = start + code_size; + let codes = self + .pq_codes()? + .get(start..end) + .ok_or_else(|| invalid_data("DiskANN PQ codes are truncated"))?; + Ok(pq.distance_from_table(distance_table, codes)) + } + + fn rerank( + &mut self, + query: &[f32], + candidates: &[SearchCandidate], + top_k: usize, + window_cache: &mut VectorWindowCache, + window_buffers: &mut WindowBufferPool, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let record_size = self.header.vector_record_size as usize; + let encoding = self.header.raw_vector_encoding(); + let planner = + VectorWindowPlanner::new(self.read_plan(), self.header.sections.vectors, record_size)?; + let windows = planner.plan_nodes(candidates.iter().map(|candidate| candidate.node)); + self.last_search_stats.rerank_candidate_references = self + .last_search_stats + .rerank_candidate_references + .saturating_add(candidates.len()); + self.last_search_stats.rerank_unique_windows = self + .last_search_stats + .rerank_unique_windows + .saturating_add(windows.len()); + self.last_search_stats.rerank_chunks = self + .last_search_stats + .rerank_chunks + .saturating_add(usize::from(!windows.is_empty())); + let raw_vector_cache_bytes = self.options().raw_vector_cache_bytes; + let (retain_vector_windows, preparation_evictions) = prepare_vector_window_cache( + &windows, + window_cache, + window_buffers, + raw_vector_cache_bytes, + ); + let distance_kernel = selected_raw_vector_distance_kernel(query.len()); + let metric = self.header.metric_type(); + self.last_search_stats.raw_vector_cache_evictions = self + .last_search_stats + .raw_vector_cache_evictions + .saturating_add(preparation_evictions); + let cache_load = self.load_vector_windows(&windows, window_cache, window_buffers)?; + self.last_search_stats.raw_vector_cache_hits = self + .last_search_stats + .raw_vector_cache_hits + .saturating_add(cache_load.hits); + self.last_search_stats.raw_vector_cache_misses = self + .last_search_stats + .raw_vector_cache_misses + .saturating_add(cache_load.misses); + self.last_search_stats.raw_vector_cache_evictions = self + .last_search_stats + .raw_vector_cache_evictions + .saturating_add(cache_load.evictions); + let result = (|| { + let mut exact = BinaryHeap::new(); + exact + .try_reserve(top_k.min(candidates.len())) + .map_err(|_| invalid_input("DiskANN exact result allocation failed"))?; + for candidate in candidates { + let window = planner + .window_for_node(candidate.node) + .ok_or_else(|| invalid_data("DiskANN raw-vector record is out of range"))?; + let payload = window_cache + .get(&window.offset) + .ok_or_else(|| invalid_data("DiskANN vector window is not loaded"))?; + let record = planner.record(window, payload, candidate.node)?; + let distance = + raw_vector_distance(query, record, encoding, metric, distance_kernel)?; + let row_id = self.row_id(candidate.node)?; + push_bounded_exact_result( + &mut exact, + ExactSearchResult { row_id, distance }, + top_k, + )?; + } + let exact = exact.into_sorted_vec(); + let mut ids = exact.iter().map(|result| result.row_id).collect::<Vec<_>>(); + let mut distances = exact + .iter() + .map(|result| result.distance) + .collect::<Vec<_>>(); + ids.resize(top_k, -1); + distances.resize(top_k, f32::MAX); + Ok((ids, distances)) + })(); + if retain_vector_windows && result.is_ok() { + window_cache.touch_windows(&windows); + let evictions = window_cache.trim(window_buffers, raw_vector_cache_bytes); + self.last_search_stats.raw_vector_cache_evictions = self + .last_search_stats + .raw_vector_cache_evictions + .saturating_add(evictions); + } else { + window_cache.recycle(window_buffers); + } + result + } + + fn rerank_interleaved( + &mut self, + query: &[f32], + candidates: &[SearchCandidate], + top_k: usize, + adjacency_windows: &mut AdjacencyWindowCache, + loaded_pages: &mut HashSet<usize>, + window_buffers: &mut WindowBufferPool, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let nodes = candidates + .iter() + .map(|candidate| candidate.node) + .collect::<Vec<_>>(); + self.load_adjacency_pages(&nodes, adjacency_windows, loaded_pages, window_buffers)?; + let planner = ReadWindowPlanner::new(self.read_plan(), self.header.sections.adjacency); + let windows = nodes + .iter() + .map(|&node| { + let page = self.adjacency_locator(node)?.page_index as usize; + planner + .window_for_logical_page(page) + .map(|window| window.offset) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range")) + }) + .collect::<io::Result<BTreeSet<_>>>()?; + self.last_search_stats.rerank_candidate_references = self + .last_search_stats + .rerank_candidate_references + .saturating_add(candidates.len()); + self.last_search_stats.rerank_unique_windows = self + .last_search_stats + .rerank_unique_windows + .saturating_add(windows.len()); + self.last_search_stats.rerank_chunks = self + .last_search_stats + .rerank_chunks + .saturating_add(usize::from(!windows.is_empty())); + let record_size = self.header.vector_record_size as usize; + let encoding = self.header.raw_vector_encoding(); + let distance_kernel = selected_raw_vector_distance_kernel(query.len()); + let metric = self.header.metric_type(); + let mut exact = BinaryHeap::new(); + exact + .try_reserve(top_k.min(candidates.len())) + .map_err(|_| invalid_input("DiskANN exact result allocation failed"))?; + for candidate in candidates { + let locator = self.adjacency_locator(candidate.node)?; + let page_index = locator.page_index as usize; + let window = planner + .window_for_logical_page(page_index) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range"))?; + let payload = if let Some(hot) = self.hot_adjacency_window(window.offset, window.length) + { + hot + } else { + adjacency_windows + .get(&window.offset) + .map(WindowPayload::as_slice) + .ok_or_else(|| invalid_data("DiskANN adjacency rerank window is not loaded"))? + }; + let page_offset = self.header.sections.adjacency.offset + + page_index as u64 * DISKANN_PAGE_SIZE as u64 + - window.offset; + let record_offset = (page_offset as usize) + .checked_add(locator.byte_offset as usize) + .and_then(|offset| offset.checked_sub(record_size)) + .ok_or_else(|| invalid_data("DiskANN interleaved vector offset underflows"))?; + let record = payload + .get(record_offset..record_offset + record_size) + .ok_or_else(|| invalid_data("DiskANN interleaved raw vector is truncated"))?; + let distance = raw_vector_distance(query, record, encoding, metric, distance_kernel)?; + push_bounded_exact_result( + &mut exact, + ExactSearchResult { + row_id: self.row_id(candidate.node)?, + distance, + }, + top_k, + )?; + } + let exact = exact.into_sorted_vec(); + let mut ids = exact.iter().map(|result| result.row_id).collect::<Vec<_>>(); + let mut distances = exact + .iter() + .map(|result| result.distance) + .collect::<Vec<_>>(); + ids.resize(top_k, -1); + distances.resize(top_k, f32::MAX); + Ok((ids, distances)) + } + + fn rerank_with_query_scratch( + &mut self, + query: &[f32], + candidates: &[SearchCandidate], + top_k: usize, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + let mut scratch = std::mem::take(&mut self.query_scratch); + let result = { + scratch.begin_rerank(); + if self.header.is_interleaved() { + self.rerank_interleaved( + query, + candidates, + top_k, + &mut scratch.adjacency_windows, + &mut scratch.loaded_adjacency_pages, + &mut scratch.window_buffers, + ) + } else { + self.rerank( + query, + candidates, + top_k, + &mut scratch.vector_windows, + &mut scratch.window_buffers, + ) + } + }; + scratch.recycle_adjacency_windows(); + self.query_scratch = scratch; + result + } + + fn rerank_candidate_batch_streaming( + &mut self, + queries: &[f32], + top_k: usize, + partitions: Vec<CandidatePartition>, + ) -> io::Result<(Vec<i64>, Vec<f32>)> { + if self.header.is_interleaved() { + return Err(invalid_data( + "DiskANN interleaved rerank must run inside a search session", + )); + } + let dimension = self.header.dimension as usize; + let query_count = queries.len() / dimension; + let distance_kernel = selected_raw_vector_distance_kernel(dimension); + let metric = self.header.metric_type(); + let mut candidate_sets = (0..query_count).map(|_| None).collect::<Vec<_>>(); + for (query_index, candidates) in partitions.into_iter().flatten() { + let slot = candidate_sets + .get_mut(query_index) + .ok_or_else(|| invalid_data("DiskANN filtered batch query index is invalid"))?; + if slot.replace(candidates).is_some() { + return Err(invalid_data( + "DiskANN filtered batch contains duplicate query results", + )); + } + } + if candidate_sets.iter().any(Option::is_none) { + return Err(invalid_data( + "DiskANN filtered batch is missing query candidates", + )); + } + + let record_size = self.header.vector_record_size as usize; + let encoding = self.header.raw_vector_encoding(); + let planner = + VectorWindowPlanner::new(self.read_plan(), self.header.sections.vectors, record_size)?; + let mut grouped = HashMap::<u64, (ReadWindow, Vec<(usize, usize)>)>::new(); + for (query_index, candidates) in candidate_sets.into_iter().enumerate() { + for node in candidates.expect("validated DiskANN batch candidates") { + let window = planner + .window_for_node(node) + .ok_or_else(|| invalid_data("DiskANN raw-vector record is out of range"))?; + grouped + .entry(window.offset) + .or_insert_with(|| (window, Vec::new())) + .1 + .push((query_index, node)); + } + } + let mut window_groups = grouped.into_values().collect::<Vec<_>>(); + window_groups.sort_unstable_by_key(|(window, _)| window.offset); + let windows = window_groups + .iter() + .map(|(window, _)| *window) + .collect::<Vec<_>>(); + let chunks = plan_streaming_window_chunks(&windows); + self.last_search_stats.rerank_candidate_references = self + .last_search_stats + .rerank_candidate_references + .saturating_add( + window_groups + .iter() + .map(|(_, references)| references.len()) + .sum::<usize>(), + ); + self.last_search_stats.rerank_unique_windows = self + .last_search_stats + .rerank_unique_windows + .saturating_add(windows.len()); + self.last_search_stats.rerank_chunks = self + .last_search_stats + .rerank_chunks + .saturating_add(chunks.len()); + let mut exact_heaps = (0..query_count) + .map(|_| BinaryHeap::new()) + .collect::<Vec<BinaryHeap<ExactSearchResult>>>(); + let raw_vector_cache_bytes = self.options().raw_vector_cache_bytes; + let mut scratch = std::mem::take(&mut self.query_scratch); + scratch.begin_rerank(); + let result = (|| { + for chunk in chunks { + let chunk_windows = &windows[chunk.clone()]; + let cache_load = self.load_vector_windows( + chunk_windows, + &mut scratch.vector_windows, + &mut scratch.window_buffers, + )?; + self.last_search_stats.raw_vector_cache_hits = self + .last_search_stats + .raw_vector_cache_hits + .saturating_add(cache_load.hits); + self.last_search_stats.raw_vector_cache_misses = self + .last_search_stats + .raw_vector_cache_misses + .saturating_add(cache_load.misses); + self.last_search_stats.raw_vector_cache_evictions = self + .last_search_stats + .raw_vector_cache_evictions + .saturating_add(cache_load.evictions); + let chunk_reference_count = window_groups[chunk.clone()] + .iter() + .map(|(_, references)| references.len()) + .sum::<usize>(); + let parallel = rayon::current_num_threads() > 1 + && chunk_reference_count.saturating_mul(dimension) + >= PARALLEL_EXACT_RERANK_MIN_COMPONENTS; + if parallel { + self.last_search_stats.parallel_exact_rerank_chunks = self + .last_search_stats + .parallel_exact_rerank_chunks + .saturating_add(1); + self.last_search_stats.parallel_exact_rerank_references = self + .last_search_stats + .parallel_exact_rerank_references + .saturating_add(chunk_reference_count); + let mut rerank_references = Vec::with_capacity(chunk_reference_count); + for (window, references) in &window_groups[chunk.clone()] { + let payload = scratch + .vector_windows + .get(&window.offset) + .ok_or_else(|| invalid_data("DiskANN vector window is not loaded"))?; + for &(query_index, node) in references { + rerank_references.push(ExactRerankReference { + query_index, + row_id: self.row_id(node)?, + record: planner.record(*window, payload, node)?, + }); + } + } + let exact_results = rerank_references + .par_iter() + .map(|reference| { + let query = &queries[reference.query_index * dimension + ..(reference.query_index + 1) * dimension]; + Ok::<_, io::Error>(( + reference.query_index, + ExactSearchResult { + row_id: reference.row_id, + distance: raw_vector_distance( + query, + reference.record, + encoding, + metric, + distance_kernel, + )?, + }, + )) + }) + .collect::<io::Result<Vec<_>>>()?; + for (query_index, exact) in exact_results { + push_bounded_exact_result(&mut exact_heaps[query_index], exact, top_k)?; + } + } else { + for (window, references) in &window_groups[chunk.clone()] { + let payload = scratch + .vector_windows + .get(&window.offset) + .ok_or_else(|| invalid_data("DiskANN vector window is not loaded"))?; + for &(query_index, node) in references { + let query = + &queries[query_index * dimension..(query_index + 1) * dimension]; + push_bounded_exact_result( + &mut exact_heaps[query_index], + ExactSearchResult { + row_id: self.row_id(node)?, + distance: raw_vector_distance( + query, + planner.record(*window, payload, node)?, + encoding, + metric, + distance_kernel, + )?, + }, + top_k, + )?; + } + } + } + scratch.vector_windows.touch_windows(chunk_windows); + let evictions = scratch + .vector_windows + .trim(&mut scratch.window_buffers, raw_vector_cache_bytes); + self.last_search_stats.raw_vector_cache_evictions = self + .last_search_stats + .raw_vector_cache_evictions + .saturating_add(evictions); + } + + let mut ids = Vec::with_capacity(query_count.saturating_mul(top_k)); + let mut distances = Vec::with_capacity(query_count.saturating_mul(top_k)); + for heap in exact_heaps { + let exact = heap.into_sorted_vec(); + ids.extend(exact.iter().map(|result| result.row_id)); + distances.extend(exact.iter().map(|result| result.distance)); + ids.resize(ids.len() + top_k - exact.len(), -1); + distances.resize(distances.len() + top_k - exact.len(), f32::MAX); + } + Ok((ids, distances)) + })(); + if result.is_err() { + scratch.recycle_window_caches(); + } else { + scratch.recycle_adjacency_windows(); + } + self.query_scratch = scratch; + result + } + + fn load_adjacency_pages( + &mut self, + nodes: &[usize], + window_cache: &mut AdjacencyWindowCache, + loaded_pages: &mut HashSet<usize>, + window_buffers: &mut WindowBufferPool, + ) -> io::Result<()> { + let planner = ReadWindowPlanner::new(self.read_plan(), self.header.sections.adjacency); + let mut pages = BTreeSet::new(); + let mut required_windows = Vec::with_capacity(nodes.len()); + for &node in nodes { + let page = self.adjacency_locator(node)?.page_index as usize; + let window = planner + .window_for_logical_page(page) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range"))?; + let window_is_hot = self + .hot_adjacency_window(window.offset, window.length) + .is_some(); + if !window_is_hot { + required_windows.push(window); + } + let window_is_available = window_is_hot || window_cache.contains_key(&window.offset); + if !loaded_pages.contains(&page) || !window_is_available { + pages.insert(page); + } + } + if pages.is_empty() { + return Ok(()); + } + let windows = planner.plan_logical_pages(pages.iter().copied()); + let cold_windows = windows + .iter() + .copied() + .filter(|window| { + self.hot_adjacency_window(window.offset, window.length) + .is_none() + }) + .collect::<Vec<_>>(); + let incoming_bytes = cold_windows + .iter() + .filter(|window| !window_cache.contains_key(&window.offset)) + .fold(0usize, |total, window| total.saturating_add(window.length)); + let preparation_evictions = prepare_adjacency_window_cache( + &required_windows, + incoming_bytes, + window_cache, + window_buffers, + ); + self.last_search_stats.query_adjacency_cache_evictions = self + .last_search_stats + .query_adjacency_cache_evictions + .saturating_add(preparation_evictions); + self.load_adjacency_windows(&cold_windows, window_cache, window_buffers)?; + self.last_search_stats.query_adjacency_cache_peak_bytes = self + .last_search_stats + .query_adjacency_cache_peak_bytes + .max(window_cache.retained_capacity()); + for page_index in pages { + let window = planner + .window_for_logical_page(page_index) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range"))?; + let payload = if let Some(hot) = self.hot_adjacency_window(window.offset, window.length) + { + hot + } else { + window_cache + .get(&window.offset) + .map(WindowPayload::as_slice) + .ok_or_else(|| { + invalid_data("DiskANN adjacency validation window is not loaded") + })? + }; + let page_offset = self.header.sections.adjacency.offset + + page_index as u64 * DISKANN_PAGE_SIZE as u64 + - window.offset; + let page_start = page_offset as usize; + let page_end = page_start + DISKANN_PAGE_SIZE as usize; + self.validate_adjacency_page( + page_index, + payload + .get(page_start..page_end) + .ok_or_else(|| invalid_data("DiskANN adjacency page is truncated"))?, + )?; + loaded_pages.insert(page_index); + } + Ok(()) + } + + fn decode_adjacency_neighbors( + &self, + node: usize, + window_cache: &AdjacencyWindowCache, + neighbors: &mut Vec<u32>, + ) -> io::Result<()> { + let locator = self.adjacency_locator(node)?; + let page_index = locator.page_index as usize; + let planner = ReadWindowPlanner::new(self.read_plan(), self.header.sections.adjacency); + let window = planner + .window_for_logical_page(page_index) + .ok_or_else(|| invalid_data("DiskANN adjacency page is out of range"))?; + let payload = if let Some(hot) = self.hot_adjacency_window(window.offset, window.length) { + hot + } else { + window_cache + .get(&window.offset) + .map(WindowPayload::as_slice) + .ok_or_else(|| invalid_data("DiskANN adjacency decode window is not loaded"))? + }; + let page_offset = self.header.sections.adjacency.offset + + page_index as u64 * DISKANN_PAGE_SIZE as u64 + - window.offset; + let start = page_offset as usize + locator.byte_offset as usize; + let bytes = payload + .get(start..) + .ok_or_else(|| invalid_data("DiskANN adjacency list is truncated"))?; + decode_adjacency_list(bytes, locator.degree(), locator.encoding(), neighbors)?; + Ok(()) + } + + fn load_adjacency_windows( + &mut self, + windows: &[ReadWindow], + local_cache: &mut AdjacencyWindowCache, + window_buffers: &mut WindowBufferPool, + ) -> io::Result<()> { + let mut pending = windows + .iter() + .copied() + .filter(|window| !local_cache.contains_key(&window.offset)) + .collect::<Vec<_>>(); + if pending.is_empty() { + return Ok(()); + } + if self.options().adjacency_cache_bytes == 0 { + self.last_search_stats.adjacency_cache_misses = self + .last_search_stats + .adjacency_cache_misses + .saturating_add(pending.len()); + let mut payloads = Vec::with_capacity(pending.len()); + for window in &pending { + payloads.push(window_buffers.take(window.length)?); + } + let read_result = { + let mut requests = pending + .iter() + .zip(payloads.iter_mut()) + .map(|(window, payload)| ReadRequest::new(window.offset, payload)) + .collect::<Vec<_>>(); + self.pread_ranges(&mut requests) + }; + if let Err(error) = read_result { + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + for (window, payload) in pending.into_iter().zip(payloads) { + local_cache.insert(window.offset, WindowPayload::Owned(payload)); + } + return Ok(()); + } + + while !pending.is_empty() { + let mut reserved = Vec::new(); + let mut waiting = Vec::new(); + for window in &pending { + let (lookup, lock_metrics) = self + .adjacency_cache()? + .lookup_or_reserve(window.offset, window.length)?; + self.last_search_stats + .record_adjacency_cache_lock(lock_metrics); + match lookup { + SharedWindowCacheLookup::Hit(payload) => { + local_cache.insert(window.offset, WindowPayload::Shared(payload)); + self.last_search_stats.adjacency_cache_hits = self + .last_search_stats + .adjacency_cache_hits + .saturating_add(1); + } + SharedWindowCacheLookup::Reserved => { + reserved.push(*window); + self.last_search_stats.adjacency_cache_misses = self + .last_search_stats + .adjacency_cache_misses + .saturating_add(1); + } + SharedWindowCacheLookup::Loading => { + waiting.push(*window); + self.last_search_stats.adjacency_cache_waits = self + .last_search_stats + .adjacency_cache_waits + .saturating_add(1); + } + } + } + + if !reserved.is_empty() { + let mut payloads = Vec::with_capacity(reserved.len()); + for window in &reserved { + match window_buffers.take(window.length) { + Ok(payload) => payloads.push(payload), + Err(error) => { + let lock_metrics = self.adjacency_cache()?.cancel( + &reserved + .iter() + .map(|window| window.offset) + .collect::<Vec<_>>(), + )?; + self.last_search_stats + .record_adjacency_cache_lock(lock_metrics); + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + } + } + let read_result = { + let mut requests = reserved + .iter() + .zip(payloads.iter_mut()) + .map(|(window, payload)| ReadRequest::new(window.offset, payload)) + .collect::<Vec<_>>(); + self.pread_ranges(&mut requests) + }; + if let Err(error) = read_result { + let lock_metrics = self.adjacency_cache()?.cancel( + &reserved + .iter() + .map(|window| window.offset) + .collect::<Vec<_>>(), + )?; + self.last_search_stats + .record_adjacency_cache_lock(lock_metrics); + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + for (window, payload) in reserved.into_iter().zip(payloads) { + let payload = share_window_payload(payload); + local_cache.insert(window.offset, WindowPayload::Shared(Arc::clone(&payload))); + let (evictions, lock_metrics) = + self.adjacency_cache()?.publish(window.offset, payload)?; + self.last_search_stats + .record_adjacency_cache_lock(lock_metrics); + self.last_search_stats.adjacency_cache_evictions = self + .last_search_stats + .adjacency_cache_evictions + .saturating_add(evictions); + } + } + + for window in waiting { + let (payload, lock_metrics) = self + .adjacency_cache()? + .wait_for(window.offset, window.length)?; + self.last_search_stats + .record_adjacency_cache_lock(lock_metrics); + if let Some(payload) = payload { + local_cache.insert(window.offset, WindowPayload::Shared(payload)); + self.last_search_stats.adjacency_cache_hits = self + .last_search_stats + .adjacency_cache_hits + .saturating_add(1); + } + } + pending.retain(|window| !local_cache.contains_key(&window.offset)); + } + Ok(()) + } + + fn load_vector_windows( + &mut self, + windows: &[ReadWindow], + cache: &mut VectorWindowCache, + window_buffers: &mut WindowBufferPool, + ) -> io::Result<VectorWindowLoadStats> { + let mut pending = windows + .iter() + .copied() + .filter(|window| !cache.contains_key(&window.offset)) + .collect::<Vec<_>>(); + let mut stats = VectorWindowLoadStats { + hits: windows.len().saturating_sub(pending.len()), + ..VectorWindowLoadStats::default() + }; + if pending.is_empty() { + return Ok(stats); + } + + if self.options().raw_vector_cache_bytes == 0 { + stats.misses = pending.len(); + let mut payloads = Vec::with_capacity(pending.len()); + for window in &pending { + payloads.push(window_buffers.take(window.length)?); + } + let read_result = { + let mut requests = pending + .iter() + .zip(payloads.iter_mut()) + .map(|(window, payload)| ReadRequest::new(window.offset, payload)) + .collect::<Vec<_>>(); + self.pread_ranges(&mut requests) + }; + if let Err(error) = read_result { + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + for (window, payload) in pending.into_iter().zip(payloads) { + cache.insert(window.offset, payload); + } + return Ok(stats); + } + + while !pending.is_empty() { + let mut reserved = Vec::new(); + let mut waiting = Vec::new(); + for window in &pending { + match self + .raw_vector_cache()? + .lookup_or_reserve(window.offset, window.length)? + .0 + { + SharedWindowCacheLookup::Hit(payload) => { + cache.insert(window.offset, WindowPayload::Shared(payload)); + stats.hits = stats.hits.saturating_add(1); + } + SharedWindowCacheLookup::Reserved => { + reserved.push(*window); + stats.misses = stats.misses.saturating_add(1); + } + SharedWindowCacheLookup::Loading => waiting.push(*window), + } + } + + if !reserved.is_empty() { + let mut payloads = Vec::with_capacity(reserved.len()); + for window in &reserved { + match window_buffers.take(window.length) { + Ok(payload) => payloads.push(payload), + Err(error) => { + self.raw_vector_cache()?.cancel( + &reserved + .iter() + .map(|window| window.offset) + .collect::<Vec<_>>(), + )?; + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + } + } + let read_result = { + let mut requests = reserved + .iter() + .zip(payloads.iter_mut()) + .map(|(window, payload)| ReadRequest::new(window.offset, payload)) + .collect::<Vec<_>>(); + self.pread_ranges(&mut requests) + }; + if let Err(error) = read_result { + self.raw_vector_cache()?.cancel( + &reserved + .iter() + .map(|window| window.offset) + .collect::<Vec<_>>(), + )?; + for payload in payloads { + window_buffers.recycle(payload); + } + return Err(error); + } + for (window, payload) in reserved.into_iter().zip(payloads) { + let payload = share_window_payload(payload); + cache.insert(window.offset, WindowPayload::Shared(Arc::clone(&payload))); + stats.evictions = stats.evictions.saturating_add( + self.raw_vector_cache()?.publish(window.offset, payload)?.0, + ); + } + } + + for window in waiting { + if let Some(payload) = self + .raw_vector_cache()? + .wait_for(window.offset, window.length)? + .0 + { + cache.insert(window.offset, WindowPayload::Shared(payload)); + stats.hits = stats.hits.saturating_add(1); + } + } + pending.retain(|window| !cache.contains_key(&window.offset)); + } + Ok(stats) + } +} + +fn sort_candidates(candidates: &mut [SearchCandidate]) { + candidates.sort_by(|left, right| { + left.distance + .total_cmp(&right.distance) + .then_with(|| left.node.cmp(&right.node)) + }); +} + +fn desired_filtered_candidate_count(matching_count: usize, top_k: usize) -> usize { + matching_count.min(top_k.saturating_mul(4).max(64)) +} + +fn resolve_diskann_l_search(top_k: usize, l_search: usize) -> usize { + let configured = if l_search == 0 { + top_k.saturating_mul(2).max(100) + } else { + l_search + }; + top_k.max(configured) +} + +fn topk_result_stability(left: &[i64], right: &[i64], top_k: usize) -> f32 { + if top_k == 0 || left.len() != right.len() || !left.len().is_multiple_of(top_k) { + return 0.0; + } + let mut overlap = 0usize; + let mut denominator = 0usize; + for (left_query, right_query) in left.chunks_exact(top_k).zip(right.chunks_exact(top_k)) { Review Comment: Calibration mishandles two supported row-ID cases. First, every negative row ID is treated as padding: DiskANN's codec supports the full `i64` range (the new golden fixture even stores `i64::MIN`), but this loop only counts `row_id >= 0`. Two completely disjoint negative-ID result sets therefore report stability `1.0` because the denominator stays zero; I reproduced this with `[-2, -3]` versus `[-4, -5]`. Second, duplicate multiplicity is overcounted because each left occurrence independently uses `right_query.contains`: `[7, 7]` versus `[7, 8]` also reports `1.0`, although one Top-K slot changed, and the new filtered-search tests explicitly support duplicate row IDs. `calibrate_l_search` can therefore select the smallest width while its actual result rows are still changing. Could we distinguish padding from valid negative IDs and compute a multiset-aware overlap? -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
