From 1226b7099dff68dcd4a7a0907a49a2562e4d2f00 Mon Sep 17 00:00:00 2001 From: zoey Date: Wed, 27 May 2026 00:44:55 +0800 Subject: [PATCH 1/8] perf(knn): reduce memory for batch flat vector search Avoid retaining whole scan RecordBatches (and vectors) in batch flat KNN heaps when the final projection does not include the vector column. When vectors are requested, retain only a per-row copy. --- rust/lance/src/dataset/scanner.rs | 28 ++ rust/lance/src/io/exec/knn.rs | 427 ++++++++++++++++++++++++++---- 2 files changed, 410 insertions(+), 45 deletions(-) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index d0f5983b3ce..76c9b040340 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -4424,6 +4424,14 @@ impl Scanner { } else { input }; + let retain_vector = if self.is_batch_nearest { + let vector_field_id = self.dataset.schema().field_id(q.column.as_str())?; + self.projection_plan + .physical_projection + .contains_field_id(vector_field_id) + } else { + false + }; let flat_dist = Arc::new(KNNVectorDistanceExec::try_new_batch( input, &q.column, @@ -4435,6 +4443,7 @@ impl Scanner { lower_bound: q.lower_bound, upper_bound: q.upper_bound, distance_type: metric_type, + retain_vector, }, )?); @@ -5975,6 +5984,25 @@ mod test { } assert_batch_matches_single_queries(dataset, &batch, &query_values, k, false, None).await; + let mut scan_with_vec = dataset.scan(); + scan_with_vec.nearest("vec", &queries, k).unwrap(); + scan_with_vec.use_index(false); + scan_with_vec.project(&["i", "vec"]).unwrap(); + let batch_with_vec = scan_with_vec.try_into_batch().await.unwrap(); + assert!( + batch_with_vec.schema().column_with_name("vec").is_some(), + "batch flat KNN should return vector column when projected" + ); + assert_batch_matches_single_queries( + dataset, + &batch_with_vec, + &query_values, + k, + false, + None, + ) + .await; + let query_values_one = (32..64).map(|v| v as f32).collect::>(); let queries_one = FixedSizeListArray::try_new_from_values( Float32Array::from(query_values_one.clone()), diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index 71239b4e34b..665b8b53865 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -156,6 +156,7 @@ pub struct KNNVectorDistanceExec { pub upper_bound: Option, pub column: String, pub distance_type: DistanceType, + retain_vector: bool, input_schema: SchemaRef, output_schema: SchemaRef, @@ -171,10 +172,11 @@ pub struct KnnBatchParams { pub lower_bound: Option, pub upper_bound: Option, pub distance_type: DistanceType, + pub retain_vector: bool, } struct BatchKnnConfig { - input_schema: SchemaRef, + stored_schema: SchemaRef, output_schema: SchemaRef, column: String, query: ArrayRef, @@ -183,6 +185,7 @@ struct BatchKnnConfig { lower_bound: Option, upper_bound: Option, distance_type: DistanceType, + retain_vector: bool, } impl DisplayAs for KNNVectorDistanceExec { @@ -235,6 +238,7 @@ impl KNNVectorDistanceExec { lower_bound: None, upper_bound: None, distance_type, + retain_vector: false, }, ) } @@ -252,6 +256,7 @@ impl KNNVectorDistanceExec { lower_bound, upper_bound, distance_type, + retain_vector, } = params; if query_count == 0 { return Err(Error::invalid_input( @@ -286,13 +291,19 @@ impl KNNVectorDistanceExec { "batch KNN cannot run when the input already contains reserved column '{QUERY_INDEX_COL}'" ))); } - let input_schema = Arc::new(input_schema); + + let stored_schema = if is_batch && !retain_vector { + Arc::new(Schema::new(vec![ROW_ID_FIELD.clone()])) + } else { + Arc::new(input_schema) + }; + let output_schema = if is_batch { - input_schema + stored_schema .as_ref() .try_with_column_at(0, query_index_field())? } else { - input_schema.as_ref().clone() + stored_schema.as_ref().clone() }; let output_schema = Arc::new(output_schema.try_with_column(Field::new( DIST_COL, @@ -329,19 +340,104 @@ impl KNNVectorDistanceExec { upper_bound, column: column.to_string(), distance_type, - input_schema, + retain_vector, + input_schema: stored_schema, output_schema, properties, metrics: ExecutionPlanMetricsSet::new(), }) } + fn take_vector_row( + batch: &RecordBatch, + column: &str, + row_index: u32, + ) -> DataFusionResult { + let vectors = batch.column_by_name(column).ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected vector column '{column}' in scan batch" + )) + })?; + let indices = UInt32Array::from(vec![row_index]); + arrow::compute::take(vectors, &indices, None) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) + } + + fn assemble_batch_output( + results: &[BatchKnnCandidate], + stored_schema: &Schema, + column: &str, + ) -> DataFusionResult { + let mut columns: Vec = Vec::with_capacity(stored_schema.fields().len()); + for field in stored_schema.fields() { + if field.name() == ROW_ID { + let row_ids = UInt64Array::from_iter( + results.iter().map(|candidate| Some(candidate.row_id())), + ); + columns.push(Arc::new(row_ids)); + } else if field.name() == column { + let vector_rows: Vec<&dyn Array> = results + .iter() + .map(|candidate| { + let BatchKnnCandidate::WithVector { vector_row, .. } = candidate else { + return Err(DataFusionError::Internal( + "batch KNN expected vector rows in candidate heap".to_string(), + )); + }; + Ok(vector_row.as_ref()) + }) + .collect::>>()?; + let vector_column = arrow::compute::concat(&vector_rows) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + columns.push(vector_column); + } else { + let mut row_batches = Vec::with_capacity(results.len()); + for candidate in results { + let BatchKnnCandidate::WithVector { + batch, row_index, .. + } = candidate + else { + return Err(DataFusionError::Internal( + "batch KNN expected slim batch in candidate heap".to_string(), + )); + }; + let indices = UInt32Array::from(vec![*row_index]); + row_batches.push( + arrow_select::take::take_record_batch(batch.as_ref(), &indices).map_err( + |e| { + DataFusionError::ArrowError( + Box::new(e), + Some("take top-k row".to_string()), + ) + }, + )?, + ); + } + let taken = concat_batches(&row_batches[0].schema(), &row_batches) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + columns.push( + taken + .column_by_name(field.name()) + .ok_or_else(|| { + DataFusionError::Internal(format!( + "column '{}' missing from slim batch", + field.name() + )) + })? + .clone(), + ); + } + } + RecordBatch::try_new(Arc::new(stored_schema.clone()), columns) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) + } + async fn execute_batch( input: SendableRecordBatchStream, config: BatchKnnConfig, ) -> DataFusionResult { let BatchKnnConfig { - input_schema, + stored_schema, output_schema, column, query, @@ -350,6 +446,7 @@ impl KNNVectorDistanceExec { lower_bound, upper_bound, distance_type, + retain_vector, } = config; let query_dim = query.len() / query_count; let mut heaps = (0..query_count) @@ -373,6 +470,16 @@ impl KNNVectorDistanceExec { .as_primitive::() .clone(); + let slim_batch = if retain_vector { + Some(Arc::new( + batch + .drop_column(&column) + .map_err(|e| DataFusionError::External(Box::new(e)))?, + )) + } else { + None + }; + for (query_index, heap) in heaps.iter_mut().enumerate().take(query_count) { let key = query.slice(query_index * query_dim, query_dim); let with_distances = compute_distance(key, distance_type, &column, batch.clone()) @@ -397,13 +504,23 @@ impl KNNVectorDistanceExec { } let query_index = query_index as i32; let row_id = row_ids.value(row_index); - let row_index = row_index as u32; - let candidate = BatchKnnCandidate { - query_index, - distance, - row_id, - batch: batch.clone(), - row_index, + let candidate = if retain_vector { + let row_index = row_index as u32; + let vector_row = Self::take_vector_row(&batch, &column, row_index)?; + BatchKnnCandidate::WithVector { + query_index, + distance, + row_id, + batch: Arc::clone(slim_batch.as_ref().expect("slim batch")), + row_index, + vector_row, + } + } else { + BatchKnnCandidate::RowIdOnly { + query_index, + distance, + row_id, + } }; if heap.len() < k { heap.push(candidate); @@ -423,10 +540,10 @@ impl KNNVectorDistanceExec { .flat_map(BinaryHeap::into_vec) .collect::>(); results.sort_by(|left, right| { - left.query_index - .cmp(&right.query_index) - .then_with(|| left.distance.total_cmp(&right.distance)) - .then_with(|| left.row_id.cmp(&right.row_id)) + left.query_index() + .cmp(&right.query_index()) + .then_with(|| left.distance().total_cmp(&right.distance())) + .then_with(|| left.row_id().cmp(&right.row_id())) }); if results.is_empty() { @@ -435,20 +552,19 @@ impl KNNVectorDistanceExec { let mut query_indices = Int32Builder::with_capacity(results.len()); let mut distances = Float32Builder::with_capacity(results.len()); - let mut row_batches = Vec::with_capacity(results.len()); - for result in results { - query_indices.append_value(result.query_index); - distances.append_value(result.distance); - let indices = UInt32Array::from(vec![result.row_index]); - row_batches.push( - arrow_select::take::take_record_batch(&result.batch, &indices).map_err(|e| { - DataFusionError::ArrowError(Box::new(e), Some("take top-k row".to_string())) - })?, - ); + for result in &results { + query_indices.append_value(result.query_index()); + distances.append_value(result.distance()); } - let output = concat_batches(&input_schema, &row_batches) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + let output = if retain_vector { + Self::assemble_batch_output(&results, stored_schema.as_ref(), &column)? + } else { + let row_ids = UInt64Array::from_iter(results.iter().map(|c| Some(c.row_id()))); + RecordBatch::try_new(stored_schema.clone(), vec![Arc::new(row_ids) as ArrayRef]) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))? + }; + output .try_with_column_at(0, query_index_field(), Arc::new(query_indices.finish())) .and_then(|batch| { @@ -500,6 +616,7 @@ impl ExecutionPlan for KNNVectorDistanceExec { lower_bound: self.lower_bound, upper_bound: self.upper_bound, distance_type: self.distance_type, + retain_vector: self.retain_vector, }, )?)) } @@ -514,7 +631,7 @@ impl ExecutionPlan for KNNVectorDistanceExec { let stream = stream::once(Self::execute_batch( input_stream, BatchKnnConfig { - input_schema: self.input_schema.clone(), + stored_schema: self.input_schema.clone(), output_schema: self.output_schema.clone(), column: self.column.clone(), query: self.query.clone(), @@ -523,6 +640,7 @@ impl ExecutionPlan for KNNVectorDistanceExec { lower_bound: self.lower_bound, upper_bound: self.upper_bound, distance_type: self.distance_type, + retain_vector: self.retain_vector, }, )); let schema = self.schema(); @@ -646,20 +764,49 @@ impl ExecutionPlan for KNNVectorDistanceExec { } #[derive(Clone)] -struct BatchKnnCandidate { - query_index: i32, - distance: f32, - row_id: u64, - batch: RecordBatch, - row_index: u32, +enum BatchKnnCandidate { + RowIdOnly { + query_index: i32, + distance: f32, + row_id: u64, + }, + WithVector { + query_index: i32, + distance: f32, + row_id: u64, + batch: Arc, + row_index: u32, + vector_row: ArrayRef, + }, +} + +impl BatchKnnCandidate { + fn query_index(&self) -> i32 { + match self { + Self::RowIdOnly { query_index, .. } | Self::WithVector { query_index, .. } => { + *query_index + } + } + } + + fn distance(&self) -> f32 { + match self { + Self::RowIdOnly { distance, .. } | Self::WithVector { distance, .. } => *distance, + } + } + + fn row_id(&self) -> u64 { + match self { + Self::RowIdOnly { row_id, .. } | Self::WithVector { row_id, .. } => *row_id, + } + } } impl PartialEq for BatchKnnCandidate { fn eq(&self, other: &Self) -> bool { - self.query_index == other.query_index - && self.distance == other.distance - && self.row_id == other.row_id - && self.row_index == other.row_index + self.query_index() == other.query_index() + && self.distance() == other.distance() + && self.row_id() == other.row_id() } } @@ -673,11 +820,201 @@ impl PartialOrd for BatchKnnCandidate { impl Ord for BatchKnnCandidate { fn cmp(&self, other: &Self) -> CmpOrdering { - self.distance - .total_cmp(&other.distance) - .then_with(|| self.row_id.cmp(&other.row_id)) - .then_with(|| self.query_index.cmp(&other.query_index)) - .then_with(|| self.row_index.cmp(&other.row_index)) + self.distance() + .total_cmp(&other.distance()) + .then_with(|| self.row_id().cmp(&other.row_id())) + .then_with(|| self.query_index().cmp(&other.query_index())) + } +} + +#[cfg(test)] +fn batch_knn_candidate_heap_memory(candidates: &[BatchKnnCandidate]) -> usize { + use std::collections::HashSet; + + let mut seen_slim_batches = HashSet::new(); + candidates + .iter() + .map(|candidate| match candidate { + BatchKnnCandidate::RowIdOnly { .. } => 0, + BatchKnnCandidate::WithVector { + batch, vector_row, .. + } => { + let slim_mem = if seen_slim_batches.insert(Arc::as_ptr(batch)) { + batch + .columns() + .iter() + .map(|col| col.get_array_memory_size()) + .sum::() + } else { + 0 + }; + slim_mem + vector_row.get_array_memory_size() + } + }) + .sum() +} + +#[cfg(test)] +fn batch_knn_baseline_heap_memory(candidates: &[(RecordBatch, u32)]) -> usize { + candidates + .iter() + .map(|(batch, _)| { + batch + .columns() + .iter() + .map(|col| col.get_array_memory_size()) + .sum::() + }) + .sum() +} + +#[cfg(test)] +mod batch_knn_memory_tests { + use std::time::Instant; + + use arrow::array::{FixedSizeListArray, Float32Array, UInt64Array}; + use arrow_array::RecordBatch; + use arrow_schema::{DataType, Field, Schema}; + + use super::*; + + fn make_scan_batch(num_rows: usize, dim: usize) -> RecordBatch { + let vectors = + Float32Array::from_iter((0..num_rows * dim).map(|i| (i % 1000) as f32 / 1000.0)); + let vectors = FixedSizeListArray::try_new_from_values(vectors, dim as i32).unwrap(); + let row_ids = UInt64Array::from_iter((0..num_rows as u64).map(Some)); + let schema = Arc::new(Schema::new(vec![ + Field::new( + "vec", + DataType::FixedSizeList( + Field::new("item", DataType::Float32, true).into(), + dim as i32, + ), + true, + ), + ROW_ID_FIELD.clone(), + ])); + RecordBatch::try_new( + schema, + vec![Arc::new(vectors) as ArrayRef, Arc::new(row_ids) as ArrayRef], + ) + .unwrap() + } + + #[test] + fn test_batch_knn_heap_memory() { + let num_rows = 4096; + let dim = 512; + let m = 10; + let k = 10; + let batch = make_scan_batch(num_rows, dim); + let slim_batch = Arc::new(batch.drop_column("vec").unwrap()); + + let mut baseline = Vec::with_capacity(m * k); + let mut optimized_row_id_only = Vec::with_capacity(m * k); + let mut optimized_with_vector = Vec::with_capacity(m * k); + + for query_index in 0..m { + for row_index in 0..k { + let row_index_u32 = (query_index * k + row_index) as u32; + baseline.push((batch.clone(), row_index_u32)); + optimized_row_id_only.push(BatchKnnCandidate::RowIdOnly { + query_index: query_index as i32, + distance: row_index as f32, + row_id: row_index_u32 as u64, + }); + let vector_row = + KNNVectorDistanceExec::take_vector_row(&batch, "vec", row_index_u32).unwrap(); + optimized_with_vector.push(BatchKnnCandidate::WithVector { + query_index: query_index as i32, + distance: row_index as f32, + row_id: row_index_u32 as u64, + batch: Arc::clone(&slim_batch), + row_index: row_index_u32, + vector_row, + }); + } + } + + let baseline_mem = batch_knn_baseline_heap_memory(&baseline); + let row_id_only_mem = batch_knn_candidate_heap_memory(&optimized_row_id_only); + let with_vector_mem = batch_knn_candidate_heap_memory(&optimized_with_vector); + + assert!( + row_id_only_mem < baseline_mem / 50, + "RowIdOnly heap memory ({row_id_only_mem}) should be much smaller than baseline ({baseline_mem})" + ); + assert!( + with_vector_mem < baseline_mem / 10, + "WithVector heap memory ({with_vector_mem}) should be much smaller than baseline ({baseline_mem})" + ); + + let expected_vector_bytes = m * k * dim * std::mem::size_of::(); + assert!( + with_vector_mem >= expected_vector_bytes / 2 + && with_vector_mem <= expected_vector_bytes * 4, + "WithVector heap memory ({with_vector_mem}) should be on the order of m*k*d*4 ({expected_vector_bytes})" + ); + } + + #[test] + fn test_batch_knn_assemble_perf_smoke() { + let num_rows = 4096; + let dim = 128; + let m = 10; + let k = 10; + let batch = make_scan_batch(num_rows, dim); + let slim_batch = Arc::new(batch.drop_column("vec").unwrap()); + let stored_schema = batch.schema(); + + let mut candidates = Vec::with_capacity(m * k); + for query_index in 0..m { + for row_index in 0..k { + let row_index_u32 = (query_index * k + row_index) as u32; + let vector_row = + KNNVectorDistanceExec::take_vector_row(&batch, "vec", row_index_u32).unwrap(); + candidates.push(BatchKnnCandidate::WithVector { + query_index: query_index as i32, + distance: 0.0, + row_id: row_index_u32 as u64, + batch: Arc::clone(&slim_batch), + row_index: row_index_u32, + vector_row, + }); + } + } + + let baseline_batches: Vec<_> = (0..candidates.len()) + .map(|i| (batch.clone(), i as u32)) + .collect(); + + let start = Instant::now(); + for _ in 0..50 { + let mut row_batches = Vec::with_capacity(baseline_batches.len()); + for (full_batch, row_index) in &baseline_batches { + let indices = UInt32Array::from(vec![*row_index]); + row_batches + .push(arrow_select::take::take_record_batch(full_batch, &indices).unwrap()); + } + let _ = concat_batches(&row_batches[0].schema(), &row_batches).unwrap(); + } + let baseline_elapsed = start.elapsed(); + + let start = Instant::now(); + for _ in 0..50 { + let _ = KNNVectorDistanceExec::assemble_batch_output( + &candidates, + stored_schema.as_ref(), + "vec", + ) + .unwrap(); + } + let optimized_elapsed = start.elapsed(); + + assert!( + optimized_elapsed <= baseline_elapsed * 2, + "optimized assemble ({optimized_elapsed:?}) should not be much slower than baseline ({baseline_elapsed:?})" + ); } } From 2699f9b5dc078f20094b55b59147e18bb59f0d57 Mon Sep 17 00:00:00 2001 From: zoey Date: Wed, 27 May 2026 13:22:02 +0800 Subject: [PATCH 2/8] fix(knn): address batch flat KNN review feedback Refactor heap candidates to struct + extra enum, defer vector and slim batch copies until rows enter top-k, and use slice/batched take for assembly. Add correctness tests asserting vector column is omitted when not projected. --- rust/lance/src/dataset/scanner.rs | 50 ++++ rust/lance/src/io/exec/knn.rs | 445 +++++++++--------------------- 2 files changed, 188 insertions(+), 307 deletions(-) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 76c9b040340..9c344c95778 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -5888,6 +5888,14 @@ mod test { assert!(!field.is_nullable()); } + fn assert_batch_knn_output_has_no_vector(batch: &RecordBatch) { + assert!( + batch.schema().column_with_name("vec").is_none(), + "batch flat KNN output must not include vector column when vec is not projected; columns: {:?}", + batch.schema().field_names() + ); + } + async fn assert_batch_matches_single_queries( dataset: &Dataset, batch: &RecordBatch, @@ -5962,6 +5970,7 @@ mod test { let batch = scan.try_into_batch().await.unwrap(); assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch); assert_eq!( batch.num_rows(), 2 * k, @@ -6028,12 +6037,53 @@ mod test { let batch = scan.try_into_batch().await.unwrap(); assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch); assert_eq!( batch[QUERY_INDEX_COL].as_primitive::().values(), &[0, 0] ); } + #[tokio::test] + async fn test_batch_knn_flat_omits_vector_without_projection() { + let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) + .await + .unwrap(); + let dataset = &test_ds.dataset; + let k = 2; + let (queries, query_values) = batch_knn_two_queries(); + + let mut scan = dataset.scan(); + scan.nearest("vec", &queries, k).unwrap(); + scan.use_index(false); + scan.project(&["i"]).unwrap(); + let batch = scan.try_into_batch().await.unwrap(); + assert_batch_knn_output_has_no_vector(&batch); + assert_query_index_field(&batch); + assert!(batch.schema().column_with_name("i").is_some()); + assert!(batch.schema().column_with_name(DIST_COL).is_some()); + assert_batch_matches_single_queries(dataset, &batch, &query_values, k, false, None).await; + + let mut scan_rowid_only = dataset.scan(); + scan_rowid_only.nearest("vec", &queries, k).unwrap(); + scan_rowid_only.use_index(false); + scan_rowid_only.project(&[ROW_ID]).unwrap(); + let batch_rowid_only = scan_rowid_only.try_into_batch().await.unwrap(); + assert_batch_knn_output_has_no_vector(&batch_rowid_only); + assert!(batch_rowid_only.schema().column_with_name(ROW_ID).is_some()); + assert!(batch_rowid_only.schema().column_with_name("i").is_none()); + + let mut scan_with_vec = dataset.scan(); + scan_with_vec.nearest("vec", &queries, k).unwrap(); + scan_with_vec.use_index(false); + scan_with_vec.project(&["vec"]).unwrap(); + let batch_with_vec = scan_with_vec.try_into_batch().await.unwrap(); + assert!( + batch_with_vec.schema().column_with_name("vec").is_some(), + "batch flat KNN must include vector column when vec is projected" + ); + } + #[tokio::test] async fn test_primitive_query_length_multiple_of_dim_is_rejected() { let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index 665b8b53865..f8b506f46bb 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -17,7 +17,6 @@ use arrow_array::{ cast::AsArray, }; use arrow_schema::{DataType, Field, Schema, SchemaRef}; -use arrow_select::concat::concat_batches; use datafusion::physical_plan::PlanProperties; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::physical_plan::{ @@ -348,18 +347,62 @@ impl KNNVectorDistanceExec { }) } - fn take_vector_row( - batch: &RecordBatch, - column: &str, - row_index: u32, + fn take_vector_row(vectors: &dyn Array, row_index: u32) -> ArrayRef { + vectors.slice(row_index as usize, 1) + } + + fn take_slim_batch_field( + results: &[BatchKnnCandidate], + field_name: &str, ) -> DataFusionResult { - let vectors = batch.column_by_name(column).ok_or_else(|| { - DataFusionError::Internal(format!( - "batch KNN expected vector column '{column}' in scan batch" - )) - })?; - let indices = UInt32Array::from(vec![row_index]); - arrow::compute::take(vectors, &indices, None) + use std::collections::HashMap; + + type SlimBatchGroup = (Arc, Vec<(usize, u32)>); + let mut groups: HashMap<*const RecordBatch, SlimBatchGroup> = HashMap::new(); + for (result_index, candidate) in results.iter().enumerate() { + let BatchKnnExtra::WithVector { + slim_batch, + row_index, + .. + } = &candidate.extra + else { + return Err(DataFusionError::Internal( + "batch KNN expected slim batch in candidate heap".to_string(), + )); + }; + groups + .entry(Arc::as_ptr(slim_batch)) + .or_insert_with(|| (Arc::clone(slim_batch), Vec::new())) + .1 + .push((result_index, *row_index)); + } + + let mut ordered: Vec> = vec![None; results.len()]; + for (_, (slim_batch, entries)) in groups { + let indices = + UInt32Array::from_iter(entries.iter().map(|(_, row_index)| Some(*row_index))); + let taken = arrow_select::take::take_record_batch(slim_batch.as_ref(), &indices) + .map_err(|e| { + DataFusionError::ArrowError(Box::new(e), Some("take top-k rows".to_string())) + })?; + let column = taken.column_by_name(field_name).ok_or_else(|| { + DataFusionError::Internal(format!("column '{field_name}' missing from slim batch")) + })?; + for (offset, (result_index, _)) in entries.iter().enumerate() { + ordered[*result_index] = Some(column.slice(offset, 1)); + } + } + + let row_arrays: Vec<&dyn Array> = ordered + .iter() + .map(|array| { + array + .as_ref() + .expect("every result mapped from slim batch") + .as_ref() + }) + .collect(); + arrow::compute::concat(&row_arrays) .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) } @@ -371,15 +414,14 @@ impl KNNVectorDistanceExec { let mut columns: Vec = Vec::with_capacity(stored_schema.fields().len()); for field in stored_schema.fields() { if field.name() == ROW_ID { - let row_ids = UInt64Array::from_iter( - results.iter().map(|candidate| Some(candidate.row_id())), - ); + let row_ids = + UInt64Array::from_iter(results.iter().map(|candidate| Some(candidate.row_id))); columns.push(Arc::new(row_ids)); } else if field.name() == column { let vector_rows: Vec<&dyn Array> = results .iter() .map(|candidate| { - let BatchKnnCandidate::WithVector { vector_row, .. } = candidate else { + let BatchKnnExtra::WithVector { vector_row, .. } = &candidate.extra else { return Err(DataFusionError::Internal( "batch KNN expected vector rows in candidate heap".to_string(), )); @@ -391,41 +433,7 @@ impl KNNVectorDistanceExec { .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; columns.push(vector_column); } else { - let mut row_batches = Vec::with_capacity(results.len()); - for candidate in results { - let BatchKnnCandidate::WithVector { - batch, row_index, .. - } = candidate - else { - return Err(DataFusionError::Internal( - "batch KNN expected slim batch in candidate heap".to_string(), - )); - }; - let indices = UInt32Array::from(vec![*row_index]); - row_batches.push( - arrow_select::take::take_record_batch(batch.as_ref(), &indices).map_err( - |e| { - DataFusionError::ArrowError( - Box::new(e), - Some("take top-k row".to_string()), - ) - }, - )?, - ); - } - let taken = concat_batches(&row_batches[0].schema(), &row_batches) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; - columns.push( - taken - .column_by_name(field.name()) - .ok_or_else(|| { - DataFusionError::Internal(format!( - "column '{}' missing from slim batch", - field.name() - )) - })? - .clone(), - ); + columns.push(Self::take_slim_batch_field(results, field.name())?); } } RecordBatch::try_new(Arc::new(stored_schema.clone()), columns) @@ -470,12 +478,18 @@ impl KNNVectorDistanceExec { .as_primitive::() .clone(); - let slim_batch = if retain_vector { - Some(Arc::new( + let mut slim_batch: Option> = None; + let vectors = if retain_vector { + Some( batch - .drop_column(&column) - .map_err(|e| DataFusionError::External(Box::new(e)))?, - )) + .column_by_name(&column) + .ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected vector column '{column}' in scan batch" + )) + })? + .clone(), + ) } else { None }; @@ -504,30 +518,39 @@ impl KNNVectorDistanceExec { } let query_index = query_index as i32; let row_id = row_ids.value(row_index); - let candidate = if retain_vector { + if !would_enter_heap(heap, k, distance, row_id, query_index) { + continue; + } + + let extra = if retain_vector { let row_index = row_index as u32; - let vector_row = Self::take_vector_row(&batch, &column, row_index)?; - BatchKnnCandidate::WithVector { - query_index, - distance, - row_id, - batch: Arc::clone(slim_batch.as_ref().expect("slim batch")), + if slim_batch.is_none() { + slim_batch = Some(Arc::new( + batch + .drop_column(&column) + .map_err(|e| DataFusionError::External(Box::new(e)))?, + )); + } + let slim_batch = slim_batch.as_ref().expect("slim batch"); + let vector_row = + Self::take_vector_row(vectors.as_ref().expect("vectors"), row_index); + BatchKnnExtra::WithVector { + slim_batch: Arc::clone(slim_batch), row_index, vector_row, } } else { - BatchKnnCandidate::RowIdOnly { - query_index, - distance, - row_id, - } + BatchKnnExtra::RowIdOnly + }; + let candidate = BatchKnnCandidate { + query_index, + distance, + row_id, + extra, }; if heap.len() < k { heap.push(candidate); - } else if heap - .peek() - .is_some_and(|worst| candidate.cmp(worst).is_lt()) - { + } else { heap.pop(); heap.push(candidate); } @@ -540,10 +563,10 @@ impl KNNVectorDistanceExec { .flat_map(BinaryHeap::into_vec) .collect::>(); results.sort_by(|left, right| { - left.query_index() - .cmp(&right.query_index()) - .then_with(|| left.distance().total_cmp(&right.distance())) - .then_with(|| left.row_id().cmp(&right.row_id())) + left.query_index + .cmp(&right.query_index) + .then_with(|| left.distance.total_cmp(&right.distance)) + .then_with(|| left.row_id.cmp(&right.row_id)) }); if results.is_empty() { @@ -553,14 +576,14 @@ impl KNNVectorDistanceExec { let mut query_indices = Int32Builder::with_capacity(results.len()); let mut distances = Float32Builder::with_capacity(results.len()); for result in &results { - query_indices.append_value(result.query_index()); - distances.append_value(result.distance()); + query_indices.append_value(result.query_index); + distances.append_value(result.distance); } let output = if retain_vector { Self::assemble_batch_output(&results, stored_schema.as_ref(), &column)? } else { - let row_ids = UInt64Array::from_iter(results.iter().map(|c| Some(c.row_id()))); + let row_ids = UInt64Array::from_iter(results.iter().map(|c| Some(c.row_id))); RecordBatch::try_new(stored_schema.clone(), vec![Arc::new(row_ids) as ArrayRef]) .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))? }; @@ -764,49 +787,48 @@ impl ExecutionPlan for KNNVectorDistanceExec { } #[derive(Clone)] -enum BatchKnnCandidate { - RowIdOnly { - query_index: i32, - distance: f32, - row_id: u64, - }, +struct BatchKnnCandidate { + query_index: i32, + distance: f32, + row_id: u64, + extra: BatchKnnExtra, +} + +#[derive(Clone)] +enum BatchKnnExtra { + RowIdOnly, WithVector { - query_index: i32, - distance: f32, - row_id: u64, - batch: Arc, + slim_batch: Arc, row_index: u32, vector_row: ArrayRef, }, } -impl BatchKnnCandidate { - fn query_index(&self) -> i32 { - match self { - Self::RowIdOnly { query_index, .. } | Self::WithVector { query_index, .. } => { - *query_index - } - } - } - - fn distance(&self) -> f32 { - match self { - Self::RowIdOnly { distance, .. } | Self::WithVector { distance, .. } => *distance, - } - } - - fn row_id(&self) -> u64 { - match self { - Self::RowIdOnly { row_id, .. } | Self::WithVector { row_id, .. } => *row_id, - } - } +fn would_enter_heap( + heap: &BinaryHeap, + k: usize, + distance: f32, + row_id: u64, + query_index: i32, +) -> bool { + if heap.len() < k { + return true; + } + let worst = heap.peek().expect("heap non-empty when len >= k"); + let probe = BatchKnnCandidate { + query_index, + distance, + row_id, + extra: BatchKnnExtra::RowIdOnly, + }; + probe.cmp(worst).is_lt() } impl PartialEq for BatchKnnCandidate { fn eq(&self, other: &Self) -> bool { - self.query_index() == other.query_index() - && self.distance() == other.distance() - && self.row_id() == other.row_id() + self.query_index == other.query_index + && self.distance == other.distance + && self.row_id == other.row_id } } @@ -820,201 +842,10 @@ impl PartialOrd for BatchKnnCandidate { impl Ord for BatchKnnCandidate { fn cmp(&self, other: &Self) -> CmpOrdering { - self.distance() - .total_cmp(&other.distance()) - .then_with(|| self.row_id().cmp(&other.row_id())) - .then_with(|| self.query_index().cmp(&other.query_index())) - } -} - -#[cfg(test)] -fn batch_knn_candidate_heap_memory(candidates: &[BatchKnnCandidate]) -> usize { - use std::collections::HashSet; - - let mut seen_slim_batches = HashSet::new(); - candidates - .iter() - .map(|candidate| match candidate { - BatchKnnCandidate::RowIdOnly { .. } => 0, - BatchKnnCandidate::WithVector { - batch, vector_row, .. - } => { - let slim_mem = if seen_slim_batches.insert(Arc::as_ptr(batch)) { - batch - .columns() - .iter() - .map(|col| col.get_array_memory_size()) - .sum::() - } else { - 0 - }; - slim_mem + vector_row.get_array_memory_size() - } - }) - .sum() -} - -#[cfg(test)] -fn batch_knn_baseline_heap_memory(candidates: &[(RecordBatch, u32)]) -> usize { - candidates - .iter() - .map(|(batch, _)| { - batch - .columns() - .iter() - .map(|col| col.get_array_memory_size()) - .sum::() - }) - .sum() -} - -#[cfg(test)] -mod batch_knn_memory_tests { - use std::time::Instant; - - use arrow::array::{FixedSizeListArray, Float32Array, UInt64Array}; - use arrow_array::RecordBatch; - use arrow_schema::{DataType, Field, Schema}; - - use super::*; - - fn make_scan_batch(num_rows: usize, dim: usize) -> RecordBatch { - let vectors = - Float32Array::from_iter((0..num_rows * dim).map(|i| (i % 1000) as f32 / 1000.0)); - let vectors = FixedSizeListArray::try_new_from_values(vectors, dim as i32).unwrap(); - let row_ids = UInt64Array::from_iter((0..num_rows as u64).map(Some)); - let schema = Arc::new(Schema::new(vec![ - Field::new( - "vec", - DataType::FixedSizeList( - Field::new("item", DataType::Float32, true).into(), - dim as i32, - ), - true, - ), - ROW_ID_FIELD.clone(), - ])); - RecordBatch::try_new( - schema, - vec![Arc::new(vectors) as ArrayRef, Arc::new(row_ids) as ArrayRef], - ) - .unwrap() - } - - #[test] - fn test_batch_knn_heap_memory() { - let num_rows = 4096; - let dim = 512; - let m = 10; - let k = 10; - let batch = make_scan_batch(num_rows, dim); - let slim_batch = Arc::new(batch.drop_column("vec").unwrap()); - - let mut baseline = Vec::with_capacity(m * k); - let mut optimized_row_id_only = Vec::with_capacity(m * k); - let mut optimized_with_vector = Vec::with_capacity(m * k); - - for query_index in 0..m { - for row_index in 0..k { - let row_index_u32 = (query_index * k + row_index) as u32; - baseline.push((batch.clone(), row_index_u32)); - optimized_row_id_only.push(BatchKnnCandidate::RowIdOnly { - query_index: query_index as i32, - distance: row_index as f32, - row_id: row_index_u32 as u64, - }); - let vector_row = - KNNVectorDistanceExec::take_vector_row(&batch, "vec", row_index_u32).unwrap(); - optimized_with_vector.push(BatchKnnCandidate::WithVector { - query_index: query_index as i32, - distance: row_index as f32, - row_id: row_index_u32 as u64, - batch: Arc::clone(&slim_batch), - row_index: row_index_u32, - vector_row, - }); - } - } - - let baseline_mem = batch_knn_baseline_heap_memory(&baseline); - let row_id_only_mem = batch_knn_candidate_heap_memory(&optimized_row_id_only); - let with_vector_mem = batch_knn_candidate_heap_memory(&optimized_with_vector); - - assert!( - row_id_only_mem < baseline_mem / 50, - "RowIdOnly heap memory ({row_id_only_mem}) should be much smaller than baseline ({baseline_mem})" - ); - assert!( - with_vector_mem < baseline_mem / 10, - "WithVector heap memory ({with_vector_mem}) should be much smaller than baseline ({baseline_mem})" - ); - - let expected_vector_bytes = m * k * dim * std::mem::size_of::(); - assert!( - with_vector_mem >= expected_vector_bytes / 2 - && with_vector_mem <= expected_vector_bytes * 4, - "WithVector heap memory ({with_vector_mem}) should be on the order of m*k*d*4 ({expected_vector_bytes})" - ); - } - - #[test] - fn test_batch_knn_assemble_perf_smoke() { - let num_rows = 4096; - let dim = 128; - let m = 10; - let k = 10; - let batch = make_scan_batch(num_rows, dim); - let slim_batch = Arc::new(batch.drop_column("vec").unwrap()); - let stored_schema = batch.schema(); - - let mut candidates = Vec::with_capacity(m * k); - for query_index in 0..m { - for row_index in 0..k { - let row_index_u32 = (query_index * k + row_index) as u32; - let vector_row = - KNNVectorDistanceExec::take_vector_row(&batch, "vec", row_index_u32).unwrap(); - candidates.push(BatchKnnCandidate::WithVector { - query_index: query_index as i32, - distance: 0.0, - row_id: row_index_u32 as u64, - batch: Arc::clone(&slim_batch), - row_index: row_index_u32, - vector_row, - }); - } - } - - let baseline_batches: Vec<_> = (0..candidates.len()) - .map(|i| (batch.clone(), i as u32)) - .collect(); - - let start = Instant::now(); - for _ in 0..50 { - let mut row_batches = Vec::with_capacity(baseline_batches.len()); - for (full_batch, row_index) in &baseline_batches { - let indices = UInt32Array::from(vec![*row_index]); - row_batches - .push(arrow_select::take::take_record_batch(full_batch, &indices).unwrap()); - } - let _ = concat_batches(&row_batches[0].schema(), &row_batches).unwrap(); - } - let baseline_elapsed = start.elapsed(); - - let start = Instant::now(); - for _ in 0..50 { - let _ = KNNVectorDistanceExec::assemble_batch_output( - &candidates, - stored_schema.as_ref(), - "vec", - ) - .unwrap(); - } - let optimized_elapsed = start.elapsed(); - - assert!( - optimized_elapsed <= baseline_elapsed * 2, - "optimized assemble ({optimized_elapsed:?}) should not be much slower than baseline ({baseline_elapsed:?})" - ); + self.distance + .total_cmp(&other.distance) + .then_with(|| self.row_id.cmp(&other.row_id)) + .then_with(|| self.query_index.cmp(&other.query_index)) } } From e593bb7815b547b579de1dc5b53d70bc01a3bd8d Mon Sep 17 00:00:00 2001 From: zoey Date: Wed, 27 May 2026 21:43:44 +0800 Subject: [PATCH 3/8] fix(knn): drop vector while preserving non-vector batch columns For batch flat KNN, keep stored schema as input without the vector column instead of row-id-only, and carry slim batches when needed so non-vector columns can be assembled without retaining vectors. Add a filter-based batch KNN regression test to verify non-vector outputs remain correct. --- rust/lance/src/dataset/scanner.rs | 44 +++++++++++++++++++++++++++++++ rust/lance/src/io/exec/knn.rs | 37 +++++++++++++++----------- 2 files changed, 65 insertions(+), 16 deletions(-) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 9c344c95778..b34b01ea17a 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -6084,6 +6084,50 @@ mod test { ); } + #[tokio::test] + async fn test_batch_knn_flat_filter_keeps_non_vector_columns() { + let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) + .await + .unwrap(); + let dataset = &test_ds.dataset; + let k = 2; + let (queries, query_values) = batch_knn_two_queries(); + + let mut scan = dataset.scan(); + scan.nearest("vec", &queries, k).unwrap(); + scan.use_index(false); + scan.filter("i >= 0").unwrap(); + scan.project(&["i"]).unwrap(); + let batch = scan.try_into_batch().await.unwrap(); + + assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch); + assert!(batch.schema().column_with_name("i").is_some()); + + let query_indices = batch[QUERY_INDEX_COL].as_primitive::(); + for query_index in 0..2 { + let query = + Float32Array::from(query_values[query_index * 32..(query_index + 1) * 32].to_vec()); + let mut single = dataset.scan(); + single.nearest("vec", &query, k).unwrap(); + single.use_index(false); + single.filter("i >= 0").unwrap(); + single.project(&["i"]).unwrap(); + let single_batch = single.try_into_batch().await.unwrap(); + + let mask = BooleanArray::from_iter( + query_indices + .iter() + .map(|value| value.map(|value| value == query_index as i32)), + ); + let batch_slice = arrow::compute::filter_record_batch(&batch, &mask).unwrap(); + assert_eq!( + batch_slice["i"].as_primitive::().values(), + single_batch["i"].as_primitive::().values() + ); + } + } + #[tokio::test] async fn test_primitive_query_length_multiple_of_dim_is_rejected() { let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index f8b506f46bb..8a4721048ed 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -292,7 +292,7 @@ impl KNNVectorDistanceExec { } let stored_schema = if is_batch && !retain_vector { - Arc::new(Schema::new(vec![ROW_ID_FIELD.clone()])) + Arc::new(input_schema.without_column(column)) } else { Arc::new(input_schema) }; @@ -360,7 +360,7 @@ impl KNNVectorDistanceExec { type SlimBatchGroup = (Arc, Vec<(usize, u32)>); let mut groups: HashMap<*const RecordBatch, SlimBatchGroup> = HashMap::new(); for (result_index, candidate) in results.iter().enumerate() { - let BatchKnnExtra::WithVector { + let BatchKnnExtra::WithSlimBatch { slim_batch, row_index, .. @@ -421,7 +421,11 @@ impl KNNVectorDistanceExec { let vector_rows: Vec<&dyn Array> = results .iter() .map(|candidate| { - let BatchKnnExtra::WithVector { vector_row, .. } = &candidate.extra else { + let BatchKnnExtra::WithSlimBatch { + vector_row: Some(vector_row), + .. + } = &candidate.extra + else { return Err(DataFusionError::Internal( "batch KNN expected vector rows in candidate heap".to_string(), )); @@ -457,6 +461,7 @@ impl KNNVectorDistanceExec { retain_vector, } = config; let query_dim = query.len() / query_count; + let needs_slim_batch = stored_schema.fields().iter().any(|f| f.name() != ROW_ID); let mut heaps = (0..query_count) .map(|_| BinaryHeap::::with_capacity(k)) .collect::>(); @@ -522,7 +527,7 @@ impl KNNVectorDistanceExec { continue; } - let extra = if retain_vector { + let extra = if retain_vector || needs_slim_batch { let row_index = row_index as u32; if slim_batch.is_none() { slim_batch = Some(Arc::new( @@ -532,9 +537,15 @@ impl KNNVectorDistanceExec { )); } let slim_batch = slim_batch.as_ref().expect("slim batch"); - let vector_row = - Self::take_vector_row(vectors.as_ref().expect("vectors"), row_index); - BatchKnnExtra::WithVector { + let vector_row = if retain_vector { + Some(Self::take_vector_row( + vectors.as_ref().expect("vectors"), + row_index, + )) + } else { + None + }; + BatchKnnExtra::WithSlimBatch { slim_batch: Arc::clone(slim_batch), row_index, vector_row, @@ -580,13 +591,7 @@ impl KNNVectorDistanceExec { distances.append_value(result.distance); } - let output = if retain_vector { - Self::assemble_batch_output(&results, stored_schema.as_ref(), &column)? - } else { - let row_ids = UInt64Array::from_iter(results.iter().map(|c| Some(c.row_id))); - RecordBatch::try_new(stored_schema.clone(), vec![Arc::new(row_ids) as ArrayRef]) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))? - }; + let output = Self::assemble_batch_output(&results, stored_schema.as_ref(), &column)?; output .try_with_column_at(0, query_index_field(), Arc::new(query_indices.finish())) @@ -797,10 +802,10 @@ struct BatchKnnCandidate { #[derive(Clone)] enum BatchKnnExtra { RowIdOnly, - WithVector { + WithSlimBatch { slim_batch: Arc, row_index: u32, - vector_row: ArrayRef, + vector_row: Option, }, } From d69c68d260985eba6851daa4a30a54d2d4033f3b Mon Sep 17 00:00:00 2001 From: zoey Date: Thu, 28 May 2026 05:54:24 +0800 Subject: [PATCH 4/8] test(knn): cover nested vector batch flat KNN projection Add regression tests for payload.vec batch KNN omit/retain behavior and use qualified batch column access when retaining nested vectors. --- rust/lance/src/dataset/scanner.rs | 102 +++++++++++++++++++++++++++--- rust/lance/src/io/exec/knn.rs | 2 +- 2 files changed, 95 insertions(+), 9 deletions(-) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index b34b01ea17a..96aeb616771 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -5880,6 +5880,60 @@ mod test { (queries, query_values) } + async fn nested_vector_test_dataset(dim: u32) -> (TempStrDir, Dataset) { + let path = TempStrDir::default(); + let vec_field = ArrowField::new( + "vec", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + dim as i32, + ), + true, + ); + let payload_field = ArrowField::new( + "payload", + DataType::Struct(vec![vec_field.clone()].into()), + true, + ); + let schema = Arc::new(ArrowSchema::new(vec![ + ArrowField::new("i", DataType::Int32, true), + payload_field.clone(), + ])); + + let batches: Vec = (0..5) + .map(|batch_idx| { + let vector_values: Float32Array = (0..dim * 80).map(|v| v as f32).collect(); + let vectors = + FixedSizeListArray::try_new_from_values(vector_values, dim as i32).unwrap(); + let payload = StructArray::from(vec![( + Arc::new(vec_field.clone()), + Arc::new(vectors) as ArrayRef, + )]); + RecordBatch::try_new( + schema.clone(), + vec![ + Arc::new(Int32Array::from_iter_values( + batch_idx * 80..(batch_idx + 1) * 80, + )), + Arc::new(payload), + ], + ) + .unwrap() + }) + .collect(); + + let params = WriteParams { + max_rows_per_group: 10, + max_rows_per_file: 200, + data_storage_version: Some(LanceFileVersion::Stable), + enable_stable_row_ids: true, + ..Default::default() + }; + let reader = RecordBatchIterator::new(batches.into_iter().map(Ok), schema); + let dataset = Dataset::write(reader, &path, Some(params)).await.unwrap(); + (path, dataset) + } + fn assert_query_index_field(batch: &RecordBatch) { let schema = batch.schema(); let field = schema.field(0); @@ -5888,10 +5942,10 @@ mod test { assert!(!field.is_nullable()); } - fn assert_batch_knn_output_has_no_vector(batch: &RecordBatch) { + fn assert_batch_knn_output_has_no_vector(batch: &RecordBatch, vector_column: &str) { assert!( - batch.schema().column_with_name("vec").is_none(), - "batch flat KNN output must not include vector column when vec is not projected; columns: {:?}", + batch.schema().column_with_name(vector_column).is_none(), + "batch flat KNN output must not include vector column '{vector_column}' when it is not projected; columns: {:?}", batch.schema().field_names() ); } @@ -5970,7 +6024,7 @@ mod test { let batch = scan.try_into_batch().await.unwrap(); assert_query_index_field(&batch); - assert_batch_knn_output_has_no_vector(&batch); + assert_batch_knn_output_has_no_vector(&batch, "vec"); assert_eq!( batch.num_rows(), 2 * k, @@ -6037,7 +6091,7 @@ mod test { let batch = scan.try_into_batch().await.unwrap(); assert_query_index_field(&batch); - assert_batch_knn_output_has_no_vector(&batch); + assert_batch_knn_output_has_no_vector(&batch, "vec"); assert_eq!( batch[QUERY_INDEX_COL].as_primitive::().values(), &[0, 0] @@ -6058,7 +6112,7 @@ mod test { scan.use_index(false); scan.project(&["i"]).unwrap(); let batch = scan.try_into_batch().await.unwrap(); - assert_batch_knn_output_has_no_vector(&batch); + assert_batch_knn_output_has_no_vector(&batch, "vec"); assert_query_index_field(&batch); assert!(batch.schema().column_with_name("i").is_some()); assert!(batch.schema().column_with_name(DIST_COL).is_some()); @@ -6069,7 +6123,7 @@ mod test { scan_rowid_only.use_index(false); scan_rowid_only.project(&[ROW_ID]).unwrap(); let batch_rowid_only = scan_rowid_only.try_into_batch().await.unwrap(); - assert_batch_knn_output_has_no_vector(&batch_rowid_only); + assert_batch_knn_output_has_no_vector(&batch_rowid_only, "vec"); assert!(batch_rowid_only.schema().column_with_name(ROW_ID).is_some()); assert!(batch_rowid_only.schema().column_with_name("i").is_none()); @@ -6101,7 +6155,7 @@ mod test { let batch = scan.try_into_batch().await.unwrap(); assert_query_index_field(&batch); - assert_batch_knn_output_has_no_vector(&batch); + assert_batch_knn_output_has_no_vector(&batch, "vec"); assert!(batch.schema().column_with_name("i").is_some()); let query_indices = batch[QUERY_INDEX_COL].as_primitive::(); @@ -6128,6 +6182,38 @@ mod test { } } + #[tokio::test] + async fn test_batch_knn_flat_nested_vector_projection() { + const VECTOR_COLUMN: &str = "payload.vec"; + let (_tmp, dataset) = nested_vector_test_dataset(32).await; + let k = 2; + let (queries, _query_values) = batch_knn_two_queries(); + + let mut scan = dataset.scan(); + scan.nearest(VECTOR_COLUMN, &queries, k).unwrap(); + scan.use_index(false); + scan.project(&["i"]).unwrap(); + let batch = scan.try_into_batch().await.unwrap(); + assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch, VECTOR_COLUMN); + assert_eq!(batch.num_rows(), 2 * k); + assert!(batch.schema().column_with_name("i").is_some()); + + let mut scan_with_vec = dataset.scan(); + scan_with_vec.nearest(VECTOR_COLUMN, &queries, k).unwrap(); + scan_with_vec.use_index(false); + scan_with_vec.project(&[VECTOR_COLUMN]).unwrap(); + let batch_with_vec = scan_with_vec.try_into_batch().await.unwrap(); + assert!( + batch_with_vec + .schema() + .column_with_name(VECTOR_COLUMN) + .is_some(), + "batch flat KNN must include nested vector column when projected; columns: {:?}", + batch_with_vec.schema().field_names() + ); + } + #[tokio::test] async fn test_primitive_query_length_multiple_of_dim_is_rejected() { let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index 8a4721048ed..31c29b55e57 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -487,7 +487,7 @@ impl KNNVectorDistanceExec { let vectors = if retain_vector { Some( batch - .column_by_name(&column) + .column_by_qualified_name(&column) .ok_or_else(|| { DataFusionError::Internal(format!( "batch KNN expected vector column '{column}' in scan batch" From 806fe9871beeca8e3b8334d0e45427a478d24583 Mon Sep 17 00:00:00 2001 From: zoey Date: Thu, 28 May 2026 19:21:02 +0800 Subject: [PATCH 5/8] fix(knn): align batch stats and trim nested vector retention Rebuild batch KNN partition statistics from output schema, remove vector columns via nested path-aware helpers, and switch projected vector row capture from slice to take/copy to avoid retaining full buffers. --- rust/lance/src/io/exec/knn.rs | 259 ++++++++++++++++++++++++++++++---- 1 file changed, 228 insertions(+), 31 deletions(-) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index 31c29b55e57..c36cc2b37b6 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -217,6 +217,82 @@ impl DisplayAs for KNNVectorDistanceExec { } impl KNNVectorDistanceExec { + fn remove_field_path_from_fields( + fields: &[Arc], + path: &[String], + ) -> DataFusionResult>> { + if path.is_empty() { + return Ok(fields.to_vec()); + } + let mut removed = false; + let mut new_fields = Vec::with_capacity(fields.len()); + for field in fields { + if field.name() != &path[0] { + new_fields.push(field.clone()); + continue; + } + removed = true; + if path.len() == 1 { + continue; + } + match field.data_type() { + DataType::Struct(children) => { + let child_fields = children.iter().cloned().collect::>(); + let projected_children = + Self::remove_field_path_from_fields(&child_fields, &path[1..])?; + if projected_children.is_empty() { + continue; + } + let updated = Field::new( + field.name(), + DataType::Struct(projected_children.into()), + field.is_nullable(), + ) + .with_metadata(field.metadata().clone()); + new_fields.push(Arc::new(updated)); + } + _ => { + return Err(DataFusionError::Internal(format!( + "batch KNN cannot remove nested path '{}': '{}' is not a struct", + path.join("."), + field.name() + ))); + } + } + } + if !removed { + return Err(DataFusionError::Internal(format!( + "batch KNN expected vector column '{}' in scan batch schema", + path.join(".") + ))); + } + Ok(new_fields) + } + + fn remove_vector_from_schema(schema: &Schema, column: &str) -> DataFusionResult { + let path = lance_core::datatypes::parse_field_path(column).map_err(|err| { + DataFusionError::Internal(format!( + "batch KNN failed to parse vector column path '{column}': {err}" + )) + })?; + let fields = schema.fields().iter().cloned().collect::>(); + let updated_fields = Self::remove_field_path_from_fields(&fields, &path)?; + Ok(Schema::new_with_metadata( + updated_fields, + schema.metadata().clone(), + )) + } + + fn remove_vector_from_batch( + batch: &RecordBatch, + column: &str, + ) -> DataFusionResult { + let slim_schema = Self::remove_vector_from_schema(batch.schema().as_ref(), column)?; + batch + .project_by_schema(&slim_schema) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) + } + /// Create a new [`KNNVectorDistanceExec`] node. /// /// Returns an error if the preconditions are not met. @@ -292,7 +368,7 @@ impl KNNVectorDistanceExec { } let stored_schema = if is_batch && !retain_vector { - Arc::new(input_schema.without_column(column)) + Arc::new(Self::remove_vector_from_schema(&input_schema, column)?) } else { Arc::new(input_schema) }; @@ -347,8 +423,10 @@ impl KNNVectorDistanceExec { }) } - fn take_vector_row(vectors: &dyn Array, row_index: u32) -> ArrayRef { - vectors.slice(row_index as usize, 1) + fn take_vector_row(vectors: &dyn Array, row_index: u32) -> DataFusionResult { + let indices = UInt32Array::from_iter([Some(row_index)]); + arrow_select::take::take(vectors, &indices, None) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) } fn take_slim_batch_field( @@ -530,18 +608,22 @@ impl KNNVectorDistanceExec { let extra = if retain_vector || needs_slim_batch { let row_index = row_index as u32; if slim_batch.is_none() { - slim_batch = Some(Arc::new( - batch - .drop_column(&column) - .map_err(|e| DataFusionError::External(Box::new(e)))?, - )); + let should_drop_vector_from_slim = + !retain_vector || !column.contains('.'); + let slim = if should_drop_vector_from_slim { + Self::remove_vector_from_batch(&batch, &column)? + } else { + // keep nested projected vectors in slim batch until assembly supports nested reinsertion + batch.clone() + }; + slim_batch = Some(Arc::new(slim)); } let slim_batch = slim_batch.as_ref().expect("slim batch"); let vector_row = if retain_vector { Some(Self::take_vector_row( vectors.as_ref().expect("vectors"), row_index, - )) + )?) } else { None }; @@ -734,35 +816,41 @@ impl ExecutionPlan for KNNVectorDistanceExec { fn partition_statistics(&self, partition: Option) -> DataFusionResult { let inner_stats = self.input.partition_statistics(partition)?; - let schema = self.input.schema(); - let dist_stats = inner_stats + let input_schema = self.input.schema(); + let input_stats_by_name = inner_stats .column_statistics .iter() - .zip(schema.fields()) - .find(|(_, field)| field.name() == &self.column) - .map(|(stats, _)| ColumnStatistics { + .zip(input_schema.fields()) + .map(|(stats, field)| (field.name().as_str(), stats.clone())) + .collect::>(); + let vector_root = lance_core::datatypes::parse_field_path(&self.column) + .ok() + .and_then(|parts| parts.first().cloned()) + .unwrap_or_else(|| self.column.clone()); + let dist_stats = input_stats_by_name + .get(vector_root.as_str()) + .map(|stats| ColumnStatistics { null_count: stats.null_count, ..Default::default() }) .unwrap_or_default(); - let column_statistics = inner_stats - .column_statistics - .into_iter() - .zip(schema.fields()) - .filter(|(_, field)| field.name() != DIST_COL) - .map(|(stats, _)| stats) + let column_statistics = self + .output_schema + .fields() + .iter() + .map(|field| { + if field.name() == QUERY_INDEX_COL { + ColumnStatistics::default() + } else if field.name() == DIST_COL { + dist_stats.clone() + } else { + input_stats_by_name + .get(field.name().as_str()) + .cloned() + .unwrap_or_default() + } + }) .collect::>(); - let column_statistics = if self.is_batch { - std::iter::once(ColumnStatistics::default()) - .chain(column_statistics) - .chain(std::iter::once(dist_stats)) - .collect::>() - } else { - column_statistics - .into_iter() - .chain(std::iter::once(dist_stats)) - .collect::>() - }; Ok(Statistics { num_rows: inner_stats.num_rows, column_statistics, @@ -2802,6 +2890,115 @@ mod tests { ); } + #[test] + fn test_batch_partition_statistics_aligns_with_output_schema() { + let schema = Arc::new(ArrowSchema::new(vec![ + ArrowField::new("i", DataType::Int32, true), + ArrowField::new( + "vec", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + 4, + ), + true, + ), + ROW_ID_FIELD.clone(), + ])); + let batch = RecordBatch::new_empty(schema); + let input: Arc = Arc::new(TestingExec::new(vec![batch])); + let query = Arc::new(Float32Array::from(vec![0.0, 1.0, 2.0, 3.0])) as ArrayRef; + let plan = KNNVectorDistanceExec::try_new_batch( + input, + "vec", + query, + KnnBatchParams { + is_batch: true, + query_count: 1, + k: 2, + lower_bound: None, + upper_bound: None, + distance_type: DistanceType::L2, + retain_vector: false, + }, + ) + .unwrap(); + let stats = plan.partition_statistics(None).unwrap(); + assert_eq!( + stats.column_statistics.len(), + plan.schema().fields().len(), + "partition stats must align with output schema" + ); + let schema = plan.schema(); + let query_index_pos = schema + .column_with_name(QUERY_INDEX_COL) + .expect("query_index must exist") + .0; + let dist_pos = schema + .column_with_name(DIST_COL) + .expect("distance must exist") + .0; + assert_eq!( + stats.column_statistics[query_index_pos], + ColumnStatistics::default(), + ); + assert_eq!( + stats.column_statistics[dist_pos].null_count, + stats.column_statistics[schema.column_with_name("i").unwrap().0].null_count, + "distance null-count should be derived from vector/input nullability and remain aligned" + ); + } + + #[test] + fn test_remove_vector_from_schema_nested_path() { + let payload_field = ArrowField::new( + "payload", + DataType::Struct( + vec![ + ArrowField::new( + "vec", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + 4, + ), + true, + ), + ArrowField::new("tag", DataType::Utf8, true), + ] + .into(), + ), + true, + ); + let schema = ArrowSchema::new(vec![ + ArrowField::new("i", DataType::Int32, true), + payload_field, + ROW_ID_FIELD.clone(), + ]); + let without_vec = + KNNVectorDistanceExec::remove_vector_from_schema(&schema, "payload.vec").unwrap(); + let payload = without_vec.field_with_name("payload").unwrap(); + let DataType::Struct(children) = payload.data_type() else { + panic!("payload should remain struct"); + }; + assert!(children.iter().all(|f| f.name() != "vec")); + assert!(children.iter().any(|f| f.name() == "tag")); + } + + #[test] + fn test_take_vector_row_copies_single_row() { + let vectors = FixedSizeListArray::try_new_from_values( + Float32Array::from((0..12).map(|v| v as f32).collect::>()), + 4, + ) + .unwrap(); + let row = KNNVectorDistanceExec::take_vector_row(&vectors, 2).unwrap(); + assert_eq!(row.len(), 1); + assert_eq!( + row.to_data().offset(), + 0, + "take/copy should not retain row offset into the full input buffer" + ); + } + #[tokio::test] async fn test_multivector_score() { let query = Query { From de2e9f6abb04d458c163d9d4b52fa5848f5663ec Mon Sep 17 00:00:00 2001 From: zoey Date: Thu, 28 May 2026 19:32:08 +0800 Subject: [PATCH 6/8] test(knn): cover row id and row addr batch output Add a batch flat KNN regression test for projecting row_id and row_addr without vectors to ensure system columns are materialized correctly. --- rust/lance/src/dataset/scanner.rs | 29 +++++++++++++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 96aeb616771..ae16e3841e9 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -6214,6 +6214,35 @@ mod test { ); } + #[tokio::test] + async fn test_batch_knn_flat_projects_row_id_and_row_addr_without_vector() { + let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) + .await + .unwrap(); + let dataset = &test_ds.dataset; + let k = 2; + let (queries, _) = batch_knn_two_queries(); + + let mut scan = dataset.scan(); + scan.nearest("vec", &queries, k).unwrap(); + scan.use_index(false); + scan.project(&[ROW_ID]).unwrap(); + scan.with_row_address(); + + let batch = scan.try_into_batch().await.unwrap(); + assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch, "vec"); + assert_eq!(batch.num_rows(), 2 * k); + assert!(batch.schema().column_with_name(ROW_ID).is_some()); + assert!(batch.schema().column_with_name(ROW_ADDR).is_some()); + assert!(batch.schema().column_with_name(DIST_COL).is_some()); + assert_eq!( + batch[ROW_ADDR].as_primitive::().null_count(), + 0, + "row addresses should be materialized for all top-k rows" + ); + } + #[tokio::test] async fn test_primitive_query_length_multiple_of_dim_is_rejected() { let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) From 8c28cc3b65a3404bd08ee4a616359ae2abb36379 Mon Sep 17 00:00:00 2001 From: zoey Date: Thu, 28 May 2026 22:34:41 +0800 Subject: [PATCH 7/8] fix(knn): support escaped nested paths in batch retained vectors Use parse_field_path-based vector resolution for batch retained vectors and remove the nested projected fallback that cloned full batches, then add regressions for escaped nested vector projection and nested slim-batch vector removal. --- rust/lance/src/dataset/scanner.rs | 86 ++++++++++++ rust/lance/src/io/exec/knn.rs | 226 +++++++++++++++++++++++++----- 2 files changed, 274 insertions(+), 38 deletions(-) diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index ae16e3841e9..bfec901d0ca 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -5934,6 +5934,60 @@ mod test { (path, dataset) } + async fn escaped_nested_vector_test_dataset(dim: u32) -> (TempStrDir, Dataset) { + let path = TempStrDir::default(); + let vec_field = ArrowField::new( + "vec.with.dot", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + dim as i32, + ), + true, + ); + let payload_field = ArrowField::new( + "payload", + DataType::Struct(vec![vec_field.clone()].into()), + true, + ); + let schema = Arc::new(ArrowSchema::new(vec![ + ArrowField::new("i", DataType::Int32, true), + payload_field.clone(), + ])); + + let batches: Vec = (0..5) + .map(|batch_idx| { + let vector_values: Float32Array = (0..dim * 80).map(|v| v as f32).collect(); + let vectors = + FixedSizeListArray::try_new_from_values(vector_values, dim as i32).unwrap(); + let payload = StructArray::from(vec![( + Arc::new(vec_field.clone()), + Arc::new(vectors) as ArrayRef, + )]); + RecordBatch::try_new( + schema.clone(), + vec![ + Arc::new(Int32Array::from_iter_values( + batch_idx * 80..(batch_idx + 1) * 80, + )), + Arc::new(payload), + ], + ) + .unwrap() + }) + .collect(); + + let params = WriteParams { + max_rows_per_group: 10, + max_rows_per_file: 200, + data_storage_version: Some(LanceFileVersion::Stable), + enable_stable_row_ids: true, + ..Default::default() + }; + let reader = RecordBatchIterator::new(batches.into_iter().map(Ok), schema); + let dataset = Dataset::write(reader, &path, Some(params)).await.unwrap(); + (path, dataset) + } + fn assert_query_index_field(batch: &RecordBatch) { let schema = batch.schema(); let field = schema.field(0); @@ -6214,6 +6268,38 @@ mod test { ); } + #[tokio::test] + async fn test_batch_knn_flat_escaped_nested_vector_projection() { + const VECTOR_COLUMN: &str = "payload.`vec.with.dot`"; + let (_tmp, dataset) = escaped_nested_vector_test_dataset(32).await; + let k = 2; + let (queries, _) = batch_knn_two_queries(); + + let mut scan = dataset.scan(); + scan.nearest(VECTOR_COLUMN, &queries, k).unwrap(); + scan.use_index(false); + scan.project(&["i"]).unwrap(); + let batch = scan.try_into_batch().await.unwrap(); + assert_query_index_field(&batch); + assert_batch_knn_output_has_no_vector(&batch, VECTOR_COLUMN); + assert_eq!(batch.num_rows(), 2 * k); + assert!(batch.schema().column_with_name("i").is_some()); + + let mut scan_with_vec = dataset.scan(); + scan_with_vec.nearest(VECTOR_COLUMN, &queries, k).unwrap(); + scan_with_vec.use_index(false); + scan_with_vec.project(&[VECTOR_COLUMN]).unwrap(); + let batch_with_vec = scan_with_vec.try_into_batch().await.unwrap(); + assert!( + batch_with_vec + .schema() + .column_with_name(VECTOR_COLUMN) + .is_some(), + "batch flat KNN must include escaped nested vector column when projected; columns: {:?}", + batch_with_vec.schema().field_names() + ); + } + #[tokio::test] async fn test_batch_knn_flat_projects_row_id_and_row_addr_without_vector() { let test_ds = TestVectorDataset::new(LanceFileVersion::Stable, true) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index c36cc2b37b6..b2ac2ad4960 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -293,6 +293,44 @@ impl KNNVectorDistanceExec { .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) } + fn resolve_vector_column(batch: &RecordBatch, column: &str) -> DataFusionResult { + if let Some(col) = batch.column_by_name(column) { + return Ok(col.clone()); + } + let parts = lance_core::datatypes::parse_field_path(column).map_err(|e| { + DataFusionError::Internal(format!( + "batch KNN failed to parse vector column path '{column}': {e}" + )) + })?; + if parts.is_empty() { + return Err(DataFusionError::Internal(format!( + "batch KNN has invalid empty vector column path '{column}'" + ))); + } + let mut current = batch.column_by_name(&parts[0]).cloned().ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected vector column '{column}' in scan batch (missing root field '{}')", + parts[0] + )) + })?; + for part in &parts[1..] { + let struct_array = current + .as_any() + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected struct while resolving '{column}', but parent of '{part}' was not a struct" + )) + })?; + current = struct_array.column_by_name(part).cloned().ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected vector column '{column}' in scan batch (missing nested field '{part}')" + )) + })?; + } + Ok(current) + } + /// Create a new [`KNNVectorDistanceExec`] node. /// /// Returns an error if the preconditions are not met. @@ -484,36 +522,97 @@ impl KNNVectorDistanceExec { .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) } + fn build_struct_column_for_path( + field: &Field, + path: &[String], + leaf_column: ArrayRef, + ) -> DataFusionResult { + if path.is_empty() { + return Ok(leaf_column); + } + let DataType::Struct(children) = field.data_type() else { + return Err(DataFusionError::Internal(format!( + "batch KNN expected struct field '{}' while rebuilding nested vector path '{}'", + field.name(), + path.join(".") + ))); + }; + let mut columns = Vec::with_capacity(children.len()); + for child in children.iter() { + if child.name() == &path[0] { + if path.len() == 1 { + columns.push(leaf_column.clone()); + } else { + columns.push(Self::build_struct_column_for_path( + child, + &path[1..], + leaf_column.clone(), + )?); + } + } else { + columns.push(arrow_array::new_null_array( + child.data_type(), + leaf_column.len(), + )); + } + } + let struct_array = arrow_array::StructArray::try_new(children.clone(), columns, None) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + Ok(Arc::new(struct_array)) + } + + fn take_retained_vector_column( + results: &[BatchKnnCandidate], + field: &Field, + field_path: &[String], + ) -> DataFusionResult { + let vector_rows: Vec<&dyn Array> = results + .iter() + .map(|candidate| { + let BatchKnnExtra::WithSlimBatch { + vector_row: Some(vector_row), + .. + } = &candidate.extra + else { + return Err(DataFusionError::Internal( + "batch KNN expected vector rows in candidate heap".to_string(), + )); + }; + Ok(vector_row.as_ref()) + }) + .collect::>>()?; + let leaf_column = arrow::compute::concat(&vector_rows) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + if field_path.len() <= 1 { + Ok(leaf_column) + } else { + Self::build_struct_column_for_path(field, &field_path[1..], leaf_column) + } + } + fn assemble_batch_output( results: &[BatchKnnCandidate], stored_schema: &Schema, column: &str, + retain_vector: bool, ) -> DataFusionResult { + let field_path = lance_core::datatypes::parse_field_path(column).map_err(|e| { + DataFusionError::Internal(format!( + "batch KNN failed to parse vector column path '{column}': {e}" + )) + })?; let mut columns: Vec = Vec::with_capacity(stored_schema.fields().len()); for field in stored_schema.fields() { if field.name() == ROW_ID { let row_ids = UInt64Array::from_iter(results.iter().map(|candidate| Some(candidate.row_id))); columns.push(Arc::new(row_ids)); - } else if field.name() == column { - let vector_rows: Vec<&dyn Array> = results - .iter() - .map(|candidate| { - let BatchKnnExtra::WithSlimBatch { - vector_row: Some(vector_row), - .. - } = &candidate.extra - else { - return Err(DataFusionError::Internal( - "batch KNN expected vector rows in candidate heap".to_string(), - )); - }; - Ok(vector_row.as_ref()) - }) - .collect::>>()?; - let vector_column = arrow::compute::concat(&vector_rows) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; - columns.push(vector_column); + } else if retain_vector && !field_path.is_empty() && field.name() == &field_path[0] { + columns.push(Self::take_retained_vector_column( + results, + field, + &field_path, + )?); } else { columns.push(Self::take_slim_batch_field(results, field.name())?); } @@ -563,16 +662,7 @@ impl KNNVectorDistanceExec { let mut slim_batch: Option> = None; let vectors = if retain_vector { - Some( - batch - .column_by_qualified_name(&column) - .ok_or_else(|| { - DataFusionError::Internal(format!( - "batch KNN expected vector column '{column}' in scan batch" - )) - })? - .clone(), - ) + Some(Self::resolve_vector_column(&batch, &column)?) } else { None }; @@ -608,14 +698,7 @@ impl KNNVectorDistanceExec { let extra = if retain_vector || needs_slim_batch { let row_index = row_index as u32; if slim_batch.is_none() { - let should_drop_vector_from_slim = - !retain_vector || !column.contains('.'); - let slim = if should_drop_vector_from_slim { - Self::remove_vector_from_batch(&batch, &column)? - } else { - // keep nested projected vectors in slim batch until assembly supports nested reinsertion - batch.clone() - }; + let slim = Self::remove_vector_from_batch(&batch, &column)?; slim_batch = Some(Arc::new(slim)); } let slim_batch = slim_batch.as_ref().expect("slim batch"); @@ -673,7 +756,8 @@ impl KNNVectorDistanceExec { distances.append_value(result.distance); } - let output = Self::assemble_batch_output(&results, stored_schema.as_ref(), &column)?; + let output = + Self::assemble_batch_output(&results, stored_schema.as_ref(), &column, retain_vector)?; output .try_with_column_at(0, query_index_field(), Arc::new(query_indices.finish())) @@ -2137,6 +2221,7 @@ mod tests { use arrow::datatypes::Float32Type; use arrow_array::{ ArrayRef, FixedSizeListArray, Float32Array, Int32Array, RecordBatchIterator, StringArray, + StructArray, }; use arrow_schema::{Field as ArrowField, Schema as ArrowSchema}; use async_trait::async_trait; @@ -2999,6 +3084,71 @@ mod tests { ); } + #[test] + fn test_resolve_vector_column_supports_escaped_nested_path() { + let vec_field = ArrowField::new( + "vec.with.dot", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + 4, + ), + true, + ); + let payload_field = ArrowField::new( + "payload", + DataType::Struct(vec![vec_field.clone()].into()), + true, + ); + let schema = Arc::new(ArrowSchema::new(vec![payload_field])); + let vectors = FixedSizeListArray::try_new_from_values( + Float32Array::from((0..8).map(|v| v as f32).collect::>()), + 4, + ) + .unwrap(); + let payload = StructArray::from(vec![(Arc::new(vec_field), Arc::new(vectors) as ArrayRef)]); + let batch = RecordBatch::try_new(schema, vec![Arc::new(payload)]).unwrap(); + let vector = + KNNVectorDistanceExec::resolve_vector_column(&batch, "payload.`vec.with.dot`").unwrap(); + assert_eq!(vector.len(), 2); + } + + #[test] + fn test_remove_vector_from_batch_nested_keeps_siblings() { + let vec_field = ArrowField::new( + "vec.with.dot", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + 4, + ), + true, + ); + let tag_field = ArrowField::new("tag", DataType::Utf8, true); + let payload_field = ArrowField::new( + "payload", + DataType::Struct(vec![vec_field.clone(), tag_field.clone()].into()), + true, + ); + let schema = Arc::new(ArrowSchema::new(vec![payload_field])); + let vectors = FixedSizeListArray::try_new_from_values( + Float32Array::from((0..8).map(|v| v as f32).collect::>()), + 4, + ) + .unwrap(); + let tags = StringArray::from(vec!["a", "b"]); + let payload = StructArray::from(vec![ + (Arc::new(vec_field), Arc::new(vectors) as ArrayRef), + (Arc::new(tag_field), Arc::new(tags) as ArrayRef), + ]); + let batch = RecordBatch::try_new(schema, vec![Arc::new(payload)]).unwrap(); + + let slim = + KNNVectorDistanceExec::remove_vector_from_batch(&batch, "payload.`vec.with.dot`") + .unwrap(); + let payload = slim.column_by_name("payload").unwrap().as_struct(); + assert!(payload.column_by_name("vec.with.dot").is_none()); + assert!(payload.column_by_name("tag").is_some()); + } + #[tokio::test] async fn test_multivector_score() { let query = Query { From 90a5070ff309795e2069ab15e96dd47c0f2a6244 Mon Sep 17 00:00:00 2001 From: BubbleCal Date: Mon, 22 Jun 2026 14:21:48 +0800 Subject: [PATCH 8/8] fix(knn): preserve nested batch vector siblings --- rust/lance/src/io/exec/knn.rs | 136 ++++++++++++++++++++++++++++++++-- 1 file changed, 130 insertions(+), 6 deletions(-) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index b2ac2ad4960..a2e105b2978 100644 --- a/rust/lance/src/io/exec/knn.rs +++ b/rust/lance/src/io/exec/knn.rs @@ -471,6 +471,15 @@ impl KNNVectorDistanceExec { results: &[BatchKnnCandidate], field_name: &str, ) -> DataFusionResult { + Self::take_slim_batch_field_if_present(results, field_name)?.ok_or_else(|| { + DataFusionError::Internal(format!("column '{field_name}' missing from slim batch")) + }) + } + + fn take_slim_batch_field_if_present( + results: &[BatchKnnCandidate], + field_name: &str, + ) -> DataFusionResult> { use std::collections::HashMap; type SlimBatchGroup = (Arc, Vec<(usize, u32)>); @@ -501,13 +510,21 @@ impl KNNVectorDistanceExec { .map_err(|e| { DataFusionError::ArrowError(Box::new(e), Some("take top-k rows".to_string())) })?; - let column = taken.column_by_name(field_name).ok_or_else(|| { - DataFusionError::Internal(format!("column '{field_name}' missing from slim batch")) - })?; + let Some(column) = taken.column_by_name(field_name) else { + continue; + }; for (offset, (result_index, _)) in entries.iter().enumerate() { ordered[*result_index] = Some(column.slice(offset, 1)); } } + if ordered.iter().all(Option::is_none) { + return Ok(None); + } + if ordered.iter().any(Option::is_none) { + return Err(DataFusionError::Internal(format!( + "column '{field_name}' inconsistently present in slim batches" + ))); + } let row_arrays: Vec<&dyn Array> = ordered .iter() @@ -520,12 +537,14 @@ impl KNNVectorDistanceExec { .collect(); arrow::compute::concat(&row_arrays) .map_err(|e| DataFusionError::ArrowError(Box::new(e), None)) + .map(Some) } fn build_struct_column_for_path( field: &Field, path: &[String], leaf_column: ArrayRef, + slim_column: Option<&dyn Array>, ) -> DataFusionResult { if path.is_empty() { return Ok(leaf_column); @@ -537,18 +556,39 @@ impl KNNVectorDistanceExec { path.join(".") ))); }; + let slim_struct = slim_column + .map(|column| { + column + .as_any() + .downcast_ref::() + .ok_or_else(|| { + DataFusionError::Internal(format!( + "batch KNN expected slim column '{}' to be a struct while rebuilding nested vector path '{}'", + field.name(), + path.join(".") + )) + }) + }) + .transpose()?; let mut columns = Vec::with_capacity(children.len()); for child in children.iter() { if child.name() == &path[0] { if path.len() == 1 { columns.push(leaf_column.clone()); } else { + let child_slim_column = slim_struct + .and_then(|struct_array| struct_array.column_by_name(child.name())); columns.push(Self::build_struct_column_for_path( child, &path[1..], leaf_column.clone(), + child_slim_column.map(|column| column.as_ref()), )?); } + } else if let Some(column) = + slim_struct.and_then(|struct_array| struct_array.column_by_name(child.name())) + { + columns.push(column.clone()); } else { columns.push(arrow_array::new_null_array( child.data_type(), @@ -556,8 +596,12 @@ impl KNNVectorDistanceExec { )); } } - let struct_array = arrow_array::StructArray::try_new(children.clone(), columns, None) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + let struct_array = arrow_array::StructArray::try_new( + children.clone(), + columns, + slim_struct.and_then(|struct_array| struct_array.nulls().cloned()), + ) + .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; Ok(Arc::new(struct_array)) } @@ -586,7 +630,13 @@ impl KNNVectorDistanceExec { if field_path.len() <= 1 { Ok(leaf_column) } else { - Self::build_struct_column_for_path(field, &field_path[1..], leaf_column) + let slim_column = Self::take_slim_batch_field_if_present(results, field.name())?; + Self::build_struct_column_for_path( + field, + &field_path[1..], + leaf_column, + slim_column.as_deref(), + ) } } @@ -3149,6 +3199,80 @@ mod tests { assert!(payload.column_by_name("tag").is_some()); } + #[test] + fn test_assemble_batch_output_retained_nested_vector_keeps_sibling_values() { + let vec_field = ArrowField::new( + "vec", + DataType::FixedSizeList( + Arc::new(ArrowField::new("item", DataType::Float32, true)), + 4, + ), + true, + ); + let tag_field = ArrowField::new("tag", DataType::Utf8, true); + let payload_field = ArrowField::new( + "payload", + DataType::Struct(vec![vec_field.clone(), tag_field.clone()].into()), + true, + ); + let schema = Arc::new(ArrowSchema::new(vec![payload_field, ROW_ID_FIELD.clone()])); + let vectors = FixedSizeListArray::try_new_from_values( + Float32Array::from((0..12).map(|v| v as f32).collect::>()), + 4, + ) + .unwrap(); + let tags = StringArray::from(vec!["a", "b", "c"]); + let payload = StructArray::from(vec![ + (Arc::new(vec_field), Arc::new(vectors) as ArrayRef), + (Arc::new(tag_field), Arc::new(tags) as ArrayRef), + ]); + let batch = RecordBatch::try_new( + schema.clone(), + vec![ + Arc::new(payload) as ArrayRef, + Arc::new(UInt64Array::from(vec![10, 11, 12])) as ArrayRef, + ], + ) + .unwrap(); + let slim_batch = Arc::new( + KNNVectorDistanceExec::remove_vector_from_batch(&batch, "payload.vec").unwrap(), + ); + let vectors = KNNVectorDistanceExec::resolve_vector_column(&batch, "payload.vec").unwrap(); + let results = [2, 0] + .into_iter() + .map(|row_index| BatchKnnCandidate { + query_index: 0, + distance: row_index as f32, + row_id: 10 + row_index as u64, + extra: BatchKnnExtra::WithSlimBatch { + slim_batch: Arc::clone(&slim_batch), + row_index, + vector_row: Some( + KNNVectorDistanceExec::take_vector_row(vectors.as_ref(), row_index) + .unwrap(), + ), + }, + }) + .collect::>(); + + let output = KNNVectorDistanceExec::assemble_batch_output( + &results, + schema.as_ref(), + "payload.vec", + true, + ) + .unwrap(); + + let payload = output.column_by_name("payload").unwrap().as_struct(); + let tags = payload.column_by_name("tag").unwrap().as_string::(); + assert!(tags.is_valid(0)); + assert!(tags.is_valid(1)); + assert_eq!(tags.value(0), "c"); + assert_eq!(tags.value(1), "a"); + let vectors = payload.column_by_name("vec").unwrap(); + assert_eq!(vectors.len(), 2); + } + #[tokio::test] async fn test_multivector_score() { let query = Query {