This is an automated email from the ASF dual-hosted git repository.
jerry-024 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/paimon-vector-index.git
The following commit(s) were added to refs/heads/main by this push:
new b36f8a9 ivf: Retry only incomplete queries in automatic batch search
(#72)
b36f8a9 is described below
commit b36f8a9c926c4556a595f407de44e4a65939af6c
Author: shyjsarah <[email protected]>
AuthorDate: Wed Aug 12 11:44:35 2026 +0800
ivf: Retry only incomplete queries in automatic batch search (#72)
Co-authored-by: shaoyijie <[email protected]>
---
core/src/index.rs | 384 +++++++++++++++++++++++++++++++++++++++++++++------
core/src/ivfrq_io.rs | 4 +
2 files changed, 342 insertions(+), 46 deletions(-)
diff --git a/core/src/index.rs b/core/src/index.rs
index f5b35f8..6454bd1 100644
--- a/core/src/index.rs
+++ b/core/src/index.rs
@@ -64,6 +64,7 @@ use std::io::{self, Cursor};
/// dimensions, for example `m=32` at 128 dimensions and `m=240` at 960.
pub const DEFAULT_PQ_CODE_RATIO: f64 = 0.0625;
const PERSISTED_ROW_ID_ESTIMATE_BYTES: usize = 10;
+const MAX_IVF_BATCH_RETRY_BUFFER_BYTES: usize = 64 * 1024 * 1024;
/// Resolve a concrete PQ subquantizer count from a target code/raw byte ratio.
///
@@ -1775,18 +1776,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
let total_vectors = usize::try_from(reader.total_vectors)
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe = params.resolve_ivf_nprobe(reader.nlist,
total_vectors, None)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
total_vectors,
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_ivfflat_reader(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
)
@@ -1797,18 +1799,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
let total_vectors = usize::try_from(reader.total_vectors)
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe = params.resolve_ivf_nprobe(reader.nlist,
total_vectors, None)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
total_vectors,
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_ivfsq_reader(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
)
@@ -1819,18 +1822,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
let total_vectors = usize::try_from(reader.total_vectors)
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe = params.resolve_ivf_nprobe(reader.nlist,
total_vectors, None)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
total_vectors,
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_reader_with_reuse_mode_and_budget(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
params.ivfpq_batch_table_reuse,
@@ -1843,23 +1847,34 @@ impl<R: SeekRead> VectorIndexReader<R> {
let total_vectors = usize::try_from(reader.total_vectors)
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe = params.resolve_ivf_nprobe(reader.nlist,
total_vectors, None)?;
- progressive_ivf_search(
+ let mut aggregate_stats = IVFRQSearchStats::default();
+ let result = progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
total_vectors,
- |nprobe| {
- search_batch_ivfrq_reader(
+ |active_queries, active_query_count, nprobe| {
+ let result = search_batch_ivfrq_reader(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
- )
+ );
+ if result.is_ok() {
+ aggregate_stats.merge(reader.last_search_stats());
+ }
+ result
},
- )
+ );
+ if result.is_ok() {
+ aggregate_stats.query_count = query_count;
+ reader.set_last_search_stats(aggregate_stats);
+ }
+ result
}
Self::DiskAnn(reader) => reader.search_batch(
queries,
@@ -1889,18 +1904,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe =
params.resolve_ivf_nprobe(reader.nlist, total_vectors,
matching_count)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
matching_count.unwrap_or(total_vectors),
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_ivfflat_reader_roaring_filter(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
roaring_filter_bytes,
@@ -1913,18 +1929,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe =
params.resolve_ivf_nprobe(reader.nlist, total_vectors,
matching_count)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
matching_count.unwrap_or(total_vectors),
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_ivfsq_reader_roaring_filter(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
roaring_filter_bytes,
@@ -1937,18 +1954,19 @@ impl<R: SeekRead> VectorIndexReader<R> {
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe =
params.resolve_ivf_nprobe(reader.nlist, total_vectors,
matching_count)?;
- progressive_ivf_search(
+ progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
matching_count.unwrap_or(total_vectors),
- |nprobe| {
+ |active_queries, active_query_count, nprobe| {
search_batch_reader_roaring_filter_with_reuse_mode_and_budget(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
roaring_filter_bytes,
@@ -1963,24 +1981,35 @@ impl<R: SeekRead> VectorIndexReader<R> {
.map_err(|_| invalid_input("negative IVF vector count"))?;
let nprobe =
params.resolve_ivf_nprobe(reader.nlist, total_vectors,
matching_count)?;
- progressive_ivf_search(
+ let mut aggregate_stats = IVFRQSearchStats::default();
+ let result = progressive_ivf_batch_search(
params,
reader.nlist,
nprobe,
+ queries,
query_count,
params.top_k,
matching_count.unwrap_or(total_vectors),
- |nprobe| {
- search_batch_ivfrq_reader_roaring_filter(
+ |active_queries, active_query_count, nprobe| {
+ let result = search_batch_ivfrq_reader_roaring_filter(
reader,
- queries,
- query_count,
+ active_queries,
+ active_query_count,
params.top_k,
nprobe,
roaring_filter_bytes,
- )
+ );
+ if result.is_ok() {
+ aggregate_stats.merge(reader.last_search_stats());
+ }
+ result
},
- )
+ );
+ if result.is_ok() {
+ aggregate_stats.query_count = query_count;
+ reader.set_last_search_stats(aggregate_stats);
+ }
+ result
}
Self::DiskAnn(reader) => reader.search_batch_with_roaring_filter(
queries,
@@ -2020,14 +2049,7 @@ fn progressive_ivf_search(
.1
.chunks_exact(top_k)
.take(query_count)
- .all(|distances| {
- distances
- .iter()
- .filter(|&&distance| distance != f32::MAX)
- .take(required_per_query)
- .count()
- >= required_per_query
- })
+ .all(|distances| ivf_search_result_is_complete(distances,
required_per_query))
{
return Ok(result);
}
@@ -2035,6 +2057,132 @@ fn progressive_ivf_search(
}
}
+/// Runs automatic IVF batch expansion independently for each query.
+///
+/// Queries which already produced the required number of results are removed
from later rounds.
+/// Retried queries are still searched from the first probed list; reusing
work across rounds is
+/// intentionally left to the incremental-search implementation.
+fn progressive_ivf_batch_search(
+ params: VectorSearchParams,
+ nlist: usize,
+ initial_nprobe: usize,
+ queries: &[f32],
+ query_count: usize,
+ top_k: usize,
+ available_matches: usize,
+ search: impl FnMut(&[f32], usize, usize) -> io::Result<(Vec<i64>,
Vec<f32>)>,
+) -> io::Result<(Vec<i64>, Vec<f32>)> {
+ progressive_ivf_batch_search_with_retry_buffer_limit(
+ params,
+ nlist,
+ initial_nprobe,
+ queries,
+ query_count,
+ top_k,
+ available_matches,
+ MAX_IVF_BATCH_RETRY_BUFFER_BYTES,
+ search,
+ )
+}
+
+#[allow(clippy::too_many_arguments)]
+fn progressive_ivf_batch_search_with_retry_buffer_limit(
+ params: VectorSearchParams,
+ nlist: usize,
+ initial_nprobe: usize,
+ queries: &[f32],
+ query_count: usize,
+ top_k: usize,
+ available_matches: usize,
+ retry_buffer_limit_bytes: usize,
+ mut search: impl FnMut(&[f32], usize, usize) -> io::Result<(Vec<i64>,
Vec<f32>)>,
+) -> io::Result<(Vec<i64>, Vec<f32>)> {
+ if params.search_width != SearchWidth::Auto {
+ return search(queries, query_count, initial_nprobe);
+ }
+
+ let dimension = queries.len() / query_count;
+ let required_per_query = top_k.min(available_matches);
+ let mut nprobe = initial_nprobe;
+ let (mut result_ids, mut result_distances) = search(queries, query_count,
nprobe)?;
+ let mut active_queries = result_distances
+ .chunks_exact(top_k)
+ .take(query_count)
+ .enumerate()
+ .filter_map(|(query_index, distances)| {
+ (!ivf_search_result_is_complete(distances,
required_per_query)).then_some(query_index)
+ })
+ .collect::<Vec<_>>();
+
+ loop {
+ if nprobe >= nlist || required_per_query == 0 ||
active_queries.is_empty() {
+ return Ok((result_ids, result_distances));
+ }
+
+ nprobe = nprobe.saturating_mul(2).min(nlist);
+ if active_queries.len() == query_count {
+ drop(result_ids);
+ drop(result_distances);
+ (result_ids, result_distances) = search(queries, query_count,
nprobe)?;
+ active_queries = result_distances
+ .chunks_exact(top_k)
+ .take(query_count)
+ .enumerate()
+ .filter_map(|(query_index, distances)| {
+ (!ivf_search_result_is_complete(distances,
required_per_query))
+ .then_some(query_index)
+ })
+ .collect();
+ continue;
+ }
+
+ let bytes_per_retry_query = dimension
+ .saturating_mul(std::mem::size_of::<f32>())
+ .saturating_add(
+ top_k.saturating_mul(std::mem::size_of::<i64>() +
std::mem::size_of::<f32>()),
+ );
+ let retry_chunk_size = retry_buffer_limit_bytes
+ .checked_div(bytes_per_retry_query)
+ .unwrap_or(active_queries.len())
+ .max(1);
+ let mut next_active_queries = Vec::new();
+ for active_chunk in active_queries.chunks(retry_chunk_size) {
+ let mut packed_queries = Vec::with_capacity(active_chunk.len() *
dimension);
+ for &query_index in active_chunk {
+ let start = query_index * dimension;
+ packed_queries.extend_from_slice(&queries[start..start +
dimension]);
+ }
+
+ let (round_ids, round_distances) = search(&packed_queries,
active_chunk.len(), nprobe)?;
+ for (round_index, &query_index) in active_chunk.iter().enumerate()
{
+ let round_start = round_index * top_k;
+ let result_start = query_index * top_k;
+ result_ids[result_start..result_start + top_k]
+ .copy_from_slice(&round_ids[round_start..round_start +
top_k]);
+ result_distances[result_start..result_start + top_k]
+ .copy_from_slice(&round_distances[round_start..round_start
+ top_k]);
+
+ if !ivf_search_result_is_complete(
+ &round_distances[round_start..round_start + top_k],
+ required_per_query,
+ ) {
+ next_active_queries.push(query_index);
+ }
+ }
+ }
+ active_queries = next_active_queries;
+ }
+}
+
+fn ivf_search_result_is_complete(distances: &[f32], required: usize) -> bool {
+ distances
+ .iter()
+ .filter(|&&distance| distance != f32::MAX)
+ .take(required)
+ .count()
+ >= required
+}
+
fn validate_config(config: &VectorIndexConfig) -> io::Result<()> {
validate_positive(config.dimension(), "dimension")?;
validate_positive(config.nlist(), "nlist")?;
@@ -2440,6 +2588,44 @@ mod tests {
assert!(reader.diskann_search_stats().is_none());
}
+ #[test]
+ fn automatic_filtered_ivfrq_batch_stats_cover_all_progressive_rounds() {
+ let dimension = 64;
+ let nlist = 64;
+ let (mut reader, data) = build_reader(VectorIndexConfig::IvfRq {
+ dimension,
+ nlist,
+ metric: MetricType::L2,
+ bits: 4,
+ });
+ let queries = [0, nlist - 1]
+ .into_iter()
+ .flat_map(|row| data[row * dimension..(row + 1) *
dimension].iter().copied())
+ .collect::<Vec<_>>();
+ let mut filter = RoaringTreemap::new();
+ for row_id in (0..512).step_by(nlist) {
+ filter.insert(row_id as u64);
+ }
+ let mut filter_bytes = Vec::new();
+ filter.serialize_into(&mut filter_bytes).unwrap();
+
+ reader
+ .search_batch_with_roaring_filter(
+ &queries,
+ 2,
+
VectorSearchParams::automatic(8).with_max_initial_filter_expansion_factor(1),
+ &filter_bytes,
+ )
+ .unwrap();
+
+ let stats = reader.ivfrq_search_stats().expect("IVF-RQ diagnostics");
+ assert_eq!(stats.query_count, 2);
+ assert!(
+ stats.scanned_vectors > 2 * 8,
+ "statistics should include work from progressive retries"
+ );
+ }
+
#[test]
fn
diskann_batch_search_overlaps_graph_reads_and_centralizes_filtered_rerank() {
let dimension = 8;
@@ -3066,6 +3252,112 @@ mod tests {
assert_eq!(result.0, vec![7, 8]);
}
+ #[test]
+ fn automatic_batch_search_retries_only_incomplete_queries() {
+ let queries = vec![10.0, 20.0, 30.0];
+ let mut observed = Vec::new();
+ let result = progressive_ivf_batch_search(
+ VectorSearchParams::automatic(2),
+ 8,
+ 2,
+ &queries,
+ 3,
+ 2,
+ 10,
+ |active_queries, active_query_count, nprobe| {
+ observed.push((nprobe, active_query_count,
active_queries.to_vec()));
+ match nprobe {
+ 2 => Ok((
+ vec![100, 101, 200, -1, 300, 301],
+ vec![1.0, 2.0, 1.0, f32::MAX, 1.0, 2.0],
+ )),
+ 4 => Ok((vec![200, -1], vec![1.0, f32::MAX])),
+ 8 => Ok((vec![200, 201], vec![1.0, 2.0])),
+ _ => unreachable!("unexpected nprobe {nprobe}"),
+ }
+ },
+ )
+ .unwrap();
+
+ assert_eq!(
+ observed,
+ vec![
+ (2, 3, vec![10.0, 20.0, 30.0]),
+ (4, 1, vec![20.0]),
+ (8, 1, vec![20.0]),
+ ]
+ );
+ assert_eq!(result.0, vec![100, 101, 200, 201, 300, 301]);
+ assert_eq!(result.1, vec![1.0, 2.0, 1.0, 2.0, 1.0, 2.0]);
+ }
+
+ #[test]
+ fn automatic_batch_search_chunks_partial_retries_to_bound_memory() {
+ let queries = vec![10.0, 11.0, 20.0, 21.0, 30.0, 31.0, 40.0, 41.0];
+ let mut observed = Vec::new();
+ let result = progressive_ivf_batch_search_with_retry_buffer_limit(
+ VectorSearchParams::automatic(2),
+ 4,
+ 2,
+ &queries,
+ 4,
+ 2,
+ 10,
+ 32,
+ |active_queries, active_query_count, nprobe| {
+ observed.push((nprobe, active_query_count,
active_queries.to_vec()));
+ if nprobe == 2 {
+ return Ok((
+ vec![100, 101, 200, -1, 300, -1, 400, -1],
+ vec![1.0, 2.0, 1.0, f32::MAX, 1.0, f32::MAX, 1.0,
f32::MAX],
+ ));
+ }
+
+ let query = active_queries[0] as i64;
+ Ok((vec![query, query + 1], vec![1.0, 2.0]))
+ },
+ )
+ .unwrap();
+
+ assert_eq!(
+ observed,
+ vec![
+ (2, 4, queries),
+ (4, 1, vec![20.0, 21.0]),
+ (4, 1, vec![30.0, 31.0]),
+ (4, 1, vec![40.0, 41.0]),
+ ]
+ );
+ assert_eq!(result.0, vec![100, 101, 20, 21, 30, 31, 40, 41]);
+ }
+
+ #[test]
+ fn fixed_batch_search_runs_once_with_the_full_batch() {
+ let queries = vec![10.0, 20.0, 30.0];
+ let mut observed = Vec::new();
+ let result = progressive_ivf_batch_search(
+ VectorSearchParams::new(2, 4),
+ 8,
+ 4,
+ &queries,
+ 3,
+ 2,
+ 10,
+ |active_queries, active_query_count, nprobe| {
+ observed.push((nprobe, active_query_count,
active_queries.to_vec()));
+ Ok((
+ vec![100, 101, 200, 201, 300, 301],
+ vec![1.0, 2.0, 1.0, 2.0, 1.0, 2.0],
+ ))
+ },
+ )
+ .unwrap();
+
+ assert_eq!(observed, vec![(4, 3, queries)]);
+ assert_eq!(result.0, vec![100, 101, 200, 201, 300, 301]);
+ assert_eq!(result.1, vec![1.0, 2.0, 1.0, 2.0, 1.0, 2.0]);
+ }
+
#[test]
fn capped_automatic_filtered_search_can_expand_past_the_initial_cap() {
let params =
VectorSearchParams::automatic(2).with_max_initial_filter_expansion_factor(4);
diff --git a/core/src/ivfrq_io.rs b/core/src/ivfrq_io.rs
index 5990d49..328a332 100644
--- a/core/src/ivfrq_io.rs
+++ b/core/src/ivfrq_io.rs
@@ -377,6 +377,10 @@ impl<R: SeekRead> IVFRQIndexReader<R> {
self.last_search_stats
}
+ pub(crate) fn set_last_search_stats(&mut self, stats: IVFRQSearchStats) {
+ self.last_search_stats = stats;
+ }
+
pub fn ensure_loaded(&mut self) -> io::Result<()> {
if self.loaded {
return Ok(());