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 f41224a  perf(ivfpq): speed up transposed PQ8 scans (#76)
f41224a is described below

commit f41224ac56151521a689c74dcd0f565aec59e815
Author: jerry <[email protected]>
AuthorDate: Thu Aug 13 15:45:19 2026 +0800

    perf(ivfpq): speed up transposed PQ8 scans (#76)
---
 core/Cargo.toml                         |   4 +
 core/benches/ivfpq_filter_scan_bench.rs | 157 +++++++++++++++++++++++++++++
 core/src/ivfpq.rs                       | 172 +++++++++++++++++++++-----------
 3 files changed, 273 insertions(+), 60 deletions(-)

diff --git a/core/Cargo.toml b/core/Cargo.toml
index 37b7fc3..9a08cac 100644
--- a/core/Cargo.toml
+++ b/core/Cargo.toml
@@ -58,6 +58,10 @@ harness = false
 name = "ivfpq_batch_reuse_bench"
 harness = false
 
+[[bench]]
+name = "ivfpq_filter_scan_bench"
+harness = false
+
 [[bench]]
 name = "ivfpq_train_bench"
 harness = false
diff --git a/core/benches/ivfpq_filter_scan_bench.rs 
b/core/benches/ivfpq_filter_scan_bench.rs
new file mode 100644
index 0000000..3218e44
--- /dev/null
+++ b/core/benches/ivfpq_filter_scan_bench.rs
@@ -0,0 +1,157 @@
+// 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 paimon_vindex_core::distance::MetricType;
+use paimon_vindex_core::io::{write_index, IVFPQIndexReader, PosWriter};
+use paimon_vindex_core::ivfpq::{search_batch_reader_filter, IVFPQIndex, 
RowIdFilter};
+use rand::rngs::StdRng;
+use rand::seq::SliceRandom;
+use rand::{Rng, SeedableRng};
+use std::hint::black_box;
+use std::io::Cursor;
+use std::time::{Duration, Instant};
+
+const D: usize = 768;
+const M: usize = 96;
+const NLIST: usize = 128;
+const NPROBE: usize = 64;
+const NQ: usize = 64;
+const K: usize = 3;
+const ROWS_PER_LIST: usize = 6_800;
+const WARMUPS: usize = 1;
+const ROUNDS: usize = 3;
+
+struct DensityFilter<'a> {
+    row_ranks: &'a [usize],
+    max_rank: usize,
+}
+
+impl RowIdFilter for DensityFilter<'_> {
+    fn contains(&self, id: i64) -> bool {
+        self.row_ranks
+            .get(id as usize)
+            .is_some_and(|&rank| rank < self.max_rank)
+    }
+}
+
+fn randomized_row_ranks() -> Vec<usize> {
+    let mut rng = StdRng::seed_from_u64(43);
+    let mut row_ranks = vec![0; NLIST * ROWS_PER_LIST];
+    let mut row_order = (0..ROWS_PER_LIST).collect::<Vec<_>>();
+    for list_id in 0..NLIST {
+        row_order.shuffle(&mut rng);
+        let base = list_id * ROWS_PER_LIST;
+        for (rank, &row) in row_order.iter().enumerate() {
+            row_ranks[base + row] = rank;
+        }
+    }
+    row_ranks
+}
+
+fn search(
+    reader: &mut IVFPQIndexReader<Cursor<Vec<u8>>>,
+    queries: &[f32],
+    filter: &DensityFilter,
+) -> Duration {
+    let started = Instant::now();
+    let result = search_batch_reader_filter(reader, queries, NQ, K, NPROBE, 
Some(filter)).unwrap();
+    let elapsed = started.elapsed();
+    black_box(result);
+    elapsed
+}
+
+fn median(samples: &mut [Duration]) -> Duration {
+    samples.sort_unstable();
+    samples[samples.len() / 2]
+}
+
+fn main() {
+    assert_eq!((NQ, NPROBE), (64, 64), "benchmark production query shape");
+    for name in [
+        "PAIMON_VINDEX_LOG_IVFPQ_BATCH_TIMING",
+        "PAIMON_VINDEX_LOG_IVFPQ_BATCH_REUSE",
+    ] {
+        assert!(
+            std::env::var_os(name).is_none(),
+            "unset {name} for this benchmark"
+        );
+    }
+    assert_eq!(
+        rayon::current_num_threads(),
+        1,
+        "run with RAYON_NUM_THREADS=1"
+    );
+    const { assert!(NQ >= 64, "keep production query-table reuse enabled") };
+
+    let mut rng = StdRng::seed_from_u64(42);
+    let mut index = IVFPQIndex::new(D, NLIST, M, MetricType::InnerProduct, 
false);
+    index.quantizer_centroids = (0..NLIST * D)
+        .map(|_| rng.gen_range(-1.0f32..1.0))
+        .collect();
+    index.pq.centroids = (0..M * index.pq.ksub * index.pq.dsub)
+        .map(|_| rng.gen_range(-1.0f32..1.0))
+        .collect();
+    for list_id in 0..NLIST {
+        let first_id = list_id * ROWS_PER_LIST;
+        index.ids[list_id] = (first_id..first_id + ROWS_PER_LIST)
+            .map(|id| id as i64)
+            .collect();
+        index.codes[list_id] = (0..ROWS_PER_LIST * M).map(|_| 
rng.gen()).collect();
+    }
+    let queries = (0..NQ * D)
+        .map(|_| rng.gen_range(-1.0f32..1.0))
+        .collect::<Vec<_>>();
+    let row_ranks = randomized_row_ranks();
+    let mut bytes = Vec::new();
+    write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap();
+    let densities = [
+        (1, "6.25"),
+        (2, "12.50"),
+        (3, "18.75"),
+        (4, "25.00"),
+        (8, "50.00"),
+        (16, "100.00"),
+    ];
+
+    println!(
+        "shape: d={D} m={M} nlist={NLIST} nprobe={NPROBE} nq={NQ} k={K} 
rows_per_list={ROWS_PER_LIST} threads=1 warmups={WARMUPS} rounds={ROUNDS}"
+    );
+    println!("density_percent,p50_ms");
+    for (matching_sixteenths, density) in densities {
+        let filter = DensityFilter {
+            row_ranks: &row_ranks,
+            max_rank: ROWS_PER_LIST * matching_sixteenths / 16,
+        };
+        assert_eq!(
+            (0..ROWS_PER_LIST)
+                .filter(|&id| filter.contains(id as i64))
+                .count(),
+            filter.max_rank
+        );
+        let mut reader = 
IVFPQIndexReader::open(Cursor::new(bytes.clone())).unwrap();
+        for _ in 0..WARMUPS {
+            search(&mut reader, &queries, &filter);
+        }
+        let mut samples = (0..ROUNDS)
+            .map(|_| search(&mut reader, &queries, &filter))
+            .collect::<Vec<_>>();
+        println!(
+            "{density},{:.3}",
+            median(&mut samples).as_secs_f64() * 1_000.0
+        );
+    }
+}
diff --git a/core/src/ivfpq.rs b/core/src/ivfpq.rs
index e45647f..5cf12fc 100644
--- a/core/src/ivfpq.rs
+++ b/core/src/ivfpq.rs
@@ -1219,44 +1219,18 @@ fn scan_codes_transposed_with_scratch(
     }
 
     dists.resize(count, 0.0);
