diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 62d65816b..368ed514b 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -1,7 +1,10 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors -use std::sync::Arc; +use std::{ + collections::{HashSet, VecDeque}, + sync::Arc, +}; mod lsm; @@ -20,17 +23,17 @@ use arrow_schema::{DataType, Schema}; use datafusion_common::ScalarValue; use datafusion_expr::Operator; use datafusion_physical_expr::expressions::{BinaryExpr, Column, Literal}; -use datafusion_physical_plan::ExecutionPlan; use datafusion_physical_plan::PhysicalExpr; use datafusion_physical_plan::projection::ProjectionExec; use datafusion_physical_plan::repartition::RepartitionExec; use datafusion_physical_plan::union::UnionExec; +use datafusion_physical_plan::{ExecutionPlan, with_new_children_if_necessary}; use futures::future::try_join_all; use lance::dataset::mem_wal::DatasetMemWalExt; use lance::dataset::scanner::DatasetRecordBatchStream; use lance::dataset::scanner::Scanner; use lance::index::DatasetIndexInternalExt; -use lance::io::exec::{ANNIvfSubIndexExec, KNNVectorDistanceExec}; +use lance::io::exec::ANNIvfSubIndexExec; use lance_datafusion::exec::{analyze_plan as lance_analyze_plan, execute_plan}; use lance_index::metrics::NoOpMetricsCollector; use lance_index::vector::{DIST_COL, quantizer::QuantizationType}; @@ -334,17 +337,20 @@ pub async fn create_plan( } let mut plan = scanner.create_plan().await?; - if let Some(scale) = approximate_cosine_distance_scale(plan.as_ref()).await? { - // Lance's ANN range filter sees the quantizer's internal normalized squared-L2 - // scores. Translate the public cosine bounds before rebuilding the plan. - if query.lower_bound.is_some() || query.upper_bound.is_some() { + let normalized_l2_indices = normalized_l2_ann_indices(plan.as_ref()).await?; + if !normalized_l2_indices.is_empty() { + // Rebuild only the affected ANN nodes with internal normalized squared-L2 + // bounds. Exact branches keep the public cosine bounds from `plan`. + let internal_plan = if query.lower_bound.is_some() || query.upper_bound.is_some() { scanner.distance_range( - query.lower_bound.map(|bound| bound / scale), - query.upper_bound.map(|bound| bound / scale), + query.lower_bound.map(|bound| bound / COSINE_ANN_SCALE), + query.upper_bound.map(|bound| bound / COSINE_ANN_SCALE), ); - plan = scanner.create_plan().await?; - } - plan = scale_distance_column(plan, scale)?; + scanner.create_plan().await? + } else { + plan.clone() + }; + plan = normalize_ann_branches(plan, internal_plan, &normalized_l2_indices)?; } Ok(plan) @@ -352,53 +358,138 @@ pub async fn create_plan( //Helper functions below -/// Return the conversion from an ANN index's internal score to its public cosine scale. +const COSINE_ANN_SCALE: f32 = 0.5; + +/// Find ANN index segments whose scores use normalized squared L2 for cosine search. /// /// Cosine PQ/SQ/RQ indices normalize their vectors and use squared L2 internally. This /// preserves ranking, but squared L2 over unit vectors is twice the cosine distance. Flat -/// cosine indices calculate cosine directly and exact/refined plans already recompute the -/// public distance, so those plans must not be adjusted. -async fn approximate_cosine_distance_scale(plan: &dyn ExecutionPlan) -> Result> { - if contains_plan::(plan) { - return Ok(None); - } +/// cosine indices calculate cosine directly, so they are not included. +async fn normalized_l2_ann_indices(plan: &dyn ExecutionPlan) -> Result> { + let mut ann_plans = Vec::new(); + find_ann_plans(plan, &mut ann_plans); - let Some(ann) = find_plan::(plan) else { - return Ok(None); - }; - if ann.query().metric_type != Some(LanceDistanceType::Cosine) { - return Ok(None); + let mut checked = HashSet::new(); + let mut normalized_l2 = HashSet::new(); + for ann in ann_plans { + if ann.query().metric_type != Some(LanceDistanceType::Cosine) { + continue; + } + for index in ann.indices() { + let uuid = index.uuid.to_string(); + if !checked.insert(uuid.clone()) { + continue; + } + let vector_index = ann + .dataset() + .open_vector_index(&ann.query().column, &index.uuid, &NoOpMetricsCollector) + .await?; + let (_, quantization_type) = vector_index.sub_index_type(); + if matches!( + quantization_type, + QuantizationType::Product | QuantizationType::Scalar | QuantizationType::Rabit + ) { + normalized_l2.insert(uuid); + } + } } + Ok(normalized_l2) +} - let Some(index) = ann.indices().first() else { - return Ok(None); - }; - let vector_index = ann - .dataset() - .open_vector_index(&ann.query().column, &index.uuid, &NoOpMetricsCollector) - .await?; - let (_, quantization_type) = vector_index.sub_index_type(); - if matches!( - quantization_type, - QuantizationType::Product | QuantizationType::Scalar | QuantizationType::Rabit - ) { - Ok(Some(0.5)) - } else { - Ok(None) +fn find_ann_plans<'a>(plan: &'a dyn ExecutionPlan, ann_plans: &mut Vec<&'a ANNIvfSubIndexExec>) { + if let Some(ann) = plan.downcast_ref::() { + ann_plans.push(ann); + } + for child in plan.children() { + find_ann_plans(child.as_ref(), ann_plans); } } -fn find_plan(plan: &dyn ExecutionPlan) -> Option<&T> { - if let Some(plan) = (plan as &dyn std::any::Any).downcast_ref::() { - return Some(plan); +fn collect_ann_plans( + plan: &Arc, + ann_plans: &mut VecDeque>, +) { + if plan.downcast_ref::().is_some() { + ann_plans.push_back(plan.clone()); + return; } - plan.children() + for child in plan.children() { + collect_ann_plans(child, ann_plans); + } +} + +/// Replace normalized-L2 ANN nodes with equivalent nodes that use internal bounds, then +/// convert their output to the public cosine scale before any generic plan node consumes it. +fn normalize_ann_branches( + public_plan: Arc, + internal_plan: Arc, + normalized_l2_indices: &HashSet, +) -> Result> { + let mut internal_ann_plans = VecDeque::new(); + collect_ann_plans(&internal_plan, &mut internal_ann_plans); + let normalized = + replace_ann_branches(public_plan, &mut internal_ann_plans, normalized_l2_indices)?; + if !internal_ann_plans.is_empty() { + return Err(Error::Runtime { + message: "internal and public vector plans contained different ANN branches" + .to_string(), + }); + } + Ok(normalized) +} + +fn replace_ann_branches( + public_plan: Arc, + internal_ann_plans: &mut VecDeque>, + normalized_l2_indices: &HashSet, +) -> Result> { + if let Some(public_ann) = public_plan.downcast_ref::() { + let internal_plan = internal_ann_plans + .pop_front() + .ok_or_else(|| Error::Runtime { + message: "internal vector plan was missing an ANN branch".to_string(), + })?; + let internal_ann = internal_plan + .downcast_ref::() + .expect("collected only ANN plans"); + let same_indices = public_ann + .indices() + .iter() + .map(|index| &index.uuid) + .eq(internal_ann.indices().iter().map(|index| &index.uuid)); + if public_ann.query().column != internal_ann.query().column + || public_ann.query().metric_type != internal_ann.query().metric_type + || !same_indices + { + return Err(Error::Runtime { + message: "internal and public vector plans had mismatched ANN branches".to_string(), + }); + } + + let normalized_count = public_ann + .indices() + .iter() + .filter(|index| normalized_l2_indices.contains(&index.uuid.to_string())) + .count(); + if normalized_count == 0 { + return Ok(public_plan); + } + if normalized_count != public_ann.indices().len() { + return Err(Error::Runtime { + message: "one ANN branch mixed public and normalized-L2 distance scales" + .to_string(), + }); + } + return scale_distance_column(internal_plan, COSINE_ANN_SCALE); + } + + let children = public_plan + .children() .into_iter() - .find_map(|child| find_plan::(child.as_ref())) -} - -fn contains_plan(plan: &dyn ExecutionPlan) -> bool { - find_plan::(plan).is_some() + .cloned() + .map(|child| replace_ann_branches(child, internal_ann_plans, normalized_l2_indices)) + .collect::>>()?; + Ok(with_new_children_if_necessary(public_plan, children)?) } fn scale_distance_column( @@ -1220,7 +1311,7 @@ mod tests { Field::new("vector", vectors.data_type().clone(), false), ])); let batch = RecordBatch::try_new( - schema, + schema.clone(), vec![Arc::new(Int32Array::from_iter_values(0..num_rows)), vectors], ) .unwrap(); @@ -1291,6 +1382,85 @@ mod tests { let ranged_distances = distances(&ranged); assert_eq!(ranged_distances.len(), 1); assert!((ranged_distances[0] - nearest).abs() < 1e-5); + + let refined_ranged = table + .vector_search(query_vector.as_slice()) + .unwrap() + .limit(1) + .refine_factor(1) + .distance_range(None, Some(nearest + 1e-5)) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + assert_eq!( + distances(&refined_ranged).len(), + 1, + "refinement must not apply public cosine bounds to internal ANN scores" + ); + + let aliased = table + .vector_search(query_vector.as_slice()) + .unwrap() + .limit(1) + .select(Select::dynamic(&[("aliased_distance", "_distance")])) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let batch = &aliased[0]; + let aliased_distance = batch["aliased_distance"] + .as_primitive::() + .value(0); + let public_distance = batch[DIST_COL].as_primitive::().value(0); + assert!( + (aliased_distance - public_distance).abs() < 1e-5, + "distance aliases and auto-projected distances must use the same public scale" + ); + + // Appended rows take an exact fallback branch. Its public range filter must stay + // independent of the translated ANN bounds before both branches are merged. + let mut orthogonal = normalized_vector(&mut state, dimension); + let projection = orthogonal + .iter() + .zip(&query_vector) + .map(|(left, right)| left * right) + .sum::(); + for (value, query_value) in orthogonal.iter_mut().zip(&query_vector) { + *value -= projection * query_value; + } + let norm = orthogonal + .iter() + .map(|value| value * value) + .sum::() + .sqrt(); + orthogonal.iter_mut().for_each(|value| *value /= norm); + let appended_vectors = Arc::new(fixed_size_list_array(orthogonal, dimension as i32)); + let appended = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![num_rows])), appended_vectors], + ) + .unwrap(); + table.add(appended).execute().await.unwrap(); + + let mixed = table + .vector_search(query_vector.as_slice()) + .unwrap() + .limit(5) + .distance_range(None, Some(nearest + 1e-5)) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let mixed_distances = distances(&mixed); + assert_eq!(mixed_distances.len(), 1); + assert!((mixed_distances[0] - nearest).abs() < 1e-5); } #[tokio::test]