diff --git a/rust/lancedb/src/connection.rs b/rust/lancedb/src/connection.rs index 89e59e12e..46f855d44 100644 --- a/rust/lancedb/src/connection.rs +++ b/rust/lancedb/src/connection.rs @@ -420,6 +420,11 @@ impl Connection { /// /// * `name` - The name of the table /// * `initial_data` - The initial data to write to the table + /// + /// Floating-point `List` columns named `vec`, or with `vector` or `embedding` + /// in their name, are inferred as vector columns when the first batch has a + /// uniform, non-zero list length. The inferred dimension is validated for all + /// subsequent batches and stored as a `FixedSizeList`. pub fn create_table( &self, name: impl Into, diff --git a/rust/lancedb/src/connection/create_table.rs b/rust/lancedb/src/connection/create_table.rs index 66f6dfa8d..e7f1a44dc 100644 --- a/rust/lancedb/src/connection/create_table.rs +++ b/rust/lancedb/src/connection/create_table.rs @@ -8,7 +8,7 @@ use lance_io::object_store::StorageOptionsProvider; use crate::{ Error, Result, Table, connection::{merge_storage_options, set_storage_options_provider}, - data::scannable::{Scannable, WithEmbeddingsScannable}, + data::scannable::{Scannable, WithEmbeddingsScannable, maybe_infer_vector_schema}, database::{CreateTableMode, CreateTableRequest, Database}, embeddings::{EmbeddingDefinition, EmbeddingFunction, EmbeddingRegistry}, table::WriteOptions, @@ -147,6 +147,8 @@ impl CreateTableBuilder { let embedding_registry = self.embedding_registry.clone(); let parent = self.parent.clone(); + self.request.data = maybe_infer_vector_schema(self.request.data).await?; + // If embeddings were configured via add_embedding(), wrap the data if !self.embeddings.is_empty() { let wrapped_data: Box = Box::new(WithEmbeddingsScannable::try_new( diff --git a/rust/lancedb/src/data/scannable.rs b/rust/lancedb/src/data/scannable.rs index 0f810941a..630e85352 100644 --- a/rust/lancedb/src/data/scannable.rs +++ b/rust/lancedb/src/data/scannable.rs @@ -18,13 +18,16 @@ use crate::embeddings::{ }; use crate::table::{ColumnDefinition, ColumnKind, TableDefinition}; use crate::{Error, Result}; -use arrow_array::{ArrayRef, RecordBatch, RecordBatchIterator, RecordBatchReader}; -use arrow_schema::{ArrowError, SchemaRef}; +use arrow_array::{ArrayRef, RecordBatch, RecordBatchIterator, RecordBatchReader, cast::AsArray}; +use arrow_cast::{CastOptions, cast_with_options}; +use arrow_schema::{ArrowError, DataType, Schema, SchemaRef}; use async_trait::async_trait; use futures::StreamExt; use futures::stream::once; use lance_datafusion::utils::StreamingWriteSource; +use super::inspect::infer_dimension; + pub trait Scannable: Send { /// Returns the schema of the data. fn schema(&self) -> SchemaRef; @@ -497,6 +500,146 @@ impl Scannable for PeekedScannable { } } +fn name_suggests_vector_column(name: &str) -> bool { + let name = name.to_ascii_lowercase(); + name == "vec" || name.contains("vector") || name.contains("embedding") +} + +/// Infer fixed dimensions for vector-like floating-point list columns. +/// +/// Lance vector search requires `FixedSizeList` columns, but Arrow data assembled +/// from runtime embedding models is often represented as `List`. For vector-like +/// column names, inspect the first batch and convert uniform, non-empty lists to a +/// fixed-size schema. Every subsequent batch is cast with strict length checking. +pub(crate) async fn maybe_infer_vector_schema( + data: Box, +) -> Result> { + let input_schema = data.schema(); + let candidates = input_schema + .fields() + .iter() + .enumerate() + .filter(|(_, field)| { + name_suggests_vector_column(field.name()) + && matches!( + field.data_type(), + DataType::List(item) | DataType::LargeList(item) + if item.data_type().is_floating() + ) + }) + .map(|(index, _)| index) + .collect::>(); + if candidates.is_empty() { + return Ok(data); + } + + let mut peeked = PeekedScannable::new(data); + let Some(first_batch) = peeked.peek().await else { + return Ok(Box::new(peeked)); + }; + + let mut fields = input_schema.fields().iter().cloned().collect::>(); + let mut changed = false; + for index in candidates { + let array = first_batch.column(index); + let dimension = match array.data_type() { + DataType::List(_) => { + infer_dimension::(array.as_list::())? + .map(i64::from) + } + DataType::LargeList(_) => { + infer_dimension::(array.as_list::())? + } + _ => unreachable!(), + }; + let Some(dimension) = dimension.filter(|dimension| *dimension > 0) else { + continue; + }; + let dimension = i32::try_from(dimension).map_err(|_| Error::InvalidInput { + message: format!( + "Vector column '{}' has a dimension larger than i32::MAX", + fields[index].name() + ), + })?; + let item = match fields[index].data_type() { + DataType::List(item) | DataType::LargeList(item) => item.clone(), + _ => unreachable!(), + }; + fields[index] = Arc::new( + fields[index] + .as_ref() + .clone() + .with_data_type(DataType::FixedSizeList(item, dimension)), + ); + changed = true; + } + + if !changed { + return Ok(Box::new(peeked)); + } + + let output_schema = Arc::new(Schema::new_with_metadata( + fields, + input_schema.metadata().clone(), + )); + Ok(Box::new(InferredVectorScannable { + inner: peeked, + output_schema, + })) +} + +struct InferredVectorScannable { + inner: PeekedScannable, + output_schema: SchemaRef, +} + +impl Scannable for InferredVectorScannable { + fn schema(&self) -> SchemaRef { + self.output_schema.clone() + } + + fn scan_as_stream(&mut self) -> SendableRecordBatchStream { + let output_schema = self.output_schema.clone(); + let stream_schema = output_schema.clone(); + let stream = self.inner.scan_as_stream().map(move |batch| { + let batch = batch?; + let columns = batch + .columns() + .iter() + .zip(output_schema.fields()) + .map(|(array, field)| { + if array.data_type() == field.data_type() { + Ok(array.clone()) + } else { + cast_with_options( + array, + field.data_type(), + &CastOptions { + safe: false, + ..Default::default() + }, + ) + .map_err(Error::from) + } + }) + .collect::>>()?; + Ok(RecordBatch::try_new(output_schema.clone(), columns)?) + }); + Box::pin(SimpleRecordBatchStream { + schema: stream_schema, + stream, + }) + } + + fn num_rows(&self) -> Option { + self.inner.num_rows() + } + + fn rescannable(&self) -> bool { + self.inner.rescannable() + } +} + /// Compute the number of write partitions based on data size estimates. /// /// `sample_bytes` and `sample_rows` come from a representative batch and are diff --git a/rust/lancedb/src/lib.rs b/rust/lancedb/src/lib.rs index 70d023ccc..70c9a705e 100644 --- a/rust/lancedb/src/lib.rs +++ b/rust/lancedb/src/lib.rs @@ -72,7 +72,9 @@ //! //! LanceDB uses [arrow-rs](https://github.com/apache/arrow-rs) to define schema, data types and array itself. //! It treats [`FixedSizeList`](https://docs.rs/arrow/latest/arrow/array/struct.FixedSizeListArray.html) -//! columns as vector columns. +//! columns as vector columns. When creating a table with a floating-point `List` +//! column named `vec`, or with `vector` or `embedding` in its name, LanceDB infers +//! a uniform dimension from the first batch and stores it as a `FixedSizeList`. //! //! For more details, please refer to the [LanceDB documentation](https://docs.lancedb.com). //! @@ -82,6 +84,8 @@ //! schema of the `RecordBatch` determines the schema of the table. //! //! Vector columns should be represented as `FixedSizeList` data type. +//! A vector-like `List` input is also accepted when every vector +//! has the same runtime dimension. //! //! ```rust //! # use std::sync::Arc; diff --git a/rust/lancedb/src/query.rs b/rust/lancedb/src/query.rs index b76865043..d24fc66b3 100644 --- a/rust/lancedb/src/query.rs +++ b/rust/lancedb/src/query.rs @@ -1648,8 +1648,8 @@ mod tests { use super::*; use arrow::{array::downcast_array, compute::concat_batches, datatypes::Int32Type}; use arrow_array::{ - FixedSizeListArray, Float32Array, Int32Array, RecordBatch, StringArray, cast::AsArray, - types::Float32Type, + FixedSizeListArray, Float32Array, Int32Array, ListArray, RecordBatch, StringArray, + cast::AsArray, types::Float32Type, }; use arrow_schema::{DataType, Field as ArrowField, Schema as ArrowSchema}; use futures::{StreamExt, TryStreamExt}; @@ -2282,6 +2282,91 @@ mod tests { ); } + #[tokio::test] + async fn vector_search_infers_dimension_from_list_array() { + let tmp_dir = tempdir().unwrap(); + let schema = Arc::new(ArrowSchema::new(vec![ + ArrowField::new("id", DataType::Int32, false), + ArrowField::new( + "vec", + DataType::List(Arc::new(ArrowField::new("item", DataType::Float32, true))), + true, + ), + ])); + let vectors = ListArray::from_iter_primitive::([ + Some([Some(0.0), Some(0.0)]), + Some([Some(1.0), Some(1.0)]), + ]); + let batch = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![0, 1])), Arc::new(vectors)], + ) + .unwrap(); + let table = connect(tmp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap() + .create_table("vectors", batch) + .execute() + .await + .unwrap(); + assert!(matches!( + table.schema().await.unwrap().field(1).data_type(), + DataType::FixedSizeList(_, 2) + )); + + let results = table + .vector_search(&[0.0, 0.0]) + .unwrap() + .limit(1) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + + assert_eq!(results[0]["id"].as_primitive::().value(0), 0); + } + + #[tokio::test] + async fn inferred_vector_dimension_is_validated_across_batches() { + let tmp_dir = tempdir().unwrap(); + let schema = Arc::new(ArrowSchema::new(vec![ArrowField::new( + "vec", + DataType::List(Arc::new(ArrowField::new("item", DataType::Float32, true))), + true, + )])); + let first = RecordBatch::try_new( + schema.clone(), + vec![Arc::new( + ListArray::from_iter_primitive::([Some([Some(0.0), Some(0.0)])]), + )], + ) + .unwrap(); + let wrong_dimension = RecordBatch::try_new( + schema, + vec![Arc::new( + ListArray::from_iter_primitive::([Some([ + Some(1.0), + Some(1.0), + Some(1.0), + ])]), + )], + ) + .unwrap(); + + let result = connect(tmp_dir.path().to_str().unwrap()) + .execute() + .await + .unwrap() + .create_table("vectors", vec![first, wrong_dimension]) + .execute() + .await; + + assert!(result.is_err()); + } + #[tokio::test] async fn test_fast_search_plan() { let tmp_dir = tempdir().unwrap();