Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
179 changes: 178 additions & 1 deletion rust/lance/src/dataset/scanner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2200,12 +2200,12 @@ impl Scanner {
/// Default value is false.
///
/// This is essentially a weak consistency search, only on the indexed data.
/// Row IDs are used internally but are only returned when explicitly requested.
pub fn fast_search(&mut self) -> &mut Self {
if let Some(q) = self.nearest.as_mut() {
q.use_index = true;
}
self.fast_search = true;
self.projection_plan.include_row_id(); // fast search requires _rowid
self
}

Expand Down Expand Up @@ -8521,6 +8521,7 @@ mod test {
use lance_index::scalar::inverted::query::{
BooleanQuery, BoostQuery, FtsQuery, MatchQuery, MultiMatchQuery, Occur, PhraseQuery,
};
use lance_index::scalar::inverted::tokenizer::InvertedIndexParams;
use lance_index::vector::hnsw::builder::HnswBuildParams;
use lance_index::vector::ivf::IvfBuildParams;
use lance_index::vector::pq::PQBuildParams;
Expand Down Expand Up @@ -17330,6 +17331,166 @@ full_filter=name LIKE Utf8(\"test%2\"), refine_filter=name LIKE Utf8(\"test%2\")
.unwrap();
}

#[rstest]
#[case::vector(IndexType::Vector, DIST_COL, vec![2, 1])]
#[case::fts(IndexType::Inverted, SCORE_COL, vec![0, 2])]
#[tokio::test]
async fn test_fast_search_projection(
#[case] index_type: IndexType,
#[case] scoring_column: &str,
#[case] expected_ids: Vec<i32>,
#[values(false, true)] has_projection: bool,
#[values(false, true)] is_row_id_requested: bool,
#[values(false, true)] has_stable_row_ids: bool,
) {
let mut dataset = gen_batch()
.col("id", array::step::<Int32Type>())
.col("text", array::cycle_utf8_literals(&["quick", "slow"]))
.col(
"vec",
array::cycle_vec(array::step::<Float32Type>(), Dimension::from(2)),
)
.into_ram_dataset_with_params(
FragmentCount::from(2),
FragmentRowCount::from(2),
Some(WriteParams {
max_rows_per_file: 2,
enable_stable_row_ids: has_stable_row_ids,
..Default::default()
}),
)
.await
.unwrap();

if index_type == IndexType::Vector {
let centroids = Arc::new(
FixedSizeListArray::try_new_from_values(Float32Array::from(vec![0.0, 1.0]), 2)
.unwrap(),
);
let params = VectorIndexParams::with_ivf_flat_params(
MetricType::L2,
IvfBuildParams::try_with_centroids(1, centroids).unwrap(),
);
dataset
.create_index(&["vec"], index_type, None, &params, true)
.await
.unwrap();
} else {
dataset
.create_index(
&["text"],
index_type,
None,
&InvertedIndexParams::default(),
true,
)
.await
.unwrap();
}

let mut scanner = dataset.scan();
if index_type == IndexType::Vector {
scanner
.nearest("vec", &Float32Array::from(vec![4.0, 5.0]), 2)
.unwrap();
} else {
scanner
.full_text_search(FullTextSearchQuery::new("quick".to_owned()))
.unwrap();
}
let mut expected_columns = if has_projection {
scanner.project(&["id"]).unwrap();
vec!["id"]
} else {
vec!["id", "text", "vec"]
};
expected_columns.push(scoring_column);
if is_row_id_requested {
scanner.with_row_id();
expected_columns.push(ROW_ID);
}
scanner.limit(Some(2), None).unwrap();
let normal_batch = scanner.try_into_batch().await.unwrap();

scanner.fast_search();
let fast_batch = scanner.try_into_batch().await.unwrap();
assert_eq!(
fast_batch
.schema()
.fields()
.iter()
.map(|f| f.name().as_str())
.collect::<Vec<_>>(),
expected_columns
);
assert_eq!(fast_batch, normal_batch);
assert_eq!(fast_batch.num_rows(), 2);
let recall = fast_batch["id"]
.as_primitive::<Int32Type>()
.values()
.iter()
.filter(|id| expected_ids.contains(id))
.count() as f64
/ expected_ids.len() as f64;
assert_eq!(recall, 1.0);

scanner.project(&["id", ROW_ID]).unwrap();
let explicit_row_id_batch = scanner.try_into_batch().await.unwrap();
assert_eq!(
explicit_row_id_batch
.schema()
.fields()
.iter()
.map(|f| f.name().as_str())
.collect::<Vec<_>>(),
if is_row_id_requested {
vec!["id", scoring_column, ROW_ID]
} else {
vec!["id", ROW_ID, scoring_column]
}
);
assert_eq!(&explicit_row_id_batch["id"], &fast_batch["id"]);
assert_eq!(explicit_row_id_batch[ROW_ID].null_count(), 0);
}

#[rstest]
#[tokio::test]
async fn test_fast_search_requires_row_id_for_deleted_rows(
#[values(false, true)] is_row_id_requested: bool,
) {
let mut dataset = gen_batch()
.col("id", array::step::<Int32Type>())
.into_ram_dataset(FragmentCount::from(2), FragmentRowCount::from(2))
.await
.unwrap();
dataset.delete("id = 1").await.unwrap();

let mut scanner = dataset.scan();
scanner
.project(&["id"])
.unwrap()
.fast_search()
.include_deleted_rows();
if is_row_id_requested {
scanner.with_row_id();
}

let result = scanner.try_into_batch().await;
if is_row_id_requested {
let batch = result.unwrap();
assert_eq!(batch.num_rows(), 4);
assert_eq!(batch[ROW_ID].null_count(), 1);
} else {
let error = result.unwrap_err();
assert!(matches!(error, Error::InvalidInput { .. }));
assert!(
error
.to_string()
.contains("include_deleted_rows is set but with_row_id is false")
);
}
}

#[tokio::test]
async fn test_fast_search_without_vector_index_returns_empty() {
let dataset = TestVectorDataset::new(LanceFileVersion::Stable, true)
Expand Down Expand Up @@ -17427,6 +17588,22 @@ full_filter=name LIKE Utf8(\"test%2\"), refine_filter=name LIKE Utf8(\"test%2\")

assert_eq!(normal_batch.num_rows(), 15);
assert_eq!(fast_batch.num_rows(), 5);

let mut project_first_scanner = dataset.dataset.scan();
project_first_scanner
.filter("i >= 395")
.unwrap()
.project(&["i"])
.unwrap()
.fast_search();
assert_eq!(
project_first_scanner.explain_plan(false).await.unwrap(),
scanner.explain_plan(false).await.unwrap()
);
assert_eq!(
project_first_scanner.try_into_batch().await.unwrap(),
fast_batch
);
}

fn make_scalar_filter_test_batch(schema: SchemaRef, start: i32, end: i32) -> RecordBatch {
Expand Down
Loading