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(());

Reply via email to