-    let table = &sim_table[..ksub];
-    let column = &codes[..count];
-    let mut row = 0usize;
-    while row + 8 <= count {
-        dists[row] = dis0 + table[column[row] as usize];
-        dists[row + 1] = dis0 + table[column[row + 1] as usize];
-        dists[row + 2] = dis0 + table[column[row + 2] as usize];
-        dists[row + 3] = dis0 + table[column[row + 3] as usize];
-        dists[row + 4] = dis0 + table[column[row + 4] as usize];
-        dists[row + 5] = dis0 + table[column[row + 5] as usize];
-        dists[row + 6] = dis0 + table[column[row + 6] as usize];
-        dists[row + 7] = dis0 + table[column[row + 7] as usize];
-        row += 8;
-    }
-    while row < count {
-        dists[row] = dis0 + table[column[row] as usize];
-        row += 1;
-    }
+    transposed_column_init(
+        &mut dists[..count],
+        &codes[..count],
+        &sim_table[..ksub],
+        dis0,
+    );
     for sub in 1..m {
-        let table = &sim_table[sub * ksub..(sub + 1) * ksub];
-        let col_base = sub * count;
-        let column = &codes[col_base..col_base + count];
-        let mut row = 0usize;
-        while row + 8 <= count {
-            dists[row] += table[column[row] as usize];
-            dists[row + 1] += table[column[row + 1] as usize];
-            dists[row + 2] += table[column[row + 2] as usize];
-            dists[row + 3] += table[column[row + 3] as usize];
-            dists[row + 4] += table[column[row + 4] as usize];
-            dists[row + 5] += table[column[row + 5] as usize];
-            dists[row + 6] += table[column[row + 6] as usize];
-            dists[row + 7] += table[column[row + 7] as usize];
-            row += 8;
-        }
-        while row < count {
-            dists[row] += table[column[row] as usize];
-            row += 1;
-        }
+        transposed_column_add(
+            &mut dists[..count],
+            &codes[sub * count..(sub + 1) * count],
+            &sim_table[sub * ksub..(sub + 1) * ksub],
+        );
     }
 
     if let Some(rows) = matching_rows {
@@ -1270,6 +1244,63 @@ fn scan_codes_transposed_with_scratch(
     }
 }
 
