fix: normalize cosine scores at ANN boundaries

This commit is contained in:
Gatefixer
2026-08-06 08:24:54 +00:00
parent f95d4f583d
commit 4f5c55888b
+219 -49
View File
@@ -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]