diff --git a/rust/lancedb/src/table/merge.rs b/rust/lancedb/src/table/merge.rs index b9dd5732b..564a15e17 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -1369,4 +1369,195 @@ mod lsm_tests { "LSM vector search must rank the memtable row first" ); } + + #[tokio::test] + async fn lsm_cosine_distance_scale_and_mixed_tier_ordering() { + use arrow::array::{FixedSizeListBuilder, Float32Builder}; + use arrow::datatypes::Float32Type; + + use crate::index::Index; + use crate::index::vector::IvfPqIndexBuilder; + + const DIM: usize = 8; + const N: usize = 256; + + fn normalized_vector(state: &mut u64) -> Vec { + let mut vector = (0..DIM) + .map(|_| { + *state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1); + ((*state >> 32) as u32 as f32 / u32::MAX as f32) * 2.0 - 1.0 + }) + .collect::>(); + let norm = vector.iter().map(|value| value * value).sum::().sqrt(); + vector.iter_mut().for_each(|value| *value /= norm); + vector + } + + let schema = Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int64, false), + Field::new( + "vec", + DataType::FixedSizeList( + Arc::new(Field::new("item", DataType::Float32, true)), + DIM as i32, + ), + false, + ), + ])); + let make_batch = |rows: Vec<(i64, Vec)>| { + let ids = rows.iter().map(|(id, _)| *id).collect::>(); + let mut vectors = FixedSizeListBuilder::new(Float32Builder::new(), DIM as i32); + for (_, vector) in &rows { + vectors.values().append_slice(vector); + vectors.append(true); + } + RecordBatch::try_new( + schema.clone(), + vec![Arc::new(Int64Array::from(ids)), Arc::new(vectors.finish())], + ) + .unwrap() + }; + let first_result = |batches: &[RecordBatch]| { + let batch = &batches[0]; + let id = batch["id"].as_primitive::().value(0); + let distance = batch["_distance"].as_primitive::().value(0); + (id, distance) + }; + + let mut state = 42; + let base_rows = (0..N) + .map(|id| (id as i64, normalized_vector(&mut state))) + .collect::>(); + let query = normalized_vector(&mut state); + + let dir = tempdir().unwrap(); + let conn = connect(dir.path().to_str().unwrap()) + .execute() + .await + .unwrap(); + let base = make_batch(base_rows); + let reader: Box = + Box::new(RecordBatchIterator::new(vec![Ok(base)], schema.clone())); + let table = conn + .create_table("cosine_lsm", reader) + .execute() + .await + .unwrap(); + table.set_unenforced_primary_key(["id"]).await.unwrap(); + table + .create_index( + &["vec"], + Index::IvfPq( + IvfPqIndexBuilder::default() + .distance_type(crate::DistanceType::Cosine) + .num_partitions(1) + .num_sub_vectors(1), + ), + ) + .name("vec_cosine".to_string()) + .execute() + .await + .unwrap(); + table + .set_lsm_write_spec( + LsmWriteSpec::unsharded().with_maintained_indexes(vec!["vec_cosine".to_string()]), + ) + .await + .unwrap(); + + let base_only = table + .query() + .nearest_to(query.as_slice()) + .unwrap() + .limit(1) + .use_lsm(false) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let (base_id, public_distance) = first_result(&base_only); + + let lsm = table + .query() + .nearest_to(query.as_slice()) + .unwrap() + .limit(1) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let (lsm_id, lsm_distance) = first_result(&lsm); + assert_eq!(lsm_id, base_id); + assert!( + (lsm_distance - public_distance).abs() < 1e-5, + "LSM cosine distance {lsm_distance} did not use the public scale {public_distance}" + ); + + // Add an exact memtable result whose distance lies between the public ANN + // score and its doubled internal score. Correctly normalized plans still + // rank the ANN row first; mixed units would incorrectly rank this row first. + assert!(public_distance > 0.0 && public_distance < 4.0 / 3.0); + let memtable_distance = public_distance * 1.5; + let cosine_similarity = 1.0 - memtable_distance; + let mut orthogonal = normalized_vector(&mut state); + let projection = orthogonal + .iter() + .zip(&query) + .map(|(left, right)| left * right) + .sum::(); + for (value, query_value) in orthogonal.iter_mut().zip(&query) { + *value -= projection * query_value; + } + let norm = orthogonal + .iter() + .map(|value| value * value) + .sum::() + .sqrt(); + orthogonal.iter_mut().for_each(|value| *value /= norm); + let sine = (1.0 - cosine_similarity * cosine_similarity).sqrt(); + let memtable_vector = query + .iter() + .zip(&orthogonal) + .map(|(query_value, orthogonal_value)| { + cosine_similarity * query_value + sine * orthogonal_value + }) + .collect::>(); + + let mut merge = table.merge_insert(&[]); + merge + .when_matched_update_all(None) + .when_not_matched_insert_all(); + let memtable = make_batch(vec![(N as i64, memtable_vector)]); + merge + .execute(Box::new(RecordBatchIterator::new( + vec![Ok(memtable)], + schema, + ))) + .await + .unwrap(); + + let mixed = table + .query() + .nearest_to(query.as_slice()) + .unwrap() + .limit(1) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let (mixed_id, mixed_distance) = first_result(&mixed); + assert_eq!( + mixed_id, base_id, + "mixed LSM tiers must compare ANN and exact distances in public units" + ); + assert!((mixed_distance - public_distance).abs() < 1e-5); + } } diff --git a/rust/lancedb/src/table/query.rs b/rust/lancedb/src/table/query.rs index 5f3088cf0..9b675ed45 100644 --- a/rust/lancedb/src/table/query.rs +++ b/rust/lancedb/src/table/query.rs @@ -399,6 +399,21 @@ async fn normalized_l2_ann_indices(plan: &dyn ExecutionPlan) -> Result, +) -> Result> { + let normalized_l2_indices = normalized_l2_ann_indices(plan.as_ref()).await?; + if normalized_l2_indices.is_empty() { + return Ok(plan); + } + normalize_ann_branches(plan.clone(), plan, &normalized_l2_indices) +} + 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); diff --git a/rust/lancedb/src/table/query/lsm.rs b/rust/lancedb/src/table/query/lsm.rs index 255d649b1..4500bde42 100644 --- a/rust/lancedb/src/table/query/lsm.rs +++ b/rust/lancedb/src/table/query/lsm.rs @@ -128,6 +128,10 @@ pub(super) async fn create_lsm_plan( .await? }; + // Normalize cosine ANN arms before LSM merge and sort nodes compare their + // distances with exact SSTable and memtable arms. + let plan = super::normalize_cosine_ann_branches(plan).await?; + // Lance appends the primary-key columns internally for dedup and keeps them in // the output; drop the ones the user did not request so the projection matches. restore_projection(plan, &query, &pk_columns)