diff --git a/rust/lance/src/dataset/scanner.rs b/rust/lance/src/dataset/scanner.rs index 09cd7023e74..d4b58e4783f 100644 --- a/rust/lance/src/dataset/scanner.rs +++ b/rust/lance/src/dataset/scanner.rs @@ -4481,6 +4481,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, @@ -4492,6 +4500,7 @@ impl Scanner { lower_bound: q.lower_bound, upper_bound: q.upper_bound, distance_type: metric_type, + retain_vector, }, )?); @@ -5942,6 +5951,114 @@ 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) + } + + 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); @@ -5950,6 +6067,14 @@ mod test { assert!(!field.is_nullable()); } + fn assert_batch_knn_output_has_no_vector(batch: &RecordBatch, vector_column: &str) { + assert!( + 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() + ); + } + async fn assert_batch_matches_single_queries( dataset: &Dataset, batch: &RecordBatch, @@ -6024,6 +6149,7 @@ mod test { 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, @@ -6046,6 +6172,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()), @@ -6071,12 +6216,190 @@ mod test { 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[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, "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()); + 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, "vec"); + 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_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, "vec"); + 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_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_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) + .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) diff --git a/rust/lance/src/io/exec/knn.rs b/rust/lance/src/io/exec/knn.rs index 0ceddf7c5ee..73e901aee04 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::{ @@ -157,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, @@ -172,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, @@ -184,6 +185,7 @@ struct BatchKnnConfig { lower_bound: Option, upper_bound: Option, distance_type: DistanceType, + retain_vector: bool, } impl DisplayAs for KNNVectorDistanceExec { @@ -216,6 +218,120 @@ 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)) + } + + 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. @@ -236,6 +352,7 @@ impl KNNVectorDistanceExec { lower_bound: None, upper_bound: None, distance_type, + retain_vector: false, }, ) } @@ -253,6 +370,7 @@ impl KNNVectorDistanceExec { lower_bound, upper_bound, distance_type, + retain_vector, } = params; if query_count == 0 { return Err(Error::invalid_input( @@ -287,13 +405,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(Self::remove_vector_from_schema(&input_schema, column)?) + } 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, @@ -330,19 +454,230 @@ 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(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( + 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)>); + let mut groups: HashMap<*const RecordBatch, SlimBatchGroup> = HashMap::new(); + for (result_index, candidate) in results.iter().enumerate() { + let BatchKnnExtra::WithSlimBatch { + 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 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() + .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)) + .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); + } + 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 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(), + leaf_column.len(), + )); + } + } + 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)) + } + + 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 { + 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(), + ) + } + } + + 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 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())?); + } + } + 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, @@ -351,8 +686,10 @@ impl KNNVectorDistanceExec { lower_bound, upper_bound, distance_type, + 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::>(); @@ -374,6 +711,13 @@ impl KNNVectorDistanceExec { .as_primitive::() .clone(); + let mut slim_batch: Option> = None; + let vectors = if retain_vector { + Some(Self::resolve_vector_column(&batch, &column)?) + } 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()) @@ -398,20 +742,42 @@ impl KNNVectorDistanceExec { } let query_index = query_index as i32; let row_id = row_ids.value(row_index); - let row_index = row_index as u32; + if !would_enter_heap(heap, k, distance, row_id, query_index) { + continue; + } + + let extra = if retain_vector || needs_slim_batch { + let row_index = row_index as u32; + if slim_batch.is_none() { + 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"); + 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, + } + } else { + BatchKnnExtra::RowIdOnly + }; let candidate = BatchKnnCandidate { query_index, distance, row_id, - batch: batch.clone(), - row_index, + 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); } @@ -436,20 +802,14 @@ 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 { + 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())) - })?, - ); } - let output = concat_batches(&input_schema, &row_batches) - .map_err(|e| DataFusionError::ArrowError(Box::new(e), None))?; + 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())) .and_then(|batch| { @@ -501,6 +861,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, }, )?)) } @@ -515,7 +876,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(), @@ -524,6 +885,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(); @@ -586,35 +948,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, @@ -646,8 +1014,37 @@ struct BatchKnnCandidate { query_index: i32, distance: f32, row_id: u64, - batch: RecordBatch, - row_index: u32, + extra: BatchKnnExtra, +} + +#[derive(Clone)] +enum BatchKnnExtra { + RowIdOnly, + WithSlimBatch { + slim_batch: Arc, + row_index: u32, + vector_row: Option, + }, +} + +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 { @@ -655,7 +1052,6 @@ impl PartialEq for BatchKnnCandidate { self.query_index == other.query_index && self.distance == other.distance && self.row_id == other.row_id - && self.row_index == other.row_index } } @@ -673,7 +1069,6 @@ impl Ord for BatchKnnCandidate { .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)) } } @@ -1898,6 +2293,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; @@ -2684,6 +3080,254 @@ 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" + ); + } + + #[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()); + } + + #[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 {