diff --git a/lib/segment/src/data_types/vectors.rs b/lib/segment/src/data_types/vectors.rs index dd2bb42ee2..43087fbbc2 100644 --- a/lib/segment/src/data_types/vectors.rs +++ b/lib/segment/src/data_types/vectors.rs @@ -96,6 +96,18 @@ impl TryFrom for SparseVector { } } +impl TryFrom for MultiDenseVector { + type Error = OperationError; + + fn try_from(value: Vector) -> Result { + match value { + Vector::Dense(_) => Err(OperationError::WrongMulti), + Vector::Sparse(_) => Err(OperationError::WrongSparse), + Vector::MultiDense(v) => Ok(v), + } + } +} + impl<'a> From<&'a [VectorElementType]> for VectorRef<'a> { fn from(val: &'a [VectorElementType]) -> Self { VectorRef::Dense(val) @@ -194,6 +206,18 @@ impl<'a> TryInto<&'a SparseVector> for &'a Vector { } } +impl<'a> TryInto<&'a [DenseVector]> for &'a Vector { + type Error = OperationError; + + fn try_into(self) -> Result<&'a [DenseVector], Self::Error> { + match self { + Vector::Dense(_) => Err(OperationError::WrongMulti), + Vector::Sparse(_) => Err(OperationError::WrongSparse), + Vector::MultiDense(v) => Ok(v), + } + } +} + pub fn default_vector(vec: DenseVector) -> NamedVectors<'static> { NamedVectors::from([(DEFAULT_VECTOR_NAME.to_owned(), vec)]) } diff --git a/lib/segment/src/vector_storage/memmap_dense_vector_storage.rs b/lib/segment/src/vector_storage/memmap_dense_vector_storage.rs index 6412e72ff6..8234a806de 100644 --- a/lib/segment/src/vector_storage/memmap_dense_vector_storage.rs +++ b/lib/segment/src/vector_storage/memmap_dense_vector_storage.rs @@ -285,7 +285,7 @@ mod tests { assert_eq!(borrowed_storage.total_vector_count(), 3); let vector = borrowed_storage.get_vector(1).to_owned(); - let vector: Vec<_> = vector.try_into().unwrap(); + let vector: DenseVector = vector.try_into().unwrap(); assert_eq!(points[1], vector); diff --git a/lib/segment/src/vector_storage/mod.rs b/lib/segment/src/vector_storage/mod.rs index c9fed5211d..060a27f0de 100644 --- a/lib/segment/src/vector_storage/mod.rs +++ b/lib/segment/src/vector_storage/mod.rs @@ -22,7 +22,7 @@ mod bitvec; pub mod common; pub mod query; mod query_scorer; -mod simple_multi_dense_vector_storage; +pub mod simple_multi_dense_vector_storage; pub mod simple_sparse_vector_storage; pub use raw_scorer::*; diff --git a/lib/segment/src/vector_storage/query_scorer/mod.rs b/lib/segment/src/vector_storage/query_scorer/mod.rs index 0a20616386..2990966ae8 100644 --- a/lib/segment/src/vector_storage/query_scorer/mod.rs +++ b/lib/segment/src/vector_storage/query_scorer/mod.rs @@ -28,7 +28,7 @@ pub fn score_multi( sum += multidense_b .iter() .map(|dense_b| TMetric::similarity(dense_a, dense_b)) - .fold(0.0, |a, b| if a > b { a } else { b }); + .fold(ScoreType::NEG_INFINITY, |a, b| if a > b { a } else { b }); } sum } diff --git a/lib/segment/src/vector_storage/query_scorer/multi_custom_query_scorer.rs b/lib/segment/src/vector_storage/query_scorer/multi_custom_query_scorer.rs index 98c64052a5..64b8c0ac3a 100644 --- a/lib/segment/src/vector_storage/query_scorer/multi_custom_query_scorer.rs +++ b/lib/segment/src/vector_storage/query_scorer/multi_custom_query_scorer.rs @@ -1,6 +1,7 @@ use std::marker::PhantomData; use common::types::{PointOffsetType, ScoreType}; +use itertools::Itertools; use super::score_multi; use crate::data_types::vectors::MultiDenseVector; @@ -9,7 +10,7 @@ use crate::vector_storage::query::{Query, TransformInto}; use crate::vector_storage::query_scorer::QueryScorer; use crate::vector_storage::MultiVectorStorage; -pub struct CustomQueryScorer< +pub struct MultiCustomQueryScorer< 'a, TMetric: Metric, TVectorStorage: MultiVectorStorage, @@ -24,13 +25,12 @@ impl< 'a, TMetric: Metric, TVectorStorage: MultiVectorStorage, - TQuery: Query + TransformInto, - > CustomQueryScorer<'a, TMetric, TVectorStorage, TQuery> + TQuery: Query + TransformInto, + > MultiCustomQueryScorer<'a, TMetric, TVectorStorage, TQuery> { - #[allow(dead_code)] pub fn new(query: TQuery, vector_storage: &'a TVectorStorage) -> Self { let query = query - .transform(|vector| Ok(TMetric::preprocess(vector))) + .transform(|vector| Ok(vector.into_iter().map(TMetric::preprocess).collect_vec())) .unwrap(); Self { @@ -42,7 +42,7 @@ impl< } impl<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage, TQuery: Query> - QueryScorer for CustomQueryScorer<'a, TMetric, TVectorStorage, TQuery> + QueryScorer for MultiCustomQueryScorer<'a, TMetric, TVectorStorage, TQuery> { #[inline] fn score_stored(&self, idx: PointOffsetType) -> ScoreType { diff --git a/lib/segment/src/vector_storage/query_scorer/multi_metric_query_scorer.rs b/lib/segment/src/vector_storage/query_scorer/multi_metric_query_scorer.rs index fe4cb7abe7..acd30ad9fd 100644 --- a/lib/segment/src/vector_storage/query_scorer/multi_metric_query_scorer.rs +++ b/lib/segment/src/vector_storage/query_scorer/multi_metric_query_scorer.rs @@ -8,16 +8,15 @@ use crate::spaces::metric::Metric; use crate::vector_storage::query_scorer::QueryScorer; use crate::vector_storage::MultiVectorStorage; -pub struct MetricQueryScorer<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage> { +pub struct MultiMetricQueryScorer<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage> { vector_storage: &'a TVectorStorage, query: MultiDenseVector, metric: PhantomData, } impl<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage> - MetricQueryScorer<'a, TMetric, TVectorStorage> + MultiMetricQueryScorer<'a, TMetric, TVectorStorage> { - #[allow(dead_code)] pub fn new(query: MultiDenseVector, vector_storage: &'a TVectorStorage) -> Self { Self { query: query.into_iter().map(|v| TMetric::preprocess(v)).collect(), @@ -28,7 +27,7 @@ impl<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage> } impl<'a, TMetric: Metric, TVectorStorage: MultiVectorStorage> QueryScorer - for MetricQueryScorer<'a, TMetric, TVectorStorage> + for MultiMetricQueryScorer<'a, TMetric, TVectorStorage> { #[inline] fn score_stored(&self, idx: PointOffsetType) -> ScoreType { diff --git a/lib/segment/src/vector_storage/raw_scorer.rs b/lib/segment/src/vector_storage/raw_scorer.rs index d7e87f1cd3..d640995e5a 100644 --- a/lib/segment/src/vector_storage/raw_scorer.rs +++ b/lib/segment/src/vector_storage/raw_scorer.rs @@ -9,15 +9,17 @@ use super::query::discovery_query::DiscoveryQuery; use super::query::reco_query::RecoQuery; use super::query::TransformInto; use super::query_scorer::custom_query_scorer::CustomQueryScorer; +use super::query_scorer::multi_custom_query_scorer::MultiCustomQueryScorer; use super::query_scorer::sparse_custom_query_scorer::SparseCustomQueryScorer; -use super::{DenseVectorStorage, SparseVectorStorage, VectorStorageEnum}; +use super::{DenseVectorStorage, MultiVectorStorage, SparseVectorStorage, VectorStorageEnum}; use crate::common::operation_error::{OperationError, OperationResult}; -use crate::data_types::vectors::{DenseVector, QueryVector}; +use crate::data_types::vectors::{DenseVector, MultiDenseVector, QueryVector}; use crate::spaces::metric::Metric; use crate::spaces::simple::{CosineMetric, DotProductMetric, EuclidMetric, ManhattanMetric}; use crate::spaces::tools::peek_top_largest_iterable; use crate::types::Distance; use crate::vector_storage::query_scorer::metric_query_scorer::MetricQueryScorer; +use crate::vector_storage::query_scorer::multi_metric_query_scorer::MultiMetricQueryScorer; use crate::vector_storage::query_scorer::QueryScorer; /// RawScorer composition: @@ -135,9 +137,8 @@ pub fn new_stoppable_raw_scorer<'a>( VectorStorageEnum::SparseSimple(vs) => { raw_sparse_scorer_impl(query, vs, point_deleted, is_stopped) } - VectorStorageEnum::MultiDenseSimple(_vs) => { - // TODO(colbert) - unimplemented!("multidense vector") + VectorStorageEnum::MultiDenseSimple(vs) => { + raw_multi_scorer_impl(query, vs, point_deleted, is_stopped) } } } @@ -290,6 +291,85 @@ where })) } +pub fn raw_multi_scorer_impl<'a, TVectorStorage: MultiVectorStorage>( + query: QueryVector, + vector_storage: &'a TVectorStorage, + point_deleted: &'a BitSlice, + is_stopped: &'a AtomicBool, +) -> OperationResult> { + match vector_storage.distance() { + Distance::Cosine => new_multi_scorer_with_metric::( + query, + vector_storage, + point_deleted, + is_stopped, + ), + Distance::Euclid => new_multi_scorer_with_metric::( + query, + vector_storage, + point_deleted, + is_stopped, + ), + Distance::Dot => new_multi_scorer_with_metric::( + query, + vector_storage, + point_deleted, + is_stopped, + ), + Distance::Manhattan => new_multi_scorer_with_metric::( + query, + vector_storage, + point_deleted, + is_stopped, + ), + } +} + +fn new_multi_scorer_with_metric<'a, TMetric: Metric + 'a, TVectorStorage: MultiVectorStorage>( + query: QueryVector, + vector_storage: &'a TVectorStorage, + point_deleted: &'a BitSlice, + is_stopped: &'a AtomicBool, +) -> OperationResult> { + let vec_deleted = vector_storage.deleted_vector_bitslice(); + match query { + QueryVector::Nearest(vector) => raw_scorer_from_query_scorer( + MultiMetricQueryScorer::::new(vector.try_into()?, vector_storage), + point_deleted, + vec_deleted, + is_stopped, + ), + QueryVector::Recommend(reco_query) => { + let reco_query: RecoQuery = reco_query.transform_into()?; + raw_scorer_from_query_scorer( + MultiCustomQueryScorer::::new(reco_query, vector_storage), + point_deleted, + vec_deleted, + is_stopped, + ) + } + QueryVector::Discovery(discovery_query) => { + let discovery_query: DiscoveryQuery = + discovery_query.transform_into()?; + raw_scorer_from_query_scorer( + MultiCustomQueryScorer::::new(discovery_query, vector_storage), + point_deleted, + vec_deleted, + is_stopped, + ) + } + QueryVector::Context(context_query) => { + let context_query: ContextQuery = context_query.transform_into()?; + raw_scorer_from_query_scorer( + MultiCustomQueryScorer::::new(context_query, vector_storage), + point_deleted, + vec_deleted, + is_stopped, + ) + } + } +} + impl<'a, TVector, TQueryScorer> RawScorer for RawScorerImpl<'a, TVector, TQueryScorer> where TVector: ?Sized, diff --git a/lib/segment/src/vector_storage/simple_multi_dense_vector_storage.rs b/lib/segment/src/vector_storage/simple_multi_dense_vector_storage.rs index c05da748f4..23fe6bd6f3 100644 --- a/lib/segment/src/vector_storage/simple_multi_dense_vector_storage.rs +++ b/lib/segment/src/vector_storage/simple_multi_dense_vector_storage.rs @@ -8,6 +8,7 @@ use common::types::PointOffsetType; use parking_lot::RwLock; use rocksdb::DB; +use super::MultiVectorStorage; use crate::common::operation_error::{check_process_stopped, OperationError, OperationResult}; use crate::common::rocksdb_wrapper::DatabaseColumnWrapper; use crate::common::Flusher; @@ -21,7 +22,6 @@ use crate::vector_storage::{VectorStorage, VectorStorageEnum}; type StoredMultiDenseVector = StoredRecord; /// In-memory vector storage with on-update persistence using `store` -#[allow(unused)] pub struct SimpleMultiDenseVectorStorage { dim: usize, distance: Distance, @@ -124,6 +124,12 @@ impl SimpleMultiDenseVectorStorage { } } +impl MultiVectorStorage for SimpleMultiDenseVectorStorage { + fn get_multi(&self, key: PointOffsetType) -> &MultiDenseVector { + self.vectors.get(key as usize).expect("vector not found") + } +} + impl VectorStorage for SimpleMultiDenseVectorStorage { fn vector_dim(&self) -> usize { self.dim diff --git a/lib/segment/tests/integration/main.rs b/lib/segment/tests/integration/main.rs index bd413ebe91..2325e816e6 100644 --- a/lib/segment/tests/integration/main.rs +++ b/lib/segment/tests/integration/main.rs @@ -9,6 +9,7 @@ pub mod filtrable_hnsw_test; pub mod fixtures; pub mod hnsw_discover_test; pub mod hnsw_quantized_search_test; +mod multivector_hnsw_test; pub mod nested_filtering_test; pub mod payload_index_test; pub mod scroll_filtering_test; diff --git a/lib/segment/tests/integration/multivector_hnsw_test.rs b/lib/segment/tests/integration/multivector_hnsw_test.rs new file mode 100644 index 0000000000..87120020a1 --- /dev/null +++ b/lib/segment/tests/integration/multivector_hnsw_test.rs @@ -0,0 +1,171 @@ +use std::collections::HashMap; +use std::sync::atomic::AtomicBool; +use std::sync::Arc; + +use common::cpu::CpuPermit; +use rand::prelude::StdRng; +use rand::SeedableRng; +use segment::common::rocksdb_wrapper::{open_db, DB_VECTOR_CF}; +use segment::data_types::vectors::{ + only_default_vector, QueryVector, Vector, VectorRef, DEFAULT_VECTOR_NAME, +}; +use segment::entry::entry_point::SegmentEntry; +use segment::fixtures::index_fixtures::random_vector; +use segment::fixtures::payload_fixtures::random_int_payload; +use segment::index::hnsw_index::graph_links::GraphLinksRam; +use segment::index::hnsw_index::hnsw::HNSWIndex; +use segment::index::VectorIndex; +use segment::segment_constructor::build_segment; +use segment::types::{ + Condition, Distance, FieldCondition, Filter, HnswConfig, Indexes, Payload, PayloadSchemaType, + SegmentConfig, SeqNumberType, VectorDataConfig, VectorStorageType, +}; +use segment::vector_storage::simple_multi_dense_vector_storage::open_simple_multi_dense_vector_storage; +use segment::vector_storage::VectorStorage; +use serde_json::json; +use tempfile::Builder; + +use crate::utils::path; + +#[test] +fn test_single_multi_and_dense_hnsw_equivalency() { + let num_vectors: u64 = 1_000; + let distance = Distance::Cosine; + let num_payload_values = 2; + let dim = 8; + + let mut rnd = StdRng::seed_from_u64(42); + + let dir = Builder::new().prefix("segment_dir").tempdir().unwrap(); + + let config = SegmentConfig { + vector_data: HashMap::from([( + DEFAULT_VECTOR_NAME.to_owned(), + VectorDataConfig { + size: dim, + distance, + storage_type: VectorStorageType::Memory, + index: Indexes::Plain {}, + quantization_config: None, + }, + )]), + sparse_vector_data: Default::default(), + payload_storage_type: Default::default(), + }; + + let int_key = "int"; + + let mut segment = build_segment(dir.path(), &config, true).unwrap(); + + segment + .create_field_index(0, &path(int_key), Some(&PayloadSchemaType::Integer.into())) + .unwrap(); + + let dir = Builder::new().prefix("storage_dir").tempdir().unwrap(); + let db = open_db(dir.path(), &[DB_VECTOR_CF]).unwrap(); + let multi_storage = open_simple_multi_dense_vector_storage( + db, + DB_VECTOR_CF, + dim, + distance, + &AtomicBool::new(false), + ) + .unwrap(); + + for n in 0..num_vectors { + let idx = n.into(); + let vector = random_vector(&mut rnd, dim); + let vector_multi = vec![distance.preprocess_vector(vector.clone())]; + + let int_payload = random_int_payload(&mut rnd, num_payload_values..=num_payload_values); + let payload: Payload = json!({int_key:int_payload,}).into(); + + segment + .upsert_point(n as SeqNumberType, idx, only_default_vector(&vector)) + .unwrap(); + segment + .set_full_payload(n as SeqNumberType, idx, &payload) + .unwrap(); + + let internal_id = segment.id_tracker.borrow().internal_id(idx).unwrap(); + multi_storage + .borrow_mut() + .insert_vector(internal_id, VectorRef::MultiDense(&vector_multi)) + .unwrap(); + } + + let hnsw_dir = Builder::new().prefix("hnsw_dir").tempdir().unwrap(); + + let stopped = AtomicBool::new(false); + + let m = 8; + let ef_construct = 100; + let full_scan_threshold = 10000; + + let hnsw_config = HnswConfig { + m, + ef_construct, + full_scan_threshold, + max_indexing_threads: 2, + on_disk: Some(false), + payload_m: None, + }; + + // single threaded mode to guarantee equivalency between single and multi hnsw + let permit = Arc::new(CpuPermit::dummy(1)); + + let vector_storage = &segment.vector_data[DEFAULT_VECTOR_NAME].vector_storage; + let quantized_vectors = &segment.vector_data[DEFAULT_VECTOR_NAME].quantized_vectors; + let mut hnsw_index_dense = HNSWIndex::::open( + hnsw_dir.path(), + segment.id_tracker.clone(), + vector_storage.clone(), + quantized_vectors.clone(), + segment.payload_index.clone(), + hnsw_config.clone(), + ) + .unwrap(); + hnsw_index_dense + .build_index(permit.clone(), &stopped) + .unwrap(); + + let mut hnsw_index_multi = HNSWIndex::::open( + hnsw_dir.path(), + segment.id_tracker.clone(), + multi_storage, + quantized_vectors.clone(), + segment.payload_index.clone(), + hnsw_config, + ) + .unwrap(); + hnsw_index_multi.build_index(permit, &stopped).unwrap(); + + for _ in 0..10 { + let random_vector = random_vector(&mut rnd, dim); + let query_vector = random_vector.clone().into(); + let query_vector_multi = QueryVector::Nearest(Vector::from(vec![random_vector])); + + let payload_value = random_int_payload(&mut rnd, 1..=1).pop().unwrap(); + + let filter = Filter::new_must(Condition::Field(FieldCondition::new_match( + path(int_key), + payload_value.into(), + ))); + + let search_res_dense = hnsw_index_dense + .search(&[&query_vector], Some(&filter), 10, None, &false.into()) + .unwrap(); + + let search_res_multi = hnsw_index_multi + .search( + &[&query_vector_multi], + Some(&filter), + 10, + None, + &false.into(), + ) + .unwrap(); + + assert_eq!(search_res_dense, search_res_multi); + } +}