diff --git a/rust/lancedb/src/materialized_view.rs b/rust/lancedb/src/materialized_view.rs index 067bde13d..c4456dbe1 100644 --- a/rust/lancedb/src/materialized_view.rs +++ b/rust/lancedb/src/materialized_view.rs @@ -10,6 +10,7 @@ //! plain table. Queries, indexes and search work on the view unchanged. mod grouped; +pub use grouped::IVF_PARTITION; mod query; pub mod refresh; @@ -1597,6 +1598,9 @@ async fn prepare_with( lineage, .. } = plan(source_schema.clone(), &definition, staging.as_ref())?; + if definition.is_grouped() { + grouped::check(native.dataset.get().await?.as_ref(), &definition).await?; + } // What later projections (`input_column`) are planned against: for an // unnested view the flattened schema, where the alias is a column. let planning_schema = match physical_unnest(&definition, staging.as_ref())? { diff --git a/rust/lancedb/src/materialized_view/grouped.rs b/rust/lancedb/src/materialized_view/grouped.rs index fc0dca430..00e14527a 100644 --- a/rust/lancedb/src/materialized_view/grouped.rs +++ b/rust/lancedb/src/materialized_view/grouped.rs @@ -9,28 +9,40 @@ use std::collections::HashMap; use std::sync::Arc; use std::sync::atomic::{AtomicU64, Ordering}; +use arrow_array::cast::AsArray; +use arrow_array::{Array, FixedSizeListArray, UInt32Array}; +use arrow_buffer::NullBuffer; use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema, SchemaRef}; use datafusion::catalog::default_table_source::provider_as_source; use datafusion::datasource::empty::EmptyTable; +use datafusion::execution::FunctionRegistry; use datafusion::execution::SessionStateBuilder; use datafusion::execution::context::SessionState; use datafusion::physical_plan::SendableRecordBatchStream; use datafusion::physical_plan::stream::RecordBatchStreamAdapter; use datafusion::prelude::SessionContext; +use datafusion_common::DataFusionError; use datafusion_common::TableReference; use datafusion_common::config::ConfigOptions; use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion}; use datafusion_expr::planner::{ContextProvider, ExprPlanner}; use datafusion_expr::{ - AggregateUDF, HigherOrderUDF, LogicalPlan, ScalarUDF, TableSource, WindowUDF, + AggregateUDF, ColumnarValue, Expr, HigherOrderUDF, LogicalPlan, ScalarFunctionArgs, ScalarUDF, + ScalarUDFImpl, Signature, TableSource, Volatility, WindowUDF, }; use datafusion_sql::planner::SqlToRel; use datafusion_sql::sqlparser::dialect::GenericDialect; use datafusion_sql::sqlparser::parser::Parser; use futures::StreamExt; use lance::Dataset; +use lance::index::{DatasetIndexExt, DatasetIndexInternalExt}; use lance_core::ROW_ID; use lance_datafusion::exec::SessionContextExt; +use lance_index::metrics::NoOpMetricsCollector; +use lance_index::vector::ivf::{IvfTransformer, new_ivf_transformer}; +use lance_linalg::distance::DistanceType; +use lance_linalg::kernels::normalize_fsl; +use uuid::Uuid; use super::refresh::to_view_batch; use super::{ @@ -42,8 +54,262 @@ use crate::{Error, Result}; /// The name the source is planned under; never visible to the user. const SOURCE: &str = "__source"; -fn session() -> SessionState { - SessionStateBuilder::new().with_default_features().build() +/// `ivf_partition(column)`: the partition of the IVF index on `column` that +/// each vector falls in, assigned by that index's centroids and distance +/// type. Bound per refresh, so a retrained index regroups the view. +pub const IVF_PARTITION: &str = "ivf_partition"; + +/// The index `ivf_partition(column)` assigns by. +#[derive(Debug, Clone)] +struct IvfBinding { + index: Uuid, + transformer: Arc, + /// A cosine index assigns L2 over unit vectors, as lance's index path does. + normalize: bool, +} + +/// `bindings` maps a source column to its index; planning runs unbound. +#[derive(Debug)] +struct IvfPartition { + signature: Signature, + bindings: HashMap, +} + +impl PartialEq for IvfPartition { + fn eq(&self, other: &Self) -> bool { + self.bindings.len() == other.bindings.len() + && self.bindings.iter().all(|(column, binding)| { + other + .bindings + .get(column) + .is_some_and(|b| b.index == binding.index) + }) + } +} + +impl Eq for IvfPartition {} + +impl std::hash::Hash for IvfPartition { + fn hash(&self, state: &mut H) { + let mut bound: Vec<_> = self.bindings.iter().map(|(c, b)| (c, b.index)).collect(); + bound.sort(); + bound.hash(state); + } +} + +impl ScalarUDFImpl for IvfPartition { + fn name(&self) -> &str { + IVF_PARTITION + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, arg_types: &[DataType]) -> datafusion_common::Result { + // Float vectors under L2/dot/cosine, byte vectors under Hamming; the + // binding checks the index's metric against the element type. + match arg_types { + [DataType::FixedSizeList(element, _)] + if element.data_type().is_floating() || *element.data_type() == DataType::UInt8 => + { + Ok(DataType::UInt32) + } + other => datafusion_common::plan_err!( + "{IVF_PARTITION} takes one vector column, not {other:?}" + ), + } + } + + fn invoke_with_args( + &self, + args: ScalarFunctionArgs, + ) -> datafusion_common::Result { + // The argument is a bare column (checked at planning), so its field + // names the column the binding was made for. + let column = args.arg_fields[0].name(); + let Some(binding) = self.bindings.get(column) else { + return datafusion_common::exec_err!("{IVF_PARTITION}({column}) has no index bound"); + }; + let vectors = args.args[0].to_array(args.number_rows)?; + let vectors = vectors.as_fixed_size_list(); + let normalized; + let vectors = if binding.normalize { + normalized = + normalize_fsl(vectors).map_err(|e| DataFusionError::External(Box::new(e)))?; + &normalized + } else { + vectors + }; + let partitions = binding + .transformer + .compute_partitions(vectors) + .map_err(|e| DataFusionError::External(Box::new(e)))?; + // A null vector falls in no partition, and neither does one the + // assigner could not place (a zero vector under cosine): a NULL bucket + // is honest, membership in bucket 0 is not. + let partitions = UInt32Array::new( + partitions.values().clone(), + NullBuffer::union(vectors.nulls(), partitions.nulls()), + ); + Ok(ColumnarValue::Array(Arc::new(partitions))) + } +} + +fn session(bindings: HashMap) -> SessionState { + let mut state = SessionStateBuilder::new().with_default_features().build(); + let ivf = IvfPartition { + signature: Signature::any(1, Volatility::Immutable), + bindings, + }; + state + .register_udf(Arc::new(ScalarUDF::new_from_impl(ivf))) + .expect("registering a scalar function cannot fail"); + state +} + +/// The columns `plan` calls `ivf_partition` on. The argument must be a bare +/// column: its field is how the bound function finds the column's index. +fn ivf_columns(plan: &LogicalPlan) -> Result> { + let mut columns = Vec::new(); + let mut misuse = None; + plan.apply(|node| { + for expr in node.expressions() { + expr.apply(|e| { + if let Expr::ScalarFunction(call) = e + && call.name() == IVF_PARTITION + { + match call.args.as_slice() { + [Expr::Column(column)] => columns.push(column.name.clone()), + _ => misuse = Some(e.to_string()), + } + } + Ok(TreeNodeRecursion::Continue) + })?; + } + Ok(TreeNodeRecursion::Continue) + }) + .map_err(|e| Error::InvalidInput { + message: format!("invalid grouped view: {e}"), + })?; + if let Some(call) = misuse { + return Err(Error::InvalidInput { + message: format!("{IVF_PARTITION} takes a vector column of the source, not `{call}`"), + }); + } + columns.sort(); + columns.dedup(); + Ok(columns) +} + +/// Refuse a grouped `definition` whose functions cannot be bound over +/// `source`, such as `ivf_partition` on a column without an IVF index. +pub(super) async fn check(source: &Dataset, definition: &MaterializedViewDefinition) -> Result<()> { + bind(source, definition).await.map(|_| ()) +} + +/// Bind every `ivf_partition(column)` in `definition` to the IVF index on +/// that column of `source`. +async fn bind( + source: &Dataset, + definition: &MaterializedViewDefinition, +) -> Result> { + let schema = Arc::new(ArrowSchema::from(source.schema())); + let planned = logical_plan(&session(HashMap::new()), empty_source(&schema), definition)?; + let mut bindings = HashMap::new(); + for column in ivf_columns(&planned)? { + bindings.insert(column.clone(), bind_column(source, &column).await?); + } + Ok(bindings) +} + +async fn bind_column(source: &Dataset, column: &str) -> Result { + let field = source + .schema() + .field(column) + .ok_or_else(|| Error::InvalidInput { + message: format!("{IVF_PARTITION}: the source has no column '{column}'"), + })?; + // One index, whose segments all assign by one model. A segment trained + // on its own fragments carries its own centroids, and grouping every row + // by one segment's model would bucket the other segments' rows wrongly. + let mut found: Option<(String, Uuid, DistanceType, FixedSizeListArray)> = None; + for index in source.load_indices().await?.iter() { + if index.fields != [field.id] { + continue; + } + let Ok(vector) = source + .open_vector_index(column, &index.uuid, &NoOpMetricsCollector) + .await + else { + continue; + }; + let Some(centroids) = vector.ivf_model().centroids.clone() else { + continue; + }; + let metric = vector.metric_type(); + match &found { + None => found = Some((index.name.clone(), index.uuid, metric, centroids)), + Some((name, _, seen_metric, seen)) if *name == index.name => { + if *seen_metric != metric || seen.to_data() != centroids.to_data() { + return Err(Error::InvalidInput { + message: format!( + "{IVF_PARTITION}({column}): the segments of index '{name}' were \ + trained separately and assign by different models; rebuild \ + the index before grouping by it" + ), + }); + } + } + Some((name, ..)) => { + return Err(Error::InvalidInput { + message: format!( + "{IVF_PARTITION}({column}) is ambiguous: indices '{name}' and '{}' \ + both cover it", + index.name + ), + }); + } + } + } + let byte_vectors = matches!( + field.data_type(), + DataType::FixedSizeList(element, _) if *element.data_type() == DataType::UInt8 + ); + let Some((_, uuid, metric, centroids)) = found else { + return Err(Error::InvalidInput { + message: format!( + "{IVF_PARTITION}({column}) needs an IVF vector index on '{column}'; create one first" + ), + }); + }; + // The assigner pairs byte vectors with Hamming and float vectors with the + // rest; an index of the other kind cannot place this column's values. + if byte_vectors != (metric == DistanceType::Hamming) { + return Err(Error::InvalidInput { + message: format!( + "{IVF_PARTITION}({column}): the index assigns by {metric}, which does not \ + apply to {}", + field.data_type() + ), + }); + } + // Lance's index path assigns cosine by L2 over unit vectors. + let normalize = metric == DistanceType::Cosine; + let distance = if normalize { DistanceType::L2 } else { metric }; + Ok(IvfBinding { + index: uuid, + transformer: Arc::new(new_ivf_transformer(centroids, distance, vec![])), + normalize, + }) +} + +fn empty_source(source_schema: &SchemaRef) -> Arc { + let mut fields = source_schema.fields().to_vec(); + fields.push(Arc::new(ArrowField::new(ROW_ID, DataType::UInt64, false))); + provider_as_source(Arc::new(EmptyTable::new(Arc::new(ArrowSchema::new( + fields, + ))))) } /// The query DataFusion runs: the view's projections plus the group's @@ -180,12 +446,12 @@ pub(super) fn plan( ..definition.clone() }; - let mut fields = source_schema.fields().to_vec(); - fields.push(Arc::new(ArrowField::new(ROW_ID, DataType::UInt64, false))); - let empty = provider_as_source(Arc::new(EmptyTable::new(Arc::new(ArrowSchema::new( - fields, - ))))); - let planned = logical_plan(&session(), empty, &definition)?; + let planned = logical_plan( + &session(HashMap::new()), + empty_source(&source_schema), + &definition, + )?; + ivf_columns(&planned)?; let mut exprs = Vec::new(); planned @@ -241,7 +507,7 @@ pub(super) async fn stream( scanner.project(inputs)?; let scan: SendableRecordBatchStream = scanner.try_into_stream().await?.into(); - let state = session(); + let state = session(bind(source, definition).await?); let ctx = SessionContext::new_with_state(state.clone()); let source = ctx.read_one_shot(scan)?.into_view(); let planned = logical_plan(&state, provider_as_source(source), definition)?; diff --git a/rust/lancedb/src/materialized_view/refresh.rs b/rust/lancedb/src/materialized_view/refresh.rs index 39e1d8b9c..6f19c3c91 100644 --- a/rust/lancedb/src/materialized_view/refresh.rs +++ b/rust/lancedb/src/materialized_view/refresh.rs @@ -4489,4 +4489,302 @@ mod tests { .unwrap(); assert!(grouped.input_column("x").is_err()); } + + /// Twenty nonzero vectors in two clusters far apart, ids 0-9 and 10-19. + /// Nonzero so every vector has a cosine direction. + async fn clustered_source(conn: &Connection) -> Table { + use arrow_array::types::Float32Type; + use arrow_array::{FixedSizeListArray, Int32Array}; + + let vectors = FixedSizeListArray::from_iter_primitive::( + (0..20).map(|i| { + let base = if i < 10 { 0.0 } else { 100.0 }; + Some(vec![Some(base + 1.0 + i as f32 * 0.1), Some(base)]) + }), + 2, + ); + let batch = RecordBatch::try_from_iter(vec![ + ( + "id", + Arc::new(Int32Array::from_iter_values(0..20)) as arrow_array::ArrayRef, + ), + ("vec", Arc::new(vectors) as arrow_array::ArrayRef), + ]) + .unwrap(); + conn.create_table("images", batch) + .write_options(crate::materialized_view::tests::stable_row_ids()) + .execute() + .await + .unwrap() + } + + #[tokio::test] + async fn ivf_partition_groups_by_the_index_on_the_column() { + use crate::index::vector::IvfFlatIndexBuilder; + + const BUCKETS: &str = "SELECT ivf_partition(vec) AS bucket, count(*) AS n, \ + min(id) AS lo, max(id) AS hi FROM images GROUP BY ivf_partition(vec)"; + let conn = connect("memory://").execute().await.unwrap(); + let source = clustered_source(&conn).await; + let unindexed = declare(&source, "buckets", BUCKETS).await.unwrap_err(); + assert!( + unindexed.to_string().contains("IVF vector index"), + "{unindexed}" + ); + + source + .create_index( + &["vec"], + Index::IvfFlat(IvfFlatIndexBuilder::default().num_partitions(2)), + ) + .execute() + .await + .unwrap(); + let view = declare(&source, "buckets", BUCKETS).await.unwrap(); + view.refresh().execute().await.unwrap(); + assert_eq!( + rows(view.table(), &["n", "lo", "hi"]).await, + ["10 0 9", "10 10 19"] + ); + let buckets = rows(view.table(), &["bucket"]).await; + assert_eq!(buckets, ["0", "1"]); + } + + #[tokio::test] + async fn ivf_partition_takes_a_column() { + let conn = connect("memory://").execute().await.unwrap(); + let source = clustered_source(&conn).await; + for sql in [ + "SELECT ivf_partition(id) AS b, count(*) AS n FROM images GROUP BY ivf_partition(id)", + "SELECT ivf_partition(vec, vec) AS b, count(*) AS n FROM images \ + GROUP BY ivf_partition(vec, vec)", + ] { + assert!(declare(&source, "v", sql).await.is_err(), "{sql}"); + } + } + + /// A cosine index assigns by L2 over unit vectors, as lance's own index + /// path does; forwarding cosine to the assigner panics. + #[tokio::test] + async fn ivf_partition_supports_cosine_indices() { + use crate::index::vector::IvfFlatIndexBuilder; + + const BUCKETS: &str = "SELECT ivf_partition(vec) AS bucket, count(*) AS n, \ + min(id) AS lo, max(id) AS hi FROM images GROUP BY ivf_partition(vec)"; + let conn = connect("memory://").execute().await.unwrap(); + let source = clustered_source(&conn).await; + source + .create_index( + &["vec"], + Index::IvfFlat( + IvfFlatIndexBuilder::default() + .num_partitions(2) + .distance_type(crate::DistanceType::Cosine), + ), + ) + .execute() + .await + .unwrap(); + let view = declare(&source, "cosine_buckets", BUCKETS).await.unwrap(); + view.refresh().execute().await.unwrap(); + assert_eq!( + rows(view.table(), &["n", "lo", "hi"]).await, + ["10 0 9", "10 10 19"] + ); + } + + /// Segments of one index trained separately carry different centroids; + /// grouping by one of them would bucket the other segments' rows wrongly. + #[tokio::test] + async fn ivf_partition_refuses_segments_with_different_models() { + use arrow_array::types::Float32Type; + use arrow_array::{ArrayRef, FixedSizeListArray, Int32Array}; + use lance::index::DatasetIndexExt; + use lance::index::vector::VectorIndexParams; + use lance_index::IndexType; + + let conn = connect("memory://").execute().await.unwrap(); + let source = clustered_source(&conn).await; + // A second fragment far from the first, so each segment trains its own centroids. + let far = FixedSizeListArray::from_iter_primitive::( + (0..20).map(|i| Some(vec![Some(-500.0 - i as f32), Some(-500.0)])), + 2, + ); + source + .add( + RecordBatch::try_from_iter(vec![ + ( + "id", + Arc::new(Int32Array::from_iter_values(20..40)) as ArrayRef, + ), + ("vec", Arc::new(far) as ArrayRef), + ]) + .unwrap(), + ) + .execute() + .await + .unwrap(); + let dataset = source.as_native().unwrap().dataset.get().await.unwrap(); + let mut dataset = dataset.as_ref().clone(); + let params = VectorIndexParams::ivf_flat(2, lance_linalg::distance::DistanceType::L2); + let mut segments = Vec::new(); + for fragment in dataset.get_fragments() { + let segment = dataset + .create_index_builder(&["vec"], IndexType::Vector, ¶ms) + .name("shared".to_string()) + .fragments(vec![fragment.id() as u32]) + .execute_uncommitted() + .await + .unwrap(); + segments.push(segment); + } + dataset + .commit_existing_index_segments("shared", "vec", segments) + .await + .unwrap(); + assert_eq!(dataset.load_indices().await.unwrap().len(), 2); + + let source = conn.open_table("images").execute().await.unwrap(); + let error = declare( + &source, + "buckets", + "SELECT ivf_partition(vec) AS b, count(*) AS n FROM images GROUP BY ivf_partition(vec)", + ) + .await + .unwrap_err(); + assert!(error.to_string().contains("trained separately"), "{error}"); + } + + /// A vector the assigner cannot place has no bucket: it is grouped under + /// NULL, never under partition 0. + #[tokio::test] + async fn ivf_partition_leaves_an_unassignable_vector_unbucketed() { + use crate::index::vector::IvfFlatIndexBuilder; + use arrow_array::types::Float32Type; + use arrow_array::{ArrayRef, FixedSizeListArray, Int32Array}; + + let conn = connect("memory://").execute().await.unwrap(); + let source = clustered_source(&conn).await; + source + .create_index( + &["vec"], + Index::IvfFlat( + IvfFlatIndexBuilder::default() + .num_partitions(2) + .distance_type(crate::DistanceType::Cosine), + ), + ) + .execute() + .await + .unwrap(); + // The zero vector has no direction, so cosine cannot place it. + let zero = FixedSizeListArray::from_iter_primitive::( + [Some(vec![Some(0.0), Some(0.0)])], + 2, + ); + source + .add( + RecordBatch::try_from_iter(vec![ + ("id", Arc::new(Int32Array::from(vec![20])) as ArrayRef), + ("vec", Arc::new(zero) as ArrayRef), + ]) + .unwrap(), + ) + .execute() + .await + .unwrap(); + let view = declare( + &source, + "buckets", + "SELECT ivf_partition(vec) AS bucket, count(*) AS n, min(id) AS lo, max(id) AS hi \ + FROM images GROUP BY ivf_partition(vec)", + ) + .await + .unwrap(); + view.refresh().execute().await.unwrap(); + // ArrayFormatter renders NULL as the empty string; partition numbers + // are the index's own and not asserted. + let groups = rows(view.table(), &["bucket", "n", "lo", "hi"]).await; + let (unbucketed, bucketed): (Vec<_>, Vec<_>) = + groups.iter().partition(|row| row.starts_with(' ')); + assert_eq!(unbucketed, [" 1 20 20"], "{groups:?}"); + let mut clusters: Vec<&str> = bucketed + .iter() + .map(|row| row.split_once(' ').unwrap().1) + .collect(); + clusters.sort(); + assert_eq!(clusters, ["10 0 9", "10 10 19"], "{groups:?}"); + } + + /// Byte vectors group by a Hamming index, the phash case. The centroids + /// are given, as lance's own tests do: two-means over a handful of + /// hashes can converge with one partition empty. + #[tokio::test] + async fn ivf_partition_groups_byte_vectors_by_hamming() { + use arrow_array::types::UInt8Type; + use arrow_array::{ArrayRef, FixedSizeListArray, Int32Array}; + use lance::index::DatasetIndexExt; + use lance::index::vector::VectorIndexParams; + use lance_index::IndexType; + use lance_index::vector::ivf::IvfBuildParams; + + let hashes = FixedSizeListArray::from_iter_primitive::( + (0..20u8).map(|i| { + let base = if i < 10 { 0x00 } else { 0xff }; + Some(vec![ + Some(base ^ (i % 4)), + Some(base), + Some(base), + Some(base), + ]) + }), + 4, + ); + let batch = RecordBatch::try_from_iter(vec![ + ( + "id", + Arc::new(Int32Array::from_iter_values(0..20)) as ArrayRef, + ), + ("phash", Arc::new(hashes) as ArrayRef), + ]) + .unwrap(); + let conn = connect("memory://").execute().await.unwrap(); + let source = conn + .create_table("images", batch) + .write_options(crate::materialized_view::tests::stable_row_ids()) + .execute() + .await + .unwrap(); + let centroids = FixedSizeListArray::from_iter_primitive::( + [Some(vec![Some(0); 4]), Some(vec![Some(0xff); 4])], + 4, + ); + let params = VectorIndexParams::with_ivf_flat_params( + lance_linalg::distance::DistanceType::Hamming, + IvfBuildParams::try_with_centroids(2, Arc::new(centroids)).unwrap(), + ); + let mut dataset = source + .as_native() + .unwrap() + .dataset + .get() + .await + .unwrap() + .as_ref() + .clone(); + dataset + .create_index(&["phash"], IndexType::Vector, None, ¶ms, true) + .await + .unwrap(); + let source = conn.open_table("images").execute().await.unwrap(); + + const BUCKETS: &str = "SELECT ivf_partition(phash) AS bucket, count(*) AS n, \ + min(id) AS lo, max(id) AS hi FROM images GROUP BY ivf_partition(phash)"; + let view = declare(&source, "buckets", BUCKETS).await.unwrap(); + view.refresh().execute().await.unwrap(); + assert_eq!( + rows(view.table(), &["bucket", "n", "lo", "hi"]).await, + ["0 10 0 9", "1 10 10 19"] + ); + } }