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:?})" + ); } }