feat: group a materialized view by the IVF partition of a vector column (#4223)

Grouping near neighbours together, so an all-pairs comparison runs per
bucket instead of over the whole table, needs the partition an IVF index
assigns each vector.

`ivf_partition(column)` returns that partition, assigned by the
centroids and distance type of the IVF index on the column. It is bound
when the view is declared and again at every refresh; an index retrain
commits a new source version, so the next refresh regroups. A column
without an IVF index, or with two, is refused.


Stacked on #4222.
This commit is contained in:
Wyatt Alt
2026-09-22 10:16:20 -07:00
committed by GitHub
parent 814de30c5a
commit e020e13744
3 changed files with 578 additions and 10 deletions
+4
View File
@@ -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())? {
+276 -10
View File
@@ -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<IvfTransformer>,
/// 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<String, IvfBinding>,
}
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<H: std::hash::Hasher>(&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<DataType> {
// 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<ColumnarValue> {
// 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<String, IvfBinding>) -> 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<Vec<String>> {
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<HashMap<String, IvfBinding>> {
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<IvfBinding> {
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<dyn TableSource> {
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)?;
@@ -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::<Float32Type, _, _>(
(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::<Float32Type, _, _>(
(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, &params)
.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::<Float32Type, _, _>(
[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::<UInt8Type, _, _>(
(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::<UInt8Type, _, _>(
[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, &params, 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"]
);
}
}