mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-30 00:45:37 +00:00
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:
@@ -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())? {
|
||||
|
||||
@@ -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, ¶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::<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, ¶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"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user