mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-08 06:18:59 +00:00
fix: normalize cosine scores at ANN boundaries
This commit is contained in:
+219
-49
@@ -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<Option<f32>> {
|
||||
if contains_plan::<KNNVectorDistanceExec>(plan) {
|
||||
return Ok(None);
|
||||
}
|
||||
/// cosine indices calculate cosine directly, so they are not included.
|
||||
async fn normalized_l2_ann_indices(plan: &dyn ExecutionPlan) -> Result<HashSet<String>> {
|
||||
let mut ann_plans = Vec::new();
|
||||
find_ann_plans(plan, &mut ann_plans);
|
||||
|
||||
let Some(ann) = find_plan::<ANNIvfSubIndexExec>(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::<ANNIvfSubIndexExec>() {
|
||||
ann_plans.push(ann);
|
||||
}
|
||||
for child in plan.children() {
|
||||
find_ann_plans(child.as_ref(), ann_plans);
|
||||
}
|
||||
}
|
||||
|
||||
fn find_plan<T: 'static>(plan: &dyn ExecutionPlan) -> Option<&T> {
|
||||
if let Some(plan) = (plan as &dyn std::any::Any).downcast_ref::<T>() {
|
||||
return Some(plan);
|
||||
fn collect_ann_plans(
|
||||
plan: &Arc<dyn ExecutionPlan>,
|
||||
ann_plans: &mut VecDeque<Arc<dyn ExecutionPlan>>,
|
||||
) {
|
||||
if plan.downcast_ref::<ANNIvfSubIndexExec>().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<dyn ExecutionPlan>,
|
||||
internal_plan: Arc<dyn ExecutionPlan>,
|
||||
normalized_l2_indices: &HashSet<String>,
|
||||
) -> Result<Arc<dyn ExecutionPlan>> {
|
||||
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<dyn ExecutionPlan>,
|
||||
internal_ann_plans: &mut VecDeque<Arc<dyn ExecutionPlan>>,
|
||||
normalized_l2_indices: &HashSet<String>,
|
||||
) -> Result<Arc<dyn ExecutionPlan>> {
|
||||
if let Some(public_ann) = public_plan.downcast_ref::<ANNIvfSubIndexExec>() {
|
||||
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::<ANNIvfSubIndexExec>()
|
||||
.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::<T>(child.as_ref()))
|
||||
}
|
||||
|
||||
fn contains_plan<T: 'static>(plan: &dyn ExecutionPlan) -> bool {
|
||||
find_plan::<T>(plan).is_some()
|
||||
.cloned()
|
||||
.map(|child| replace_ann_branches(child, internal_ann_plans, normalized_l2_indices))
|
||||
.collect::<Result<Vec<_>>>()?;
|
||||
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::<Vec<_>>()
|
||||
.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::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
let batch = &aliased[0];
|
||||
let aliased_distance = batch["aliased_distance"]
|
||||
.as_primitive::<Float32Type>()
|
||||
.value(0);
|
||||
let public_distance = batch[DIST_COL].as_primitive::<Float32Type>().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::<f32>();
|
||||
for (value, query_value) in orthogonal.iter_mut().zip(&query_vector) {
|
||||
*value -= projection * query_value;
|
||||
}
|
||||
let norm = orthogonal
|
||||
.iter()
|
||||
.map(|value| value * value)
|
||||
.sum::<f32>()
|
||||
.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::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
let mixed_distances = distances(&mixed);
|
||||
assert_eq!(mixed_distances.len(), 1);
|
||||
assert!((mixed_distances[0] - nearest).abs() < 1e-5);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
Reference in New Issue
Block a user