This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/paimon-rust.git


The following commit(s) were added to refs/heads/main by this push:
     new 3f8acf0b perf: vectorize raw vector search (#734)
3f8acf0b is described below

commit 3f8acf0b0e6992680e960b85b605a8827c710dff
Author: Jingsong Lee <[email protected]>
AuthorDate: Thu Aug 20 18:06:36 2026 +0800

    perf: vectorize raw vector search (#734)
---
 crates/paimon/src/table/vector_search_builder.rs | 562 +++++++++++++++++++----
 1 file changed, 468 insertions(+), 94 deletions(-)

diff --git a/crates/paimon/src/table/vector_search_builder.rs 
b/crates/paimon/src/table/vector_search_builder.rs
index 4fb89bf8..4692b6c1 100644
--- a/crates/paimon/src/table/vector_search_builder.rs
+++ b/crates/paimon/src/table/vector_search_builder.rs
@@ -64,18 +64,22 @@ use crate::vindex::{is_vindex_index_type, 
vector_search_timing_enabled, VindexVe
 use arrow_array::{Array, FixedSizeListArray, Float32Array, Int64Array, 
ListArray, RecordBatch};
 use arrow_select::interleave::interleave_record_batch;
 use futures::{stream, TryStreamExt};
+use paimon_vindex_core::blas::sgemm_a_bt;
 use paimon_vindex_core::diskann_io::DISKANN_HEADER_SIZE;
 use paimon_vindex_core::distance::MetricType;
 use paimon_vindex_core::index::VectorIndexReader as VIndexReader;
 use paimon_vindex_core::io::SeekRead;
 use roaring::RoaringTreemap;
 use std::cmp::Ordering;
-use std::collections::{BinaryHeap, HashMap, HashSet};
+use std::collections::{HashMap, HashSet};
 use std::io::Cursor;
 use std::sync::Arc;
 use std::time::{Duration, Instant};
 
 const INDEX_DIR: &str = "index";
+const RAW_SCORE_MATRIX_MIN_QUERY_COUNT: usize = 4;
+const RAW_SCORE_MATRIX_TARGET_ELEMENTS: usize = 1 << 20;
+const RAW_TOP_K_MIN_PARTITION_SIZE: usize = 1 << 12;
 
 #[derive(Clone, Copy, Debug, PartialEq, Eq)]
 enum VectorIndexBackend {
@@ -3020,23 +3024,24 @@ async fn read_raw_batch_vector_search(
 struct RawScoringPlan {
     all_query_indices: Vec<usize>,
     candidate_query_indices: HashMap<u64, Vec<usize>>,
-    query_l2_norms: Vec<f32>,
+    query_l2_squared_norms: Vec<f32>,
+    dense_query_dimension: Option<usize>,
+    dense_query_matrix: Option<Vec<f32>>,
 }
 
 impl RawScoringPlan {
     fn new(vector_searches: &[VectorSearch], metric: RawVectorMetric) -> Self {
         let mut all_query_indices = Vec::new();
         let mut candidate_query_indices: HashMap<u64, Vec<usize>> = 
HashMap::new();
-        let query_l2_norms = vector_searches
+        let query_l2_squared_norms = vector_searches
             .iter()
             .map(|vector_search| match metric {
-                RawVectorMetric::Cosine => vector_search
+                RawVectorMetric::L2 | RawVectorMetric::Cosine => vector_search
                     .vector
                     .iter()
                     .map(|value| value * value)
-                    .sum::<f32>()
-                    .sqrt(),
-                RawVectorMetric::L2 | RawVectorMetric::InnerProduct => 0.0,
+                    .sum::<f32>(),
+                RawVectorMetric::InnerProduct => 0.0,
             })
             .collect();
 
@@ -3053,10 +3058,29 @@ impl RawScoringPlan {
             }
         }
 
+        let dense_query_dimension = all_query_indices
+            .first()
+            .map(|&query_index| vector_searches[query_index].vector.len());
+        let dense_query_matrix = dense_query_dimension.and_then(|dimension| {
+            all_query_indices
+                .iter()
+                .all(|&query_index| vector_searches[query_index].vector.len() 
== dimension)
+                .then(|| {
+                    let mut matrix =
+                        
Vec::with_capacity(all_query_indices.len().saturating_mul(dimension));
+                    for &query_index in &all_query_indices {
+                        
matrix.extend_from_slice(&vector_searches[query_index].vector);
+                    }
+                    matrix
+                })
+        });
+
         Self {
             all_query_indices,
             candidate_query_indices,
-            query_l2_norms,
+            query_l2_squared_norms,
+            dense_query_dimension,
+            dense_query_matrix,
         }
     }
 }
@@ -3067,42 +3091,24 @@ struct RawScoredRow {
     score: f32,
 }
 
-impl Eq for RawScoredRow {}
-
-impl PartialOrd for RawScoredRow {
-    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
-        Some(self.cmp(other))
-    }
-}
-
-impl Ord for RawScoredRow {
-    fn cmp(&self, other: &Self) -> Ordering {
-        other
-            .score
-            .total_cmp(&self.score)
-            .then_with(|| self.row_id.cmp(&other.row_id))
-    }
-}
-
 impl RawScoredRow {
-    fn is_stronger_than(&self, other: &Self) -> bool {
-        self.score
-            .total_cmp(&other.score)
-            .then_with(|| other.row_id.cmp(&self.row_id))
-            == Ordering::Greater
+    fn strongest_first(a: &Self, b: &Self) -> Ordering {
+        b.score
+            .total_cmp(&a.score)
+            .then_with(|| a.row_id.cmp(&b.row_id))
     }
 }
 
 struct RawScoreTopK {
     limit: usize,
-    heap: BinaryHeap<RawScoredRow>,
+    candidates: Vec<RawScoredRow>,
 }
 
 impl RawScoreTopK {
     fn new(limit: usize) -> Self {
         Self {
             limit,
-            heap: BinaryHeap::with_capacity(limit.min(1024).saturating_add(1)),
+            candidates: Vec::with_capacity(limit.min(1024).saturating_add(1)),
         }
     }
 
@@ -3110,26 +3116,48 @@ impl RawScoreTopK {
         if self.limit == 0 {
             return;
         }
-        let entry = RawScoredRow { row_id, score };
-        if self.heap.len() < self.limit {
-            self.heap.push(entry);
-        } else if self
-            .heap
-            .peek()
-            .is_some_and(|weakest| entry.is_stronger_than(weakest))
-        {
-            self.heap.pop();
-            self.heap.push(entry);
+        self.candidates.push(RawScoredRow { row_id, score });
+        if self.candidates.len() >= self.partition_size() {
+            self.reduce_to_limit();
         }
     }
 
-    fn into_search_result(self) -> SearchResult {
-        let mut rows = self.heap.into_vec();
-        rows.sort_by(|a, b| {
-            b.score
-                .total_cmp(&a.score)
-                .then_with(|| a.row_id.cmp(&b.row_id))
-        });
+    fn offer_many<I>(&mut self, candidates: I)
+    where
+        I: IntoIterator<Item = RawScoredRow>,
+    {
+        if self.limit == 0 {
+            return;
+        }
+        self.candidates.extend(candidates);
+        if self.candidates.len() >= self.partition_size() {
+            self.reduce_to_limit();
+        }
+    }
+
+    fn partition_size(&self) -> usize {
+        self.limit
+            .saturating_mul(2)
+            .max(RAW_TOP_K_MIN_PARTITION_SIZE)
+    }
+
+    fn reduce_to_limit(&mut self) {
+        if self.candidates.len() <= self.limit {
+            return;
+        }
+        // Partition only after a substantial candidate block has accumulated.
+        // Each partition is linear in its input, so all reductions are O(n)
+        // amortized; only the final K survivors are fully sorted.
+        self.candidates
+            .select_nth_unstable_by(self.limit, RawScoredRow::strongest_first);
+        self.candidates.truncate(self.limit);
+    }
+
+    fn into_search_result(mut self) -> SearchResult {
+        self.reduce_to_limit();
+        self.candidates
+            .sort_unstable_by(RawScoredRow::strongest_first);
+        let rows = self.candidates;
         let mut row_ids = Vec::with_capacity(rows.len());
         let mut scores = Vec::with_capacity(rows.len());
         for row in rows {
@@ -3226,6 +3254,16 @@ fn collect_raw_batch_vector_batch(
         source: None,
     })?;
 
+    let use_dense_matrix = scoring_plan.all_query_indices.len() >= 
RAW_SCORE_MATRIX_MIN_QUERY_COUNT;
+    let dense_dimension = use_dense_matrix
+        .then_some(scoring_plan.dense_query_dimension)
+        .flatten();
+    let mut dense_row_ids = Vec::with_capacity(batch.num_rows());
+    let mut dense_vectors = Vec::with_capacity(
+        batch
+            .num_rows()
+            .saturating_mul(dense_dimension.unwrap_or_default()),
+    );
     for row in 0..batch.num_rows() {
         if row_ids.is_null(row) {
             return Err(crate::Error::DataInvalid {
@@ -3249,7 +3287,8 @@ fn collect_raw_batch_vector_batch(
             }
             VectorLayout::Fixed(a) => {
                 let len = a.value_length() as usize;
-                (row * len, (row + 1) * len)
+                let start = a.value_offset(row) as usize;
+                (start, start + len)
             }
         };
         ensure_raw_vector_values_not_null(values, start, end)?;
@@ -3260,15 +3299,29 @@ fn collect_raw_batch_vector_batch(
             start,
             end,
         };
-        for &query_index in &scoring_plan.all_query_indices {
-            offer_raw_vector_score(
-                raw_row,
-                query_index,
-                metric,
-                vector_searches,
-                scoring_plan,
-                top_k_out,
-            )?;
+        if let Some(dimension) = dense_dimension {
+            ensure_raw_vector_dimension(end - start, dimension)?;
+            if scoring_plan.dense_query_matrix.is_none() {
+                let &query_index = scoring_plan
+                    .all_query_indices
+                    .iter()
+                    .find(|&&query_index| 
vector_searches[query_index].vector.len() != dimension)
+                    .expect("a missing dense matrix requires inconsistent 
query dimensions");
+                ensure_raw_vector_dimension(dimension, 
vector_searches[query_index].vector.len())?;
+            }
+            dense_row_ids.push(row_id);
+            dense_vectors.extend_from_slice(&values.values()[start..end]);
+        } else {
+            for &query_index in &scoring_plan.all_query_indices {
+                offer_raw_vector_score(
+                    raw_row,
+                    query_index,
+                    metric,
+                    vector_searches,
+                    scoring_plan,
+                    top_k_out,
+                )?;
+            }
         }
         if let Some(query_indices) = 
scoring_plan.candidate_query_indices.get(&row_id) {
             for &query_index in query_indices {
@@ -3284,14 +3337,154 @@ fn collect_raw_batch_vector_batch(
         }
     }
 
+    if !dense_row_ids.is_empty() {
+        let query_matrix = scoring_plan
+            .dense_query_matrix
+            .as_deref()
+            .expect("dense query dimensions were validated above");
+        let dimension = dense_dimension.expect("dense rows require dense 
queries");
+        let queries_per_chunk = (RAW_SCORE_MATRIX_TARGET_ELEMENTS / 
dense_row_ids.len())
+            .max(1)
+            .min(scoring_plan.all_query_indices.len());
+        for (query_chunk_index, query_indices) in scoring_plan
+            .all_query_indices
+            .chunks(queries_per_chunk)
+            .enumerate()
+        {
+            let query_start = query_chunk_index * queries_per_chunk * 
dimension;
+            let query_end = query_start + query_indices.len() * dimension;
+            let scores = compute_raw_vector_score_matrix(
+                &dense_vectors,
+                dense_row_ids.len(),
+                &query_matrix[query_start..query_end],
+                query_indices.len(),
+                dimension,
+                &scoring_plan.query_l2_squared_norms,
+                query_indices,
+                metric,
+            )?;
+            for (matrix_query_index, &query_index) in 
query_indices.iter().enumerate() {
+                let query_scores = &scores[matrix_query_index * 
dense_row_ids.len()
+                    ..(matrix_query_index + 1) * dense_row_ids.len()];
+                top_k_out[query_index].offer_many(
+                    dense_row_ids
+                        .iter()
+                        .zip(query_scores)
+                        .map(|(&row_id, &score)| RawScoredRow { row_id, score 
}),
+                );
+            }
+        }
+    }
+
+    Ok(())
+}
+
+fn ensure_raw_vector_dimension(stored_len: usize, query_len: usize) -> 
crate::Result<()> {
+    if stored_len != query_len {
+        return Err(crate::Error::DataInvalid {
+            message: format!(
+                "Query vector dimension mismatch: raw row has {}, but query 
has {}",
+                stored_len, query_len
+            ),
+            source: None,
+        });
+    }
     Ok(())
 }
 
+#[allow(clippy::too_many_arguments)]
+fn compute_raw_vector_score_matrix(
+    stored_vectors: &[f32],
+    row_count: usize,
+    query_vectors: &[f32],
+    query_count: usize,
+    dimension: usize,
+    query_l2_squared_norms: &[f32],
+    query_indices: &[usize],
+    metric: RawVectorMetric,
+) -> crate::Result<Vec<f32>> {
+    let score_count =
+        row_count
+            .checked_mul(query_count)
+            .ok_or_else(|| crate::Error::DataInvalid {
+                message: "Vector raw search score matrix is too 
large".to_string(),
+                source: None,
+            })?;
+    debug_assert_eq!(stored_vectors.len(), row_count * dimension);
+    debug_assert_eq!(query_vectors.len(), query_count * dimension);
+    debug_assert_eq!(query_indices.len(), query_count);
+
+    let mut scores = vec![0.0; score_count];
+    // Query × stored-vector^T produces a query-major score matrix. Each 
query's
+    // scores are contiguous, which feeds partial Top-K without strided reads.
+    sgemm_a_bt(
+        query_count,
+        row_count,
+        dimension,
+        1.0,
+        query_vectors,
+        stored_vectors,
+        0.0,
+        &mut scores,
+    );
+    if metric == RawVectorMetric::InnerProduct {
+        return Ok(scores);
+    }
+
+    let stored_l2_squared_norms = stored_vectors
+        .chunks_exact(dimension)
+        .map(|vector| vector.iter().map(|value| value * value).sum::<f32>())
+        .collect::<Vec<_>>();
+    for (matrix_query_index, &query_index) in query_indices.iter().enumerate() 
{
+        for (row_index, &stored_l2_squared_norm) in 
stored_l2_squared_norms.iter().enumerate() {
+            let score = &mut scores[matrix_query_index * row_count + 
row_index];
+            let query_l2_squared_norm = query_l2_squared_norms[query_index];
+            *score = match metric {
+                RawVectorMetric::L2 => {
+                    let squared_distance =
+                        stored_l2_squared_norm + query_l2_squared_norm - 2.0 * 
*score;
+                    // The norm/dot reconstruction loses the low-order 
difference when two
+                    // large vectors are close. Estimate a conservative 
accumulation-error
+                    // bound and preserve the former scalar semantics inside 
that region.
+                    let roundoff_bound = (stored_l2_squared_norm.abs()
+                        + query_l2_squared_norm.abs()
+                        + 2.0 * score.abs())
+                        * f32::EPSILON
+                        * (dimension as f32 + 2.0)
+                        * 4.0;
+                    if !squared_distance.is_finite() || squared_distance <= 
roundoff_bound {
+                        let stored =
+                            &stored_vectors[row_index * dimension..(row_index 
+ 1) * dimension];
+                        let query = &query_vectors
+                            [matrix_query_index * 
dimension..(matrix_query_index + 1) * dimension];
+                        compute_raw_vector_l2_score(query, stored)
+                    } else {
+                        1.0 / (1.0 + squared_distance)
+                    }
+                }
+                RawVectorMetric::Cosine => {
+                    let denominator = stored_l2_squared_norm.sqrt() * 
query_l2_squared_norm.sqrt();
+                    if denominator == 0.0 {
+                        0.0
+                    } else {
+                        *score / denominator
+                    }
+                }
+                RawVectorMetric::InnerProduct => unreachable!(),
+            };
+        }
+    }
+    Ok(scores)
+}
+
 fn ensure_raw_vector_values_not_null(
     values: &Float32Array,
     start: usize,
     end: usize,
 ) -> crate::Result<()> {
+    if values.null_count() == 0 {
+        return Ok(());
+    }
     for value_index in start..end {
         if values.is_null(value_index) {
             return Err(crate::Error::DataInvalid {
@@ -3321,19 +3514,10 @@ fn offer_raw_vector_score(
 ) -> crate::Result<()> {
     let vector_search = &vector_searches[query_index];
     let stored_len = row.end - row.start;
-    if stored_len != vector_search.vector.len() {
-        return Err(crate::Error::DataInvalid {
-            message: format!(
-                "Query vector dimension mismatch: raw row has {}, but query 
has {}",
-                stored_len,
-                vector_search.vector.len()
-            ),
-            source: None,
-        });
-    }
+    ensure_raw_vector_dimension(stored_len, vector_search.vector.len())?;
     let score = compute_raw_vector_score_from_values(
         &vector_search.vector,
-        scoring_plan.query_l2_norms[query_index],
+        scoring_plan.query_l2_squared_norms[query_index],
         row.values,
         row.start,
         row.end,
@@ -3345,7 +3529,7 @@ fn offer_raw_vector_score(
 
 fn compute_raw_vector_score_from_values(
     query: &[f32],
-    query_l2_norm: f32,
+    query_l2_squared_norm: f32,
     values: &Float32Array,
     start: usize,
     end: usize,
@@ -3353,17 +3537,7 @@ fn compute_raw_vector_score_from_values(
 ) -> f32 {
     debug_assert_eq!(query.len(), end - start);
     match metric {
-        RawVectorMetric::L2 => {
-            let sum_sq = query
-                .iter()
-                .zip(start..end)
-                .map(|(q, value_index)| {
-                    let diff = q - values.value(value_index);
-                    diff * diff
-                })
-                .sum::<f32>();
-            1.0 / (1.0 + sum_sq)
-        }
+        RawVectorMetric::L2 => compute_raw_vector_l2_score(query, 
&values.values()[start..end]),
         RawVectorMetric::Cosine => {
             let mut dot = 0.0;
             let mut norm_b = 0.0;
@@ -3372,7 +3546,7 @@ fn compute_raw_vector_score_from_values(
                 dot += q * stored;
                 norm_b += stored * stored;
             }
-            let denominator = query_l2_norm * norm_b.sqrt();
+            let denominator = query_l2_squared_norm.sqrt() * norm_b.sqrt();
             if denominator == 0.0 {
                 0.0
             } else {
@@ -3387,6 +3561,18 @@ fn compute_raw_vector_score_from_values(
     }
 }
 
+fn compute_raw_vector_l2_score(query: &[f32], stored: &[f32]) -> f32 {
+    let squared_distance = query
+        .iter()
+        .zip(stored)
+        .map(|(query_value, stored_value)| {
+            let difference = query_value - stored_value;
+            difference * difference
+        })
+        .sum::<f32>();
+    1.0 / (1.0 + squared_distance)
+}
+
 fn row_id_to_u64(row_id: i64) -> crate::Result<u64> {
     u64::try_from(row_id).map_err(|_| crate::Error::DataInvalid {
         message: format!("Negative _ROW_ID {row_id} cannot be used for global 
index search"),
@@ -3397,17 +3583,7 @@ fn row_id_to_u64(row_id: i64) -> crate::Result<u64> {
 #[cfg(test)]
 fn compute_raw_vector_score(query: &[f32], stored: &[f32], metric: 
RawVectorMetric) -> f32 {
     match metric {
-        RawVectorMetric::L2 => {
-            let sum_sq = query
-                .iter()
-                .zip(stored.iter())
-                .map(|(q, s)| {
-                    let diff = q - s;
-                    diff * diff
-                })
-                .sum::<f32>();
-            1.0 / (1.0 + sum_sq)
-        }
+        RawVectorMetric::L2 => compute_raw_vector_l2_score(query, stored),
         RawVectorMetric::Cosine => {
             let mut dot = 0.0;
             let mut norm_a = 0.0;
@@ -3587,6 +3763,154 @@ mod tests {
         );
     }
 
+    #[test]
+    fn test_raw_vector_score_matrix_matches_scalar_metrics() {
+        let stored = vec![1.0, 2.0, 3.0, 4.0, 0.0, 0.0];
+        let queries = vec![1.0, 1.0, -1.0, 2.0];
+        let query_indices = vec![0, 1];
+        let query_l2_squared_norms = vec![2.0, 5.0];
+
+        for metric in [
+            RawVectorMetric::L2,
+            RawVectorMetric::Cosine,
+            RawVectorMetric::InnerProduct,
+        ] {
+            let matrix_scores = compute_raw_vector_score_matrix(
+                &stored,
+                3,
+                &queries,
+                2,
+                2,
+                &query_l2_squared_norms,
+                &query_indices,
+                metric,
+            )
+            .unwrap();
+            for (row_index, stored_vector) in 
stored.chunks_exact(2).enumerate() {
+                for (query_index, query) in 
queries.chunks_exact(2).enumerate() {
+                    let expected = compute_raw_vector_score(query, 
stored_vector, metric);
+                    let actual = matrix_scores[query_index * 3 + row_index];
+                    assert!(
+                        (actual - expected).abs() < 1e-5,
+                        "metric={metric:?}, row={row_index}, 
query={query_index}: {actual} != {expected}"
+                    );
+                }
+            }
+        }
+
+        let non_finite_score = compute_raw_vector_score_matrix(
+            &[f32::INFINITY, 0.0],
+            1,
+            &[1.0, 0.0],
+            1,
+            2,
+            &[1.0],
+            &[0],
+            RawVectorMetric::L2,
+        )
+        .unwrap()[0];
+        assert_eq!(non_finite_score, 0.0);
+    }
+
+    #[test]
+    fn test_raw_vector_score_matrix_l2_preserves_large_finite_distances() {
+        let dimension = 128;
+        let query = vec![1.0e10_f32; dimension];
+        let mut nearby = query.clone();
+        nearby[0] += 1024.0;
+        let mut stored = query.clone();
+        stored.extend_from_slice(&nearby);
+        let queries = query.repeat(4);
+        let query_l2_squared_norm = query.iter().map(|value| value * 
value).sum::<f32>();
+        let query_l2_squared_norms = vec![query_l2_squared_norm; 4];
+        let query_indices = vec![0, 1, 2, 3];
+
+        let matrix_scores = compute_raw_vector_score_matrix(
+            &stored,
+            2,
+            &queries,
+            4,
+            dimension,
+            &query_l2_squared_norms,
+            &query_indices,
+            RawVectorMetric::L2,
+        )
+        .unwrap();
+        let exact_score = compute_raw_vector_score(&query, &query, 
RawVectorMetric::L2);
+        let nearby_score = compute_raw_vector_score(&query, &nearby, 
RawVectorMetric::L2);
+
+        for query_index in 0..4 {
+            assert_eq!(matrix_scores[query_index * 2], exact_score);
+            assert_eq!(matrix_scores[query_index * 2 + 1], nearby_score);
+            assert!(matrix_scores[query_index * 2] > matrix_scores[query_index 
* 2 + 1]);
+        }
+    }
+
+    #[test]
+    fn test_raw_vector_cosine_avoids_squared_norm_product_overflow() {
+        let query = vec![1.0e15_f32, 0.0];
+        let query_l2_squared_norm = query.iter().map(|value| value * 
value).sum::<f32>();
+        assert!(query_l2_squared_norm.is_finite());
+        let values = Float32Array::from(query.clone());
+        let scalar_score = compute_raw_vector_score_from_values(
+            &query,
+            query_l2_squared_norm,
+            &values,
+            0,
+            2,
+            RawVectorMetric::Cosine,
+        );
+        assert!((scalar_score - 1.0).abs() < 1e-6);
+
+        let queries = query.repeat(4);
+        let matrix_scores = compute_raw_vector_score_matrix(
+            &query,
+            1,
+            &queries,
+            4,
+            2,
+            &[query_l2_squared_norm; 4],
+            &[0, 1, 2, 3],
+            RawVectorMetric::Cosine,
+        )
+        .unwrap();
+        assert!(matrix_scores
+            .iter()
+            .all(|score| (*score - 1.0).abs() < 1e-6));
+    }
+
+    #[test]
+    fn test_raw_score_top_k_matches_full_sort_with_linear_partial_selection() {
+        let limit = 7;
+        let mut top_k = RawScoreTopK::new(limit);
+        let mut batched_top_k = RawScoreTopK::new(limit);
+        let mut expected = Vec::new();
+        for row_id in 0..10_000 {
+            let score = ((row_id * 37) % 101) as f32 / 10.0;
+            let candidate = RawScoredRow { row_id, score };
+            expected.push(candidate);
+            top_k.offer(row_id, score);
+            batched_top_k.offer_many(std::iter::once(candidate));
+            assert!(top_k.candidates.len() < top_k.partition_size());
+            assert!(batched_top_k.candidates.len() < 
batched_top_k.partition_size());
+        }
+        expected.sort_unstable_by(RawScoredRow::strongest_first);
+        expected.truncate(limit);
+
+        let result = top_k.into_search_result();
+        let batched_result = batched_top_k.into_search_result();
+        assert_eq!(
+            result.row_ids,
+            expected.iter().map(|row| row.row_id).collect::<Vec<_>>()
+        );
+        assert_eq!(
+            result.scores,
+            expected.iter().map(|row| row.score).collect::<Vec<_>>()
+        );
+        assert_eq!(batched_result.row_ids, result.row_ids);
+        assert_eq!(batched_result.scores, result.scores);
+    }
+
     #[test]
     fn test_configured_raw_vector_metric_precedence_and_conflict_default() {
         let mut options = HashMap::new();
@@ -3780,6 +4104,8 @@ mod tests {
         let searches = vec![
             VectorSearch::new(vec![1.0, 0.0], 1, 
"embedding".to_string()).unwrap(),
             VectorSearch::new(vec![0.0, 1.0], 1, 
"embedding".to_string()).unwrap(),
+            VectorSearch::new(vec![0.8, 0.2], 1, 
"embedding".to_string()).unwrap(),
+            VectorSearch::new(vec![0.5, 0.5], 1, 
"embedding".to_string()).unwrap(),
         ];
         let scoring_plan = RawScoringPlan::new(&searches, RawVectorMetric::L2);
         let mut top_k = searches
@@ -3802,6 +4128,54 @@ mod tests {
 
         assert_eq!(results[0].row_ids, vec![10]);
         assert_eq!(results[1].row_ids, vec![11]);
+        assert_eq!(results[2].row_ids, vec![12]);
+        assert_eq!(results[3].row_ids, vec![12]);
+    }
+
+    #[test]
+    fn test_collect_raw_batch_vector_batch_respects_fixed_size_list_offset() {
+        let element_field = Arc::new(ArrowField::new("element", 
ArrowDataType::Float32, true));
+        let mut builder =
+            FixedSizeListBuilder::new(Float32Builder::new(), 
2).with_field(element_field);
+        for vector in [[1.0, 0.0], [0.0, 1.0], [0.8, 0.2]] {
+            builder.values().append_value(vector[0]);
+            builder.values().append_value(vector[1]);
+            builder.append(true);
+        }
+        let schema = Arc::new(ArrowSchema::new(vec![
+            ArrowField::new(
+                "embedding",
+                ArrowDataType::FixedSizeList(
+                    Arc::new(ArrowField::new("element", 
ArrowDataType::Float32, true)),
+                    2,
+                ),
+                true,
+            ),
+            ArrowField::new(ROW_ID_FIELD_NAME, ArrowDataType::Int64, false),
+        ]));
+        let batch = RecordBatch::try_new(
+            schema,
+            vec![
+                Arc::new(builder.finish()) as ArrayRef,
+                Arc::new(Int64Array::from(vec![10, 11, 12])) as ArrayRef,
+            ],
+        )
+        .unwrap()
+        .slice(1, 2);
+        let searches = vec![VectorSearch::new(vec![0.0, 1.0], 1, 
"embedding".to_string()).unwrap()];
+        let scoring_plan = RawScoringPlan::new(&searches, RawVectorMetric::L2);
+        let mut top_k = vec![RawScoreTopK::new(1)];
+
+        collect_raw_batch_vector_batch(
+            &batch,
+            &searches,
+            RawVectorMetric::L2,
+            &scoring_plan,
+            &mut top_k,
+        )
+        .unwrap();
+
+        assert_eq!(top_k.pop().unwrap().into_search_result().row_ids, 
vec![11]);
     }
 
     #[test]

Reply via email to