From 1226b7099dff68dcd4a7a0907a49a2562e4d2f00 Mon Sep 17 00:00:00 2001 From: zoey Date: Wed, 27 May 2026 00:44:55 +0800 Subject: [PATCH] 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:?})" + ); } }