diff --git a/core/src/index.rs b/core/src/index.rs index c2b4d10..39a6efc 100644 --- a/core/src/index.rs +++ b/core/src/index.rs @@ -33,9 +33,10 @@ use crate::ivfflat_io::{ search_batch_ivfflat_reader, search_batch_ivfflat_reader_roaring_filter, write_ivfflat_index, IVFFlatIndexReader, IVFFLAT_MAGIC, }; +pub use crate::ivfpq::IvfPqBatchTableReuseMode; use crate::ivfpq::{ - search_batch_reader, search_batch_reader_roaring_filter, search_with_reader, - search_with_reader_roaring_filter, IVFPQIndex, + search_batch_reader_roaring_filter_with_reuse_mode, search_batch_reader_with_reuse_mode, + search_with_reader, search_with_reader_roaring_filter, IVFPQIndex, }; use crate::ivfrq::IVFRQIndex; use crate::ivfrq_io::{ @@ -988,6 +989,7 @@ pub struct VectorSearchParams { pub top_k: usize, pub search_width: SearchWidth, pub width: usize, + pub ivfpq_batch_table_reuse: IvfPqBatchTableReuseMode, } impl VectorSearchParams { @@ -996,6 +998,7 @@ impl VectorSearchParams { top_k, search_width: SearchWidth::IvfNProbe, width: nprobe, + ivfpq_batch_table_reuse: IvfPqBatchTableReuseMode::Auto, } } @@ -1004,6 +1007,7 @@ impl VectorSearchParams { top_k, search_width: SearchWidth::DiskAnnLSearch, width: l_search, + ivfpq_batch_table_reuse: IvfPqBatchTableReuseMode::Auto, } } @@ -1012,9 +1016,15 @@ impl VectorSearchParams { top_k, search_width: SearchWidth::Auto, width: 0, + ivfpq_batch_table_reuse: IvfPqBatchTableReuseMode::Auto, } } + pub fn with_ivfpq_batch_table_reuse(mut self, mode: IvfPqBatchTableReuseMode) -> Self { + self.ivfpq_batch_table_reuse = mode; + self + } + pub fn configured_ivf_nprobe(self) -> Option { (self.search_width == SearchWidth::IvfNProbe).then_some(self.width) } @@ -1755,7 +1765,14 @@ impl VectorIndexReader { params.top_k, total_vectors, |nprobe| { - search_batch_reader(reader, queries, query_count, params.top_k, nprobe) + search_batch_reader_with_reuse_mode( + reader, + queries, + query_count, + params.top_k, + nprobe, + params.ivfpq_batch_table_reuse, + ) }, ) } @@ -1865,13 +1882,14 @@ impl VectorIndexReader { params.top_k, matching_count.unwrap_or(total_vectors), |nprobe| { - search_batch_reader_roaring_filter( + search_batch_reader_roaring_filter_with_reuse_mode( reader, queries, query_count, params.top_k, nprobe, roaring_filter_bytes, + params.ivfpq_batch_table_reuse, ) }, ) @@ -2870,6 +2888,21 @@ mod tests { .contains("cannot be used with a DiskANN")); } + #[test] + fn ivfpq_batch_table_reuse_is_auto_by_default_and_can_be_disabled() { + let params = VectorSearchParams::new(10, 4); + assert_eq!( + params.ivfpq_batch_table_reuse, + IvfPqBatchTableReuseMode::Auto + ); + assert_eq!( + params + .with_ivfpq_batch_table_reuse(IvfPqBatchTableReuseMode::Off) + .ivfpq_batch_table_reuse, + IvfPqBatchTableReuseMode::Off + ); + } + #[test] fn automatic_filtered_search_expands_until_results_are_filled() { let mut observed = Vec::new(); diff --git a/core/src/ivfpq.rs b/core/src/ivfpq.rs index 0f25512..f0c35ad 100644 --- a/core/src/ivfpq.rs +++ b/core/src/ivfpq.rs @@ -19,7 +19,7 @@ use crate::distance::{ fvec_inner_product, fvec_madd, fvec_normalize, pq_distance_four_codes, pq_distance_from_table, MetricType, }; -use crate::index_io_util::ivf_payload_is_oversized; +use crate::index_io_util::{ivf_payload_is_oversized, MAX_IVF_BATCH_READ_BYTES}; use crate::io::{IVFPQIndexReader, InvertedListPayload, SeekRead}; use crate::kmeans::{self, KMeansConfig}; use crate::opq::OPQMatrix; @@ -914,6 +914,166 @@ fn should_scan_sparse(count: usize, matching_rows: &MatchingRows, divisor: usize matching_rows.len().saturating_mul(divisor) <= count } +fn has_matching_rows(matching_rows: Option<&MatchingRows>) -> bool { + match matching_rows { + Some(rows) => !rows.is_empty(), + None => true, + } +} + +// Below this size, table construction and Rayon scheduling dominate the saved +// per-query/list distance-table work. +const MIN_EPHEMERAL_PRECOMPUTE_QUERIES: usize = 64; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u32)] +pub enum IvfPqBatchTableReuseMode { + Off = 0, + On = 1, + Auto = 2, +} + +fn should_use_ephemeral_precomputation( + matching_list_count: usize, + active_query_count: usize, + probe_count: usize, +) -> bool { + let setup_tables = matching_list_count.saturating_add(active_query_count); + // Require at least 2x reuse over the list-table and query-table setup work. + setup_tables > 0 && probe_count >= setup_tables.saturating_mul(2) +} + +fn ephemeral_precomputed_table_fits_budget( + matching_list_count: usize, + query_scratch_count: usize, + m: usize, + ksub: usize, +) -> bool { + if matching_list_count == 0 { + return false; + } + matching_list_count + .checked_add(1) + .and_then(|tables| tables.checked_add(query_scratch_count)) + .and_then(|tables| tables.checked_mul(m)) + .and_then(|values| values.checked_mul(ksub)) + .and_then(|values| values.checked_mul(std::mem::size_of::())) + .is_some_and(|bytes| bytes <= MAX_IVF_BATCH_READ_BYTES) +} + +#[cfg(test)] +fn fill_list_precomputed_table( + coarse_centroid: &[f32], + pq: &ProductQuantizer, + pq_norms: &[f32], + table: &mut Vec, +) { + debug_assert_eq!(coarse_centroid.len(), pq.d); + debug_assert_eq!(pq_norms.len(), pq.m * pq.ksub); + table.resize(pq.m * pq.ksub, 0.0); + for sub in 0..pq.m { + let range = pq.chunk_range(sub); + let chunk_dim = range.len(); + let pq_base = range.start * pq.ksub; + for code in 0..pq.ksub { + let pq_offset = pq_base + code * chunk_dim; + let mut inner_product = 0.0f32; + for dimension in 0..chunk_dim { + inner_product += + coarse_centroid[range.start + dimension] * pq.centroids[pq_offset + dimension]; + } + let table_offset = sub * pq.ksub + code; + table[table_offset] = pq_norms[table_offset] + 2.0 * inner_product; + } + } +} + +fn compute_stable_ephemeral_pq_norms(pq: &ProductQuantizer) -> Vec { + let mut norms = vec![0.0f64; pq.m * pq.ksub]; + for sub in 0..pq.m { + let range = pq.chunk_range(sub); + let chunk_dim = range.len(); + let pq_base = range.start * pq.ksub; + for code in 0..pq.ksub { + let pq_offset = pq_base + code * chunk_dim; + norms[sub * pq.ksub + code] = (0..chunk_dim) + .map(|dimension| { + let value = f64::from(pq.centroids[pq_offset + dimension]); + value * value + }) + .sum(); + } + } + norms +} + +fn fill_stable_ephemeral_list_table( + coarse_centroid: &[f32], + pq: &ProductQuantizer, + pq_norms: &[f64], + table: &mut Vec, +) { + table.resize(pq.m * pq.ksub, 0.0); + for sub in 0..pq.m { + let range = pq.chunk_range(sub); + let chunk_dim = range.len(); + let pq_base = range.start * pq.ksub; + for code in 0..pq.ksub { + let pq_offset = pq_base + code * chunk_dim; + let mut inner_product = 0.0f64; + for dimension in 0..chunk_dim { + let pq_value = f64::from(pq.centroids[pq_offset + dimension]); + inner_product += f64::from(coarse_centroid[range.start + dimension]) * pq_value; + } + let offset = sub * pq.ksub + code; + table[offset] = pq_norms[offset] + 2.0 * inner_product; + } + } +} + +fn fill_stable_ephemeral_query_table(query: &[f32], pq: &ProductQuantizer, table: &mut Vec) { + table.resize(pq.m * pq.ksub, 0.0); + for sub in 0..pq.m { + let range = pq.chunk_range(sub); + let chunk_dim = range.len(); + let pq_base = range.start * pq.ksub; + for code in 0..pq.ksub { + let pq_offset = pq_base + code * chunk_dim; + let mut inner_product = 0.0f64; + for dimension in 0..chunk_dim { + inner_product += f64::from(query[range.start + dimension]) + * f64::from(pq.centroids[pq_offset + dimension]); + } + table[sub * pq.ksub + code] = inner_product; + } + } +} + +fn combine_stable_ephemeral_tables( + list_table: &[f64], + query_table: &[f64], + query: &[f32], + coarse_centroid: &[f32], + pq: &ProductQuantizer, + sim_table: &mut Vec, +) { + sim_table.resize(pq.m * pq.ksub, 0.0); + for sub in 0..pq.m { + let range = pq.chunk_range(sub); + let mut residual_norm = 0.0f64; + for dimension in range { + let residual = f64::from(query[dimension]) - f64::from(coarse_centroid[dimension]); + residual_norm += residual * residual; + } + let table_base = sub * pq.ksub; + for code in 0..pq.ksub { + let offset = table_base + code; + sim_table[offset] = + (residual_norm + list_table[offset] - 2.0 * query_table[offset]).max(0.0) as f32; + } + } +} + /// Scan 4-bit packed codes using u8-domain accumulation. fn scan_codes_4bit( sim_table: &[f32], @@ -1166,6 +1326,7 @@ struct ReaderSearchContext<'a> { #[derive(Default)] struct ReaderScanScratch { sim_table: Vec, + ip_table: Vec, distances: Vec, } @@ -1522,7 +1683,25 @@ pub fn search_batch_reader( k: usize, nprobe: usize, ) -> io::Result<(Vec, Vec)> { - search_batch_reader_filter(reader, queries, nq, k, nprobe, None) + search_batch_reader_with_reuse_mode( + reader, + queries, + nq, + k, + nprobe, + IvfPqBatchTableReuseMode::Auto, + ) +} + +pub fn search_batch_reader_with_reuse_mode( + reader: &mut IVFPQIndexReader, + queries: &[f32], + nq: usize, + k: usize, + nprobe: usize, + reuse_mode: IvfPqBatchTableReuseMode, +) -> io::Result<(Vec, Vec)> { + search_batch_reader_filter_with_reuse_mode(reader, queries, nq, k, nprobe, None, reuse_mode) } /// Big batch search with an optional row-id filter. @@ -1533,6 +1712,70 @@ pub fn search_batch_reader_filter( k: usize, nprobe: usize, filter: Option<&dyn RowIdFilter>, +) -> io::Result<(Vec, Vec)> { + search_batch_reader_filter_with_reuse_mode( + reader, + queries, + nq, + k, + nprobe, + filter, + IvfPqBatchTableReuseMode::Auto, + ) +} + +pub fn search_batch_reader_filter_with_reuse_mode( + reader: &mut IVFPQIndexReader, + queries: &[f32], + nq: usize, + k: usize, + nprobe: usize, + filter: Option<&dyn RowIdFilter>, + reuse_mode: IvfPqBatchTableReuseMode, +) -> io::Result<(Vec, Vec)> { + search_batch_reader_filter_with_reuse_mode_and_observer( + reader, + queries, + nq, + k, + nprobe, + filter, + reuse_mode, + |_| {}, + ) +} + +#[cfg(test)] +fn search_batch_reader_filter_with_observer( + reader: &mut IVFPQIndexReader, + queries: &[f32], + nq: usize, + k: usize, + nprobe: usize, + filter: Option<&dyn RowIdFilter>, + mut observe_ephemeral_precomputed_lists: impl FnMut(usize), +) -> io::Result<(Vec, Vec)> { + search_batch_reader_filter_with_reuse_mode_and_observer( + reader, + queries, + nq, + k, + nprobe, + filter, + IvfPqBatchTableReuseMode::Auto, + &mut observe_ephemeral_precomputed_lists, + ) +} + +fn search_batch_reader_filter_with_reuse_mode_and_observer( + reader: &mut IVFPQIndexReader, + queries: &[f32], + nq: usize, + k: usize, + nprobe: usize, + filter: Option<&dyn RowIdFilter>, + reuse_mode: IvfPqBatchTableReuseMode, + mut observe_ephemeral_precomputed_lists: impl FnMut(usize), ) -> io::Result<(Vec, Vec)> { reader.ensure_loaded()?; let d = reader.d; @@ -1613,9 +1856,19 @@ pub fn search_batch_reader_filter( } unique_lists.sort_unstable_by_key(|&list_id| reader.list_offsets[list_id]); - let use_precomputed = - metric == MetricType::L2 && by_residual && !reader.precomputed_table.is_empty(); - + let use_precomputed = reuse_mode != IvfPqBatchTableReuseMode::Off + && metric == MetricType::L2 + && by_residual + && !reader.precomputed_table.is_empty(); + let allow_ephemeral_precomputed = reader.pq.nbits == 8 + && metric == MetricType::L2 + && by_residual + && !use_precomputed + && match reuse_mode { + IvfPqBatchTableReuseMode::Off => false, + IvfPqBatchTableReuseMode::On => true, + IvfPqBatchTableReuseMode::Auto => nq >= MIN_EPHEMERAL_PRECOMPUTE_QUERIES, + }; let all_ip_tables: Vec> = if use_precomputed { (0..nq) .into_par_iter() @@ -1630,6 +1883,7 @@ pub fn search_batch_reader_filter( } else { Vec::new() }; + let mut stable_pq_norms = None; let mut heaps = (0..nq).map(|_| TopKHeap::new(k)).collect::>(); let mut batch_start = 0usize; @@ -1701,6 +1955,71 @@ pub fn search_batch_reader_filter( .iter() .map(|list| matching_rows(&list.ids, filter)) .collect::>(); + let matching_list_count = matching_rows_by_list + .iter() + .filter(|rows| has_matching_rows(rows.as_ref())) + .count(); + let (active_query_count, probe_count) = if allow_ephemeral_precomputed { + let mut probe_count = 0usize; + let mut active_query_count = 0usize; + for probe_indices in &all_probe_indices { + let matching_probe_count = probe_indices + .iter() + .filter(|&&list_id| { + let position = list_positions[list_id]; + position != usize::MAX + && has_matching_rows(matching_rows_by_list[position].as_ref()) + }) + .count(); + probe_count += matching_probe_count; + active_query_count += usize::from(matching_probe_count > 0); + } + (active_query_count, probe_count) + } else { + (0, 0) + }; + let query_scratch_count = active_query_count.min(rayon::current_num_threads()); + let use_ephemeral_precomputed = allow_ephemeral_precomputed + && ephemeral_precomputed_table_fits_budget( + matching_list_count, + query_scratch_count, + m, + ksub, + ) + && (reuse_mode == IvfPqBatchTableReuseMode::On + || should_use_ephemeral_precomputation( + matching_list_count, + active_query_count, + probe_count, + )); + let ephemeral_precomputed_tables = if use_ephemeral_precomputed { + let pq_norms = stable_pq_norms + .get_or_insert_with(|| compute_stable_ephemeral_pq_norms(&reader.pq)); + loaded_lists + .par_iter() + .zip(&matching_rows_by_list) + .map(|(list, rows)| { + let mut table = Vec::new(); + if has_matching_rows(rows.as_ref()) { + fill_stable_ephemeral_list_table( + &reader.quantizer_centroids[list.list_id * d..(list.list_id + 1) * d], + &reader.pq, + pq_norms, + &mut table, + ); + } + table + }) + .collect::>() + } else { + Vec::new() + }; + observe_ephemeral_precomputed_lists( + ephemeral_precomputed_tables + .iter() + .filter(|table| !table.is_empty()) + .count(), + ); let rows = (0..nq) .into_par_iter() @@ -1726,24 +2045,60 @@ pub fn search_batch_reader_filter( }; let mut heap = TopKHeap::new(k); let mut scratch = ReaderScanScratch::default(); + let query_uses_ephemeral_precomputed = use_ephemeral_precomputed + && all_probe_indices[qi].iter().any(|&list_id| { + let position = list_positions[list_id]; + position != usize::MAX && !ephemeral_precomputed_tables[position].is_empty() + }); + if query_uses_ephemeral_precomputed { + fill_stable_ephemeral_query_table(query, &reader.pq, &mut scratch.ip_table); + } for (probe_rank, &list_id) in all_probe_indices[qi].iter().enumerate() { let position = list_positions[list_id]; if position == usize::MAX { continue; } - let dis0 = if use_precomputed { + let use_ephemeral_list = query_uses_ephemeral_precomputed + && !ephemeral_precomputed_tables[position].is_empty(); + let dis0 = if use_ephemeral_list { + 0.0 + } else if use_precomputed { all_coarse_dists[qi][probe_rank] } else { 0.0 }; - scan_reader_list( - &loaded_lists[position], - dis0, - &ctx, - matching_rows_by_list[position].as_ref(), - &mut scratch, - &mut heap, - ); + if use_ephemeral_list { + combine_stable_ephemeral_tables( + &ephemeral_precomputed_tables[position], + &scratch.ip_table, + query, + &reader.quantizer_centroids[list_id * d..(list_id + 1) * d], + &reader.pq, + &mut scratch.sim_table, + ); + scan_reader_codes( + &scratch.sim_table, + loaded_lists[position].codes(), + &loaded_lists[position].ids, + m, + ksub, + reader.pq.nbits, + reader.transposed_codes, + dis0, + matching_rows_by_list[position].as_ref(), + &mut scratch.distances, + &mut heap, + ); + } else { + scan_reader_list( + &loaded_lists[position], + dis0, + &ctx, + matching_rows_by_list[position].as_ref(), + &mut scratch, + &mut heap, + ); + } } heap.into_sorted() }) @@ -1778,9 +2133,37 @@ pub fn search_batch_reader_roaring_filter( k: usize, nprobe: usize, roaring_filter_bytes: &[u8], +) -> io::Result<(Vec, Vec)> { + search_batch_reader_roaring_filter_with_reuse_mode( + reader, + queries, + nq, + k, + nprobe, + roaring_filter_bytes, + IvfPqBatchTableReuseMode::Auto, + ) +} + +pub fn search_batch_reader_roaring_filter_with_reuse_mode( + reader: &mut IVFPQIndexReader, + queries: &[f32], + nq: usize, + k: usize, + nprobe: usize, + roaring_filter_bytes: &[u8], + reuse_mode: IvfPqBatchTableReuseMode, ) -> io::Result<(Vec, Vec)> { let filter = decode_roaring_filter(roaring_filter_bytes)?; - search_batch_reader_filter(reader, queries, nq, k, nprobe, Some(&filter)) + search_batch_reader_filter_with_reuse_mode( + reader, + queries, + nq, + k, + nprobe, + Some(&filter), + reuse_mode, + ) } // --- Top-K Heap --- @@ -1967,6 +2350,78 @@ mod tests { data } + fn observed_ephemeral_precomputed_lists( + nq: usize, + nprobe: usize, + filter_step: Option, + apply_filter: bool, + seed: u64, + reuse_mode: IvfPqBatchTableReuseMode, + ) -> usize { + observed_ephemeral_precomputed_lists_with_nbits( + 8, + nq, + nprobe, + filter_step, + apply_filter, + seed, + reuse_mode, + ) + } + + fn observed_ephemeral_precomputed_lists_with_nbits( + nbits: usize, + nq: usize, + nprobe: usize, + filter_step: Option, + apply_filter: bool, + seed: u64, + reuse_mode: IvfPqBatchTableReuseMode, + ) -> usize { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let k = 5; + let data = generate_clustered_data(n, d, nlist, seed); + let ids = (0..n as i64).collect::>(); + let mut index = IVFPQIndex::with_nbits(d, nlist, m, nbits, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let filter = match filter_step { + Some(step) => ids.iter().copied().step_by(step).collect::>(), + None => HashSet::new(), + }; + let filter = if apply_filter { + Some(&filter as &dyn RowIdFilter) + } else { + None + }; + let precomputed_lists = AtomicUsize::new(0); + let mut reader = IVFPQIndexReader::open(Cursor::new(bytes)).unwrap(); + + search_batch_reader_filter_with_reuse_mode_and_observer( + &mut reader, + &data[..nq * d], + nq, + k, + nprobe, + filter, + reuse_mode, + |count| { + precomputed_lists.fetch_add(count, Ordering::Relaxed); + }, + ) + .unwrap(); + + precomputed_lists.load(Ordering::Relaxed) + } + fn assert_invalid_merge(base: &IVFPQIndex, other: &IVFPQIndex, expected_message: &str) { let mut target = IVFPQIndex::from_trained(base); let before_ids = target.ids.clone(); @@ -2591,6 +3046,65 @@ mod tests { assert_eq!(actual, expected); } + #[test] + fn reader_list_precomputed_table_matches_index_table() { + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let data = generate_clustered_data(n, d, nlist, 44); + let mut index = IVFPQIndex::new(d, nlist, m, MetricType::L2, false); + index.train(&data, n); + index.build_precomputed_table(); + let pq_norms = index.pq.compute_centroid_norms(); + + for list_id in 0..nlist { + let mut actual = Vec::new(); + fill_list_precomputed_table( + &index.quantizer_centroids[list_id * d..(list_id + 1) * d], + &index.pq, + &pq_norms, + &mut actual, + ); + let table_size = m * index.pq.ksub; + assert_eq!( + actual, + index.precomputed_table[list_id * table_size..(list_id + 1) * table_size] + ); + } + } + + #[test] + fn ephemeral_precomputation_requires_matching_probe_work() { + assert!(!should_use_ephemeral_precomputation(0, 0, 0)); + } + + #[test] + fn ephemeral_precomputation_respects_batch_memory_budget() { + let max_values = + crate::index_io_util::MAX_IVF_BATCH_READ_BYTES / std::mem::size_of::(); + let max_list_values = max_values / 3; + assert!(ephemeral_precomputed_table_fits_budget( + 1, + 1, + 1, + max_list_values + )); + assert!(!ephemeral_precomputed_table_fits_budget( + 1, + 1, + 1, + max_list_values + 1 + )); + assert!(!ephemeral_precomputed_table_fits_budget(0, 1, 1, 1)); + assert!(!ephemeral_precomputed_table_fits_budget( + usize::MAX, + usize::MAX, + usize::MAX, + usize::MAX + )); + } + #[test] fn test_precomputed_table_matches_normal_search() { let d = 16; @@ -3168,6 +3682,438 @@ mod tests { ); } + #[test] + fn filtered_batch_reader_uses_ephemeral_list_precomputation() { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + use std::sync::atomic::AtomicUsize; + + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let nq = 64; + let k = 5; + let nprobe = nlist; + let data = generate_clustered_data(n, d, nlist, 45); + let ids = (0..n as i64).map(|id| 50_000 + id * 3).collect::>(); + let mut index = IVFPQIndex::new(d, nlist, m, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let filter = ids.iter().copied().step_by(5).collect::>(); + let queries = &data[..nq * d]; + let precomputed_lists = AtomicUsize::new(0); + let mut reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + + let (batch_ids, batch_dists) = search_batch_reader_filter_with_observer( + &mut reader, + queries, + nq, + k, + nprobe, + Some(&filter), + |count| { + precomputed_lists.fetch_add(count, Ordering::Relaxed); + }, + ) + .unwrap(); + + assert_eq!(precomputed_lists.load(Ordering::Relaxed), nlist); + assert!( + reader.precomputed_table.is_empty(), + "batch-local precomputation must not remain resident on the reader" + ); + for query_index in 0..nq { + let mut single_reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + let query = &queries[query_index * d..(query_index + 1) * d]; + let (single_ids, single_dists) = + search_with_reader_filter(&mut single_reader, query, k, nprobe, Some(&filter)) + .unwrap(); + assert_eq!( + &batch_ids[query_index * k..(query_index + 1) * k], + single_ids.as_slice() + ); + for (batch, single) in batch_dists[query_index * k..(query_index + 1) * k] + .iter() + .zip(&single_dists) + { + // The algebraically equivalent precomputed formula changes + // floating-point accumulation order. Allow a small absolute + // floor near zero plus a few ULPs for large distances. + let tolerance = 1e-4 + 4.0 * f32::EPSILON * single.abs(); + assert!( + (batch - single).abs() <= tolerance, + "ephemeral precomputation distance {batch} should match direct residual distance {single} within {tolerance}" + ); + } + } + } + + #[test] + fn unfiltered_batch_reader_uses_ephemeral_list_precomputation() { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + use std::sync::atomic::AtomicUsize; + + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let nq = 64; + let k = 5; + let nprobe = nlist; + let data = generate_clustered_data(n, d, nlist, 49); + let ids = (0..n as i64).map(|id| 70_000 + id * 3).collect::>(); + let mut index = IVFPQIndex::new(d, nlist, m, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let queries = &data[..nq * d]; + let precomputed_lists = AtomicUsize::new(0); + let mut reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + + let (batch_ids, batch_dists) = search_batch_reader_filter_with_observer( + &mut reader, + queries, + nq, + k, + nprobe, + None, + |count| { + precomputed_lists.fetch_add(count, Ordering::Relaxed); + }, + ) + .unwrap(); + + assert_eq!(precomputed_lists.load(Ordering::Relaxed), nlist); + assert!( + reader.precomputed_table.is_empty(), + "batch-local precomputation must not remain resident on the reader" + ); + for query_index in 0..nq { + let mut single_reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + let query = &queries[query_index * d..(query_index + 1) * d]; + let (single_ids, single_dists) = + search_with_reader_filter(&mut single_reader, query, k, nprobe, None).unwrap(); + assert_eq!( + &batch_ids[query_index * k..(query_index + 1) * k], + single_ids.as_slice() + ); + for (batch, single) in batch_dists[query_index * k..(query_index + 1) * k] + .iter() + .zip(&single_dists) + { + let tolerance = 1e-4 + 4.0 * f32::EPSILON * single.abs(); + assert!( + (batch - single).abs() <= tolerance, + "ephemeral precomputation distance {batch} should match direct residual distance {single} within {tolerance}" + ); + } + } + } + + #[test] + fn forced_ephemeral_reuse_is_stable_for_8bit_large_offsets() { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + + let d = 4; + let nlist = 4; + let m = 1; + let n = 512; + let nq = MIN_EPHEMERAL_PRECOMPUTE_QUERIES; + let k = 8; + let common = 1_000_000.0f32; + let spread = 500_000.0f32; + let data = (0..n) + .flat_map(|row| { + let cluster = row % nlist; + let point = row / nlist; + (0..d).map(move |dimension| { + common + + cluster as f32 * spread + + (((point * 17 + dimension * 13) % 31) as f32 - 15.0) * spread / 16.0 + }) + }) + .collect::>(); + let ids = (0..n as i64).collect::>(); + let mut index = IVFPQIndex::new(d, nlist, m, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let queries = &data[..nq * d]; + let mut reader = IVFPQIndexReader::open(Cursor::new(bytes)).unwrap(); + let (result_ids, result_distances) = search_batch_reader_with_reuse_mode( + &mut reader, + queries, + nq, + k, + nlist, + IvfPqBatchTableReuseMode::On, + ) + .unwrap(); + + let code_size = index.pq.code_size(); + let mut decoded = vec![0.0f32; d]; + for query_index in 0..nq { + let query = &queries[query_index * d..(query_index + 1) * d]; + for rank in 0..k { + let offset = query_index * k + rank; + let id = result_ids[offset]; + let reported = result_distances[offset]; + assert!( + reported >= 0.0, + "query {query_index} rank {rank} produced negative squared L2 distance {reported}" + ); + + let (list_id, position) = index + .ids + .iter() + .enumerate() + .find_map(|(list_id, list_ids)| { + list_ids + .iter() + .position(|candidate| *candidate == id) + .map(|position| (list_id, position)) + }) + .unwrap(); + let code_offset = position * code_size; + index.pq.decode( + &index.codes[list_id][code_offset..code_offset + code_size], + &mut decoded, + ); + let centroid = &index.quantizer_centroids[list_id * d..(list_id + 1) * d]; + let exact = (0..d) + .map(|dimension| { + let delta = f64::from(query[dimension]) + - f64::from(centroid[dimension]) + - f64::from(decoded[dimension]); + delta * delta + }) + .sum::(); + let tolerance = 1e-3 + 8.0 * f64::from(f32::EPSILON) * exact.abs().max(1.0); + assert!( + (f64::from(reported) - exact).abs() <= tolerance, + "query {query_index} rank {rank} reported {reported}, decoded oracle {exact}, tolerance {tolerance}" + ); + } + } + } + + #[test] + fn small_filtered_batch_reader_skips_ephemeral_list_precomputation() { + assert_eq!( + observed_ephemeral_precomputed_lists( + 4, + 4, + Some(5), + true, + 46, + IvfPqBatchTableReuseMode::Auto, + ), + 0, + "small batches should keep the direct residual-table path" + ); + } + + #[test] + fn single_probe_filtered_batch_reader_skips_ephemeral_list_precomputation() { + assert_eq!( + observed_ephemeral_precomputed_lists( + MIN_EPHEMERAL_PRECOMPUTE_QUERIES, + 1, + Some(5), + true, + 47, + IvfPqBatchTableReuseMode::Auto, + ), + 0, + "single-probe batches cannot amortize list precomputation" + ); + } + + #[test] + fn empty_filtered_batch_reader_skips_ephemeral_list_precomputation() { + assert_eq!( + observed_ephemeral_precomputed_lists( + MIN_EPHEMERAL_PRECOMPUTE_QUERIES, + 4, + None, + true, + 48, + IvfPqBatchTableReuseMode::Auto, + ), + 0, + "lists without matching rows should not be precomputed" + ); + } + + #[test] + fn small_unfiltered_batch_reader_skips_ephemeral_list_precomputation() { + assert_eq!( + observed_ephemeral_precomputed_lists( + 4, + 4, + None, + false, + 50, + IvfPqBatchTableReuseMode::Auto, + ), + 0, + "small unfiltered batches should keep the direct residual-table path" + ); + } + + #[test] + fn batch_table_reuse_off_never_precomputes_list_tables() { + assert_eq!( + observed_ephemeral_precomputed_lists( + MIN_EPHEMERAL_PRECOMPUTE_QUERIES, + 4, + Some(5), + true, + 51, + IvfPqBatchTableReuseMode::Off, + ), + 0, + "off mode must keep the direct residual-table path" + ); + } + + #[test] + fn batch_table_reuse_off_ignores_resident_precomputed_tables() { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let nq = 4; + let k = 5; + let data = generate_clustered_data(n, d, nlist, 53); + let ids = (0..n as i64).collect::>(); + let mut index = IVFPQIndex::new(d, nlist, m, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let queries = &data[..nq * d]; + let mut direct_reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + let expected = search_batch_reader_with_reuse_mode( + &mut direct_reader, + queries, + nq, + k, + nlist, + IvfPqBatchTableReuseMode::Off, + ) + .unwrap(); + + let mut optimized_reader = IVFPQIndexReader::open(Cursor::new(bytes)).unwrap(); + optimized_reader.optimize_for_search().unwrap(); + assert!(!optimized_reader.precomputed_table.is_empty()); + optimized_reader.precomputed_table.fill(1_000_000_000.0); + let actual = search_batch_reader_with_reuse_mode( + &mut optimized_reader, + queries, + nq, + k, + nlist, + IvfPqBatchTableReuseMode::Off, + ) + .unwrap(); + + assert_eq!(actual, expected, "Off must ignore resident reuse tables"); + } + + #[test] + fn batch_table_reuse_on_precomputes_for_small_batches() { + assert!( + observed_ephemeral_precomputed_lists( + 4, + 4, + Some(5), + true, + 52, + IvfPqBatchTableReuseMode::On, + ) > 0, + "on mode must bypass the automatic batch-size heuristic" + ); + } + + #[test] + fn four_bit_batch_table_reuse_modes_skip_ephemeral_precomputation() { + for reuse_mode in [IvfPqBatchTableReuseMode::Auto, IvfPqBatchTableReuseMode::On] { + assert_eq!( + observed_ephemeral_precomputed_lists_with_nbits( + 4, + MIN_EPHEMERAL_PRECOMPUTE_QUERIES, + 4, + None, + false, + 54, + reuse_mode, + ), + 0, + "4-bit {reuse_mode:?} must keep the existing scan path" + ); + } + } + + #[test] + fn four_bit_auto_batch_table_reuse_matches_off() { + use crate::io::{write_index, IVFPQIndexReader, PosWriter}; + + let d = 16; + let nlist = 4; + let m = 4; + let n = 600; + let nq = MIN_EPHEMERAL_PRECOMPUTE_QUERIES; + let k = 10; + let nprobe = nlist; + let data = generate_clustered_data(n, d, nlist, 55); + let ids = (0..n as i64).collect::>(); + let mut index = IVFPQIndex::with_nbits(d, nlist, m, 4, MetricType::L2, false); + index.train(&data, n); + index.add(&data, &ids, n); + + let mut bytes = Vec::new(); + write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap(); + let queries = &data[..nq * d]; + + let mut off_reader = IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap(); + let expected = search_batch_reader_with_reuse_mode( + &mut off_reader, + queries, + nq, + k, + nprobe, + IvfPqBatchTableReuseMode::Off, + ) + .unwrap(); + + let mut auto_reader = IVFPQIndexReader::open(Cursor::new(bytes)).unwrap(); + let actual = search_batch_reader_with_reuse_mode( + &mut auto_reader, + queries, + nq, + k, + nprobe, + IvfPqBatchTableReuseMode::Auto, + ) + .unwrap(); + + assert_eq!( + actual, expected, + "4-bit Auto must preserve the Off path results" + ); + } + #[test] fn test_batch_reader_empty_roaring_filter_returns_empty_results() { use crate::io::{write_index, IVFPQIndexReader, PosWriter}; diff --git a/ffi/src/lib.rs b/ffi/src/lib.rs index 7fa2c2c..7a5368b 100644 --- a/ffi/src/lib.rs +++ b/ffi/src/lib.rs @@ -19,9 +19,9 @@ use paimon_vindex_core::distance::MetricType; use paimon_vindex_core::index::{ - SearchWidth, VectorIndexConfig, VectorIndexMetadata, VectorIndexReadPlan, VectorIndexReader, - VectorIndexReaderOptions, VectorIndexTrainer, VectorIndexTraining, VectorIndexWriter, - VectorSearchParams, + IvfPqBatchTableReuseMode, SearchWidth, VectorIndexConfig, VectorIndexMetadata, + VectorIndexReadPlan, VectorIndexReader, VectorIndexReaderOptions, VectorIndexTrainer, + VectorIndexTraining, VectorIndexWriter, VectorSearchParams, }; use paimon_vindex_core::io::{ReadRequest, SeekRead, SeekReadCapabilities, SeekWrite}; use std::cell::RefCell; @@ -525,6 +525,7 @@ fn search_params_from_ffi(params: PaimonVindexSearchParams) -> Result Result Result Result { + match code { + 0 => Ok(IvfPqBatchTableReuseMode::Off), + 1 => Ok(IvfPqBatchTableReuseMode::On), + 2 => Ok(IvfPqBatchTableReuseMode::Auto), + value => Err(format!("invalid IVF-PQ batch table reuse mode: {value}")), + } +} + fn call_int_method(env: &mut JNIEnv, object: &JObject, name: &str) -> Result { env.call_method(object, name, "()I", &[]) .and_then(|value| value.i()) @@ -1023,3 +1035,25 @@ pub extern "system" fn Java_org_apache_paimon_index_vector_VectorIndexNative_fre } }) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ivfpq_batch_table_reuse_codes_map_to_core_modes() { + assert_eq!( + ivfpq_batch_table_reuse_mode(0).unwrap(), + IvfPqBatchTableReuseMode::Off + ); + assert_eq!( + ivfpq_batch_table_reuse_mode(1).unwrap(), + IvfPqBatchTableReuseMode::On + ); + assert_eq!( + ivfpq_batch_table_reuse_mode(2).unwrap(), + IvfPqBatchTableReuseMode::Auto + ); + assert!(ivfpq_batch_table_reuse_mode(3).is_err()); + } +}