jerry-024 commented on code in PR #106:
URL:
https://github.com/apache/paimon-vector-index/pull/106#discussion_r4023897141
##########
core/tests/range_search.rs:
##########
@@ -947,6 +952,596 @@ fn serialize_roaring(allowed: &HashSet<i64>) -> Vec<u8> {
bytes
}
+fn build_sq_index(dimension: usize, rows_per_list: usize, nlist: usize) ->
IVFSQIndex {
+ let mut index = IVFSQIndex::new(dimension, nlist, MetricType::L2);
+ index.set_quantizer_centroids(
+ (0..nlist)
+ .flat_map(|list_id| {
+ (0..dimension).map(move |component| list_id as f32 + component
as f32 * 0.01)
+ })
+ .collect(),
+ );
+ index.sq = ScalarQuantizer::with_bounds(dimension, -10.0, 10.0);
+ for list_id in 0..nlist {
+ index.list_sqs[list_id] = ScalarQuantizer::with_dimension_bounds(
+ dimension,
+ (0..dimension)
+ .map(|component| -0.7 - component as f32 * 0.001)
+ .collect(),
+ vec![1.3 + list_id as f32 * 0.1; dimension],
+ );
+ for row in (0..rows_per_list).rev() {
+ index.ids[list_id].push((1i64 << 33) + (list_id * rows_per_list +
row) as i64);
+ index.codes[list_id]
+ .extend((0..dimension).map(|component| ((row * 17 + component
* 13) % 256) as u8));
+ }
+ }
+ index
+}
+
+fn serialize_sq(index: &IVFSQIndex) -> Vec<u8> {
+ let mut bytes = Vec::new();
+ write_ivfsq_index(index, &mut PosWriter::new(&mut bytes)).unwrap();
+ bytes
+}
+
+#[test]
+fn ivf_sq_range_matches_sq_estimates_for_all_entry_points() {
+ let dimension = 65;
+ let nlist = 4;
+ let rows_per_list = 67;
+ let index = build_sq_index(dimension, rows_per_list, nlist);
+ let queries: Vec<f32> = [0.0, 1.5, 3.0]
+ .into_iter()
+ .flat_map(|offset| (0..dimension).map(move |component| offset +
component as f32 * 0.01))
+ .collect();
+ let mut reader =
VectorIndexReader::open(Cursor::new(serialize_sq(&index))).unwrap();
+ let allowed: HashSet<i64> = index
+ .ids
+ .iter()
+ .flatten()
+ .copied()
+ .filter(|id| id % 3 == 0)
+ .collect();
+ let filter = serialize_roaring(&allowed);
+ for nprobe in [1, 2, nlist, nlist + 10] {
+ let band = l2(10.0, 180.0);
+ let params = VectorRangeSearchParams::new(band, nprobe);
+ let batch = reader.range_search_batch(&queries, 3, params).unwrap();
+ let filtered_batch = reader
+ .range_search_batch_with_roaring_filter(&queries, 3, params,
&filter)
+ .unwrap();
+ for (query_index, query) in
queries.chunks_exact(dimension).enumerate() {
+ let (ids, distances) = top_k(&mut reader, query, rows_per_list *
nlist, nprobe);
+ let expected = bits_of(
+ ids.into_iter()
+ .zip(distances)
+ .filter(|(id, distance)| *id != -1 &&
band.admit(*distance))
+ .collect(),
+ );
+ assert!(!expected.is_empty());
+ assert_eq!(pairs_of(batch.query(query_index)), expected);
+ assert_eq!(
+ pairs_of(reader.range_search(query, params).unwrap().query(0)),
+ expected
+ );
+ let filtered: Vec<_> = expected
+ .into_iter()
+ .filter(|(id, _)| allowed.contains(id))
+ .collect();
+ assert!(!filtered.is_empty());
+ assert_eq!(pairs_of(filtered_batch.query(query_index)), filtered);
+ assert_eq!(
+ pairs_of(
+ reader
+ .range_search_with_roaring_filter(query, params,
&filter)
+ .unwrap()
+ .query(0)
+ ),
+ filtered,
+ );
+ }
+ }
+}
+
+#[test]
+fn ivf_sq_range_uses_estimated_not_original_distance() {
+ let mut index = IVFSQIndex::new(1, 1, MetricType::L2);
+ index.set_quantizer_centroids(vec![0.0]);
+ index.sq = ScalarQuantizer::with_bounds(1, 0.0, 255.0);
+ index.list_sqs[0] = index.sq.clone();
+ let vectors = [0.49, 0.51];
+ let ids = [10, 20];
+ index.add(&vectors, &ids, 2);
+ let mut reader =
VectorIndexReader::open(Cursor::new(serialize_sq(&index))).unwrap();
+ for band in [l2(0.0, 0.1), l2(0.2, 0.3)] {
+ let result = reader
+ .range_search(&[0.0], VectorRangeSearchParams::new(band, 1))
+ .unwrap();
+ let (labels, distances) = top_k(&mut reader, &[0.0], 2, 1);
+ let expected = bits_of(
+ labels
+ .into_iter()
+ .zip(distances)
+ .filter(|(_, distance)| band.admit(*distance))
+ .collect(),
+ );
+ assert_eq!(pairs_of(result.query(0)), expected);
+ assert_ne!(
+ pairs_of(result.query(0)),
+ bits_of(brute_force_band(&[0.0], &vectors, &ids, 1, band))
+ );
+ }
+}
+
+#[test]
+fn ivf_sq_range_preserves_boundaries_and_has_no_top_k_cap() {
+ let index = build_sq_index(65, 67, 3);
+ let mut reader =
VectorIndexReader::open(Cursor::new(serialize_sq(&index))).unwrap();
+ let query = vec![0.3; 65];
+ let (ids, distances) = top_k(&mut reader, &query, 201, 3);
+ let lower = distances[20];
+ let upper = distances[80];
+ assert!(lower < upper);
+ for band in [
+ l2(lower, upper),
+ l2(lower, lower),
+ l2(0.0, 0.0),
+ DistanceBand::new(Bound::Unbounded, Bound::Finite(upper),
MetricType::L2).unwrap(),
+ DistanceBand::new(Bound::Finite(lower), Bound::Unbounded,
MetricType::L2).unwrap(),
+ DistanceBand::new(Bound::Unbounded, Bound::Unbounded,
MetricType::L2).unwrap(),
+ ] {
+ let result = reader
+ .range_search(&query, VectorRangeSearchParams::new(band, 10))
+ .unwrap();
+ let expected = bits_of(
+ ids.iter()
+ .copied()
+ .zip(distances.iter().copied())
+ .filter(|(_, distance)| band.admit(*distance))
+ .collect(),
+ );
+ assert_eq!(pairs_of(result.query(0)), expected);
+ assert_eq!(result.query(0).stats.rows_committed(), expected.len());
+ if band.is_empty() {
+ assert_eq!(result.query(0).stats.lists_probed(), 0);
+ assert_eq!(result.query(0).stats.rows_scanned(), 0);
+ assert_eq!(result.call_stats().list_reads(), 0);
+ } else {
+ assert_eq!(result.query(0).stats.lists_probed(), 3);
+ assert_eq!(result.query(0).stats.rows_scanned(), 201);
+ }
+ }
+}
+
+#[test]
+fn ivf_sq_range_parallel_scans_preserve_query_order_and_membership() {
+ let index = build_sq_index(65, 2_101, 4);
+ let bytes = serialize_sq(&index);
+ let queries: Vec<f32> = [0.0, 0.5, 2.0]
+ .into_iter()
+ .flat_map(|value| vec![value; 65])
+ .collect();
+ let mut reader = VectorIndexReader::open(Cursor::new(bytes)).unwrap();
+ let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded,
MetricType::L2).unwrap();
+ let all = reader
+ .range_search_batch(&queries, 3, VectorRangeSearchParams::new(whole,
4))
+ .unwrap();
+ let band = l2(15.0, 180.0);
+ let params = VectorRangeSearchParams::new(band, 4);
+ let expected: Vec<_> = (0..3)
+ .map(|query_index| {
+ bits_of(
+ all.query(query_index)
+ .labels
+ .iter()
+ .copied()
+ .zip(all.query(query_index).distances.iter().copied())
+ .filter(|(_, distance)| band.admit(*distance))
+ .collect(),
+ )
+ })
+ .collect();
+ let mut permuted = queries[130..].to_vec();
+ permuted.extend_from_slice(&queries[..130]);
+ let allowed: HashSet<_> = index
+ .ids
+ .iter()
+ .flatten()
+ .copied()
+ .filter(|id| id % 5 == 0)
+ .collect();
+ for threads in [1, 4] {
+ rayon::ThreadPoolBuilder::new()
+ .num_threads(threads)
+ .build()
+ .unwrap()
+ .install(|| {
+ let result = reader.range_search_batch(&permuted, 3,
params).unwrap();
+ let filtered = reader
+ .range_search_batch_with_roaring_filter(
+ &queries,
+ 3,
+ params,
+ &serialize_roaring(&allowed),
+ )
+ .unwrap();
+ for (query_index, &original_index) in [2, 0,
1].iter().enumerate() {
+ assert!(!expected[original_index].is_empty());
+ assert_eq!(
+ pairs_of(result.query(query_index)),
+ expected[original_index]
+ );
+ let single = reader
+ .range_search(
+ &queries[original_index * 65..(original_index + 1)
* 65],
+ params,
+ )
+ .unwrap();
+ assert_eq!(pairs_of(single.query(0)),
expected[original_index]);
+ assert_eq!(
+ pairs_of(filtered.query(original_index)),
+ expected[original_index]
+ .iter()
+ .copied()
+ .filter(|(id, _)| allowed.contains(id))
+ .collect::<Vec<_>>()
+ );
+ }
+ });
+ }
+}
+
+#[derive(Default)]
+struct SqReadTrace {
+ calls: usize,
+ ranges: usize,
+ max_bytes: usize,
+}
+
+struct SqRecordingReader {
+ inner: Cursor<Vec<u8>>,
+ trace: Arc<Mutex<SqReadTrace>>,
+}
+
+impl SeekRead for SqRecordingReader {
+ fn pread(&mut self, ranges: &mut [ReadRequest<'_>]) -> std::io::Result<()>
{
+ let mut trace = self.trace.lock().unwrap();
+ trace.calls += 1;
+ trace.ranges += ranges.len();
+ for request in ranges.iter() {
+ trace.max_bytes = trace.max_bytes.max(request.buf.len());
+ }
+ self.inner.pread(ranges)
+ }
+}
+
+#[test]
+fn ivf_sq_range_reuses_shared_lists_and_cache_with_query_local_filters() {
+ let mut index = build_sq_index(33, 67, 4);
+ index.ids[3].clear();
+ index.codes[3].clear();
+ let bytes = serialize_sq(&index);
+ let whole = DistanceBand::new(Bound::Unbounded, Bound::Unbounded,
MetricType::L2).unwrap();
+ let params = VectorRangeSearchParams::new(whole, 4);
+ for budget in [0, 4 * 1024 * 1024] {
+ let trace = Arc::new(Mutex::new(SqReadTrace::default()));
+ let source = SqRecordingReader {
+ inner: Cursor::new(bytes.clone()),
+ trace: Arc::clone(&trace),
+ };
+ let mut reader =
+ IVFSQIndexReader::open_with_options(source,
VectorIndexReaderOptions::new(budget))
+ .unwrap();
+ *trace.lock().unwrap() = SqReadTrace::default();
+ let queries = vec![0.0; 3 * 33];
+ let batch = reader.range_search_batch(&queries, 3, params).unwrap();
+ assert_eq!(batch.call_stats().list_reads(), 3);
+ assert_eq!(trace.lock().unwrap().calls, 1);
+ assert_eq!(trace.lock().unwrap().ranges, 3);
+ assert_eq!(batch.lims(), &[0, 201, 402, 603]);
+ let empty_filter = serialize_roaring(&HashSet::new());
+ let filtered = reader
+ .range_search_batch_with_roaring_filter(&queries, 3, params,
&empty_filter)
+ .unwrap();
+ assert!(filtered.labels().is_empty());
+ let repeated = reader.range_search(&queries[..33], params).unwrap();
+ assert_eq!(pairs_of(repeated.query(0)), pairs_of(batch.query(0)));
+ if budget > 0 {
+ assert_eq!(trace.lock().unwrap().calls, 1);
+ assert_eq!(filtered.call_stats().list_reads(), 0);
+ assert_eq!(repeated.call_stats().list_reads(), 0);
+ } else {
+ assert_eq!(trace.lock().unwrap().calls, 3);
+ assert_eq!(repeated.call_stats().list_reads(), 3);
+ }
+ }
+}
+
+#[test]
+fn ivf_sq_range_respects_bounded_reads_and_deduplicates_partial_probes() {
+ struct SingleRangeReader(SqRecordingReader);
+
+ impl SeekRead for SingleRangeReader {
+ fn pread(&mut self, ranges: &mut [ReadRequest<'_>]) ->
std::io::Result<()> {
+ assert!(ranges.len() <= 1);
+ self.0.pread(ranges)
+ }
+
+ fn read_capabilities(&self) ->
paimon_vindex_core::io::SeekReadCapabilities {
+ paimon_vindex_core::io::SeekReadCapabilities {
+ max_ranges_per_pread: 1,
+ ..Default::default()
+ }
+ }
+ }
+
+ let index = build_sq_index(33, 67, 4);
+ let mut queries = index.quantizer_centroids()[..66].to_vec();
+ queries.extend_from_slice(&index.quantizer_centroids()[..33]);
+ let trace = Arc::new(Mutex::new(SqReadTrace::default()));
+ let source = SingleRangeReader(SqRecordingReader {
+ inner: Cursor::new(serialize_sq(&index)),
+ trace: Arc::clone(&trace),
+ });
+ let mut reader = IVFSQIndexReader::open(source).unwrap();
+ let band = DistanceBand::new(Bound::Unbounded, Bound::Unbounded,
MetricType::L2).unwrap();
+ for (nprobe, reads, rows) in [(1, 2, 67), (4, 4, 268)] {
+ *trace.lock().unwrap() = SqReadTrace::default();
+ let result = reader
+ .range_search_batch(&queries, 3,
VectorRangeSearchParams::new(band, nprobe))
+ .unwrap();
+ assert_eq!(result.call_stats().list_reads(), reads);
+ assert_eq!(trace.lock().unwrap().calls, reads);
+ for query_index in 0..3 {
+ assert_eq!(result.query(query_index).stats.lists_probed(), nprobe);
+ assert_eq!(result.query(query_index).stats.rows_scanned(), rows);
+ assert_eq!(result.query(query_index).labels.len(), rows);
+ }
+ assert_eq!(pairs_of(result.query(0)), pairs_of(result.query(2)));
+ }
+}
+
+#[test]
+fn ivf_sq_range_streams_oversized_lists_for_single_batch_and_filter() {
+ let dimension = 512;
+ let count = 131_073;
+ let mut index = IVFSQIndex::new(dimension, 1, MetricType::L2);
+ index.set_quantizer_centroids(vec![0.0; dimension]);
+ index.sq = ScalarQuantizer::with_bounds(dimension, 0.0, 1.0);
+ index.list_sqs[0] = index.sq.clone();
+ index.ids[0] = (0..count as i64).collect();
+ index.codes[0] = vec![255; count * dimension];
+ index.codes[0][..32 * dimension].fill(0);
+ index.codes[0][32 * dimension..64 * dimension].fill(64);
+ index.codes[0][(count - 1) * dimension..].fill(0);
+ let trace = Arc::new(Mutex::new(SqReadTrace::default()));
+ let source = SqRecordingReader {
+ inner: Cursor::new(serialize_sq(&index)),
+ trace: Arc::clone(&trace),
+ };
+ drop(index);
+ let mut reader = IVFSQIndexReader::open(source).unwrap();
+ let mut queries = vec![0.0; dimension];
+ queries.extend(vec![0.25; dimension]);
+ let params = VectorRangeSearchParams::new(l2(0.0, 1.0), 1);
+ *trace.lock().unwrap() = SqReadTrace::default();
+ let result = reader.range_search_batch(&queries, 2, params).unwrap();
+ assert_eq!(result.call_stats().list_reads(), 1);
+ assert!(trace.lock().unwrap().calls > 2);
+ assert!(trace.lock().unwrap().max_bytes < count * dimension);
+ assert!(trace.lock().unwrap().max_bytes <= 64 * 1024 * 1024);
+ let mut first_ids: Vec<_> = (0..32).collect();
+ first_ids.push((count - 1) as i64);
+ assert_eq!(result.query(0).labels, first_ids);
+ assert_eq!(result.query(1).labels, (32..64).collect::<Vec<i64>>());
+ let allowed: HashSet<i64> = (0..count as i64).filter(|id| id % 2 ==
0).collect();
+ let filter = serialize_roaring(&allowed);
+ let filtered = reader
+ .range_search_batch_with_roaring_filter(&queries, 2, params, &filter)
+ .unwrap();
+ for (query_index, query) in queries.chunks_exact(dimension).enumerate() {
+ assert_eq!(result.query(query_index).stats.rows_scanned(), count);
+ assert!(result.query(query_index).stats.early_abandoned() > count -
100);
+ let single = reader.range_search(query, params).unwrap();
+ assert_eq!(
+ pairs_of(single.query(0)),
+ pairs_of(result.query(query_index))
+ );
+ let expected: Vec<_> = pairs_of(result.query(query_index))
+ .into_iter()
+ .filter(|(id, _)| allowed.contains(id))
+ .collect();
+ assert_eq!(pairs_of(filtered.query(query_index)), expected);
+ assert_eq!(
+ pairs_of(
+ reader
+ .range_search_with_roaring_filter(query, params, &filter)
+ .unwrap()
+ .query(0)
+ ),
+ expected
+ );
+ }
+}
+
+#[test]
+fn ivf_sq_range_validates_all_entry_points_before_empty_band_shortcuts() {
+ use std::io::ErrorKind::{InvalidInput, Unsupported};
+
+ for metric in [MetricType::L2, MetricType::Cosine,
MetricType::InnerProduct] {
+ let mut index = build_sq_index(33, 35, 2);
+ index.metric = metric;
+ let bytes = serialize_sq(&index);
+ let mut unified =
VectorIndexReader::open(Cursor::new(bytes.clone())).unwrap();
+ let mut direct = IVFSQIndexReader::open(Cursor::new(bytes)).unwrap();
+ let empty = DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0),
metric).unwrap();
+ let valid_filter = serialize_roaring(&HashSet::new());
+ let mismatched_metric = if metric == MetricType::L2 {
+ MetricType::Cosine
+ } else {
+ MetricType::L2
+ };
+ let mismatched =
+ DistanceBand::new(Bound::Finite(1.0), Bound::Finite(1.0),
mismatched_metric).unwrap();
+ for (queries, query_count, band, nprobe, filter, error) in [
+ (
+ vec![0.0; 32],
+ 1,
+ empty,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![f32::NAN; 33],
+ 1,
+ empty,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![f32::INFINITY; 33],
+ 1,
+ empty,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![0.0; 33],
+ 1,
+ empty,
+ 0,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![0.0; 33],
+ 1,
+ mismatched,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![0.0; 33],
+ 1,
+ empty,
+ 2,
+ vec![0xde, 0xad],
+ Some(InvalidInput),
+ ),
+ (
+ vec![],
+ 0,
+ empty,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![0.0; 33],
+ usize::MAX,
+ empty,
+ 2,
+ valid_filter.clone(),
+ Some(InvalidInput),
+ ),
+ (
+ vec![0.0; 33],
+ 1,
+ empty,
+ 2,
+ valid_filter.clone(),
+ if metric == MetricType::L2 {
+ None
+ } else {
+ Some(Unsupported)
+ },
+ ),
+ ] {
+ let params = VectorRangeSearchParams::new(band, nprobe);
+ for low_level in [false, true] {
Review Comment:
<!-- dlf-review -->
**[Minor]** This test repeats the same validation rules across every wrapper.
The unified and direct entry points ultimately exercise the same batch
range-search validation, so the metric × reader × entry-point matrix rechecks
the dimension, non-finite, `nprobe`, and metric errors many times. Could we
keep one invalid-input matrix against the direct batch path, cover malformed
filters through the filtered path, and retain one delegation smoke case for the
remaining wrappers? That would preserve the boundary coverage while removing
most of this ~140-line test.
--
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]