+// A u8 code cannot index out of a 256-entry table, so converting the LUT to a
+// fixed-size array reference lets the compiler drop the per-lookup bounds
+// checks that otherwise dominate this hot loop for 8-bit scans.
+#[inline]
+fn transposed_column_init(dists: &mut [f32], column: &[u8], table: &[f32], 
dis0: f32) {
+    debug_assert_eq!(dists.len(), column.len());
+    if let Ok(table) = <&[f32; 256]>::try_from(table) {
+        let mut dist_chunks = dists.chunks_exact_mut(8);
+        let mut code_chunks = column.chunks_exact(8);
+        for (dist8, code8) in (&mut dist_chunks).zip(&mut code_chunks) {
+            let dist8: &mut [f32; 8] = dist8.try_into().unwrap();
+            let code8: &[u8; 8] = code8.try_into().unwrap();
+            for i in 0..8 {
+                dist8[i] = dis0 + table[code8[i] as usize];
+            }
+        }
+        for (dist, &code) in dist_chunks
+            .into_remainder()
+            .iter_mut()
+            .zip(code_chunks.remainder())
+        {
+            *dist = dis0 + table[code as usize];
+        }
+    } else {
+        for (dist, &code) in dists.iter_mut().zip(column) {
+            *dist = dis0 + table[code as usize];
+        }
+    }
+}
+
+#[inline]
+fn transposed_column_add(dists: &mut [f32], column: &[u8], table: &[f32]) {
+    debug_assert_eq!(dists.len(), column.len());
+    if let Ok(table) = <&[f32; 256]>::try_from(table) {
+        let mut dist_chunks = dists.chunks_exact_mut(8);
+        let mut code_chunks = column.chunks_exact(8);
+        for (dist8, code8) in (&mut dist_chunks).zip(&mut code_chunks) {
+            let dist8: &mut [f32; 8] = dist8.try_into().unwrap();
+            let code8: &[u8; 8] = code8.try_into().unwrap();
+            for i in 0..8 {
+                dist8[i] += table[code8[i] as usize];
+            }
+        }
+        for (dist, &code) in dist_chunks
+            .into_remainder()
+            .iter_mut()
+            .zip(code_chunks.remainder())
+        {
+            *dist += table[code as usize];
+        }
+    } else {
+        for (dist, &code) in dists.iter_mut().zip(column) {
+            *dist += table[code as usize];
+        }
+    }
+}
+
 /// Scan inverted list codes with 4-code batching for ILP (row-major layout).
 fn scan_codes_batched(
     sim_table: &[f32],
@@ -3469,7 +3500,7 @@ mod tests {
 
     #[test]
     fn transposed_scan_matches_scalar_distance_table() {
-        let count = 37;
+        let count = 40;
         let m = 7;
         let ksub = 256;
         let dis0 = 3.25;
@@ -3507,26 +3538,47 @@ mod tests {
         expected.sort_by(|left, right| left.0.total_cmp(&right.0));
         assert_eq!(heap.into_sorted(), expected);
 
-        let matching_positions = (0..count).step_by(5).collect::<Vec<_>>();
-        let matching_rows = MatchingRows::Sparse(matching_positions);
-        let mut filtered_heap = TopKHeap::new(matching_rows.len());
-        scan_codes_transposed_with_scratch(
-            &table,
-            &codes,
-            &ids,
+        for matching_count in [4, 5, 6] {
+            let matching_rows = 
MatchingRows::Sparse((0..matching_count).collect());
+            let mut filtered_heap = TopKHeap::new(matching_rows.len());
+            scan_codes_transposed_with_scratch(
+                &table,
+                &codes,
+                &ids,
+                count,
+                m,
+                ksub,
+                dis0,
+                Some(&matching_rows),
+                &mut filtered_heap,
+                &mut scratch,
+            );
+            let filtered_expected = expected
+                .iter()
+                .copied()
+                .filter(|(_, id)| *id < ids[0] + matching_count as i64)
+                .collect::<Vec<_>>();
+            assert_eq!(filtered_heap.into_sorted(), filtered_expected);
+        }
+    }
+
+    #[test]
+    fn transposed_sparse_scan_uses_configured_crossover() {
+        let count = 40;
+        let boundary = count / TRANSPOSED_SPARSE_SCAN_DIVISOR;
+        let at_boundary = MatchingRows::Sparse((0..boundary).collect());
+        let above_boundary = MatchingRows::Sparse((0..boundary + 1).collect());
+
+        assert!(should_scan_sparse(
             count,
-            m,
-            ksub,
-            dis0,
-            Some(&matching_rows),
-            &mut filtered_heap,
-            &mut scratch,
-        );
-        let filtered_expected = expected
-            .into_iter()
-            .filter(|(_, id)| (id - ids[0]) % 5 == 0)
-            .collect::<Vec<_>>();
-        assert_eq!(filtered_heap.into_sorted(), filtered_expected);
+            &at_boundary,
+            TRANSPOSED_SPARSE_SCAN_DIVISOR
+        ));
+        assert!(!should_scan_sparse(
+            count,
+            &above_boundary,
+            TRANSPOSED_SPARSE_SCAN_DIVISOR
+        ));
     }
 
     #[test]
@@ -3630,14 +3682,14 @@ mod tests {
     fn ivfpq_batch_timing_distinguishes_sparse_and_dense_scan_work() {
         let mut timing = IvfpqBatchTiming::default();
         let sparse_rows = MatchingRows::Sparse((0..10).collect());
-        let dense_rows = MatchingRows::Sparse((0..20).collect());
+        let dense_rows = MatchingRows::Sparse((0..26).collect());
 
         timing.record_scan_work(100, Some(&sparse_rows), 4, 8, true);
         timing.record_scan_work(100, Some(&dense_rows), 3, 8, true);
 
         assert_eq!(timing.unique_list_rows, 200);
-        assert_eq!(timing.matched_rows, 30);
-        assert_eq!(timing.pq_codes_evaluated, 100);
+        assert_eq!(timing.matched_rows, 36);
+        assert_eq!(timing.pq_codes_evaluated, 118);
         assert_eq!(timing.sparse_query_list_pairs, 4);
         assert_eq!(timing.dense_query_list_pairs, 3);
         assert_eq!(timing.actual_pq_codes_evaluated, 340);

Reply via email to