diff --git a/src/core/algorithm/vamana/vamana_context.h b/src/core/algorithm/vamana/vamana_context.h index a57635357..531760ce2 100644 --- a/src/core/algorithm/vamana/vamana_context.h +++ b/src/core/algorithm/vamana/vamana_context.h @@ -115,6 +115,11 @@ class VamanaContext : public IndexContext { inline VamanaDistCalculator &dist_calculator() { return dc_; } + inline void update_dist_calculator_distance( + const IndexMetric::MatrixDistance &distance, + const IndexMetric::MatrixBatchDistance &batch_distance) { + dc_.update_distance(distance, batch_distance); + } inline TopkHeap &topk_heap() { return topk_heap_; } diff --git a/src/core/algorithm/vamana/vamana_dist_calculator.h b/src/core/algorithm/vamana/vamana_dist_calculator.h index ed15b0bdf..6ecbfa434 100644 --- a/src/core/algorithm/vamana/vamana_dist_calculator.h +++ b/src/core/algorithm/vamana/vamana_dist_calculator.h @@ -65,6 +65,13 @@ class VamanaDistCalculator { dim_ = dim; } + inline void update_distance( + const IndexMetric::MatrixDistance &distance, + const IndexMetric::MatrixBatchDistance &batch_distance) { + distance_ = distance; + batch_distance_ = batch_distance; + } + inline void reset_query(const void *query) { error_ = false; query_ = query; diff --git a/src/core/algorithm/vamana/vamana_streamer.cc b/src/core/algorithm/vamana/vamana_streamer.cc index 0c08a8238..d651c486b 100644 --- a/src/core/algorithm/vamana/vamana_streamer.cc +++ b/src/core/algorithm/vamana/vamana_streamer.cc @@ -262,6 +262,18 @@ int VamanaStreamer::open(IndexStorage::Pointer stg) { return IndexError_InvalidArgument; } + add_distance_ = metric_->distance(); + add_batch_distance_ = metric_->batch_distance(); + search_distance_ = add_distance_; + search_batch_distance_ = add_batch_distance_; + + const auto query_metric = metric_->query_metric(); + if (query_metric && query_metric->distance() && + query_metric->batch_distance()) { + search_distance_ = query_metric->distance(); + search_batch_distance_ = query_metric->batch_distance(); + } + // Create algorithm based on entity storage mode switch (entity_->storage_mode()) { case VamanaStorageMode::kBufferPool: @@ -451,6 +463,7 @@ int VamanaStreamer::add_impl(uint64_t pkey, const void *query, AILEGO_DEFER([&]() { shared_mutex_.unlock_shared(); }); ctx->clear(); + ctx->update_dist_calculator_distance(add_distance_, add_batch_distance_); ctx->check_need_adjuct_ctx(entity_->doc_cnt()); if (metric_->support_train()) { @@ -522,6 +535,7 @@ int VamanaStreamer::add_with_id_impl(uint32_t id, const void *query, AILEGO_DEFER([&]() { shared_mutex_.unlock_shared(); }); ctx->clear(); + ctx->update_dist_calculator_distance(add_distance_, add_batch_distance_); ctx->check_need_adjuct_ctx(entity_->doc_cnt()); if (metric_->support_train()) { @@ -584,6 +598,8 @@ int VamanaStreamer::search_impl(const void *query, const IndexQueryMeta &qmeta, } ctx->clear(); + ctx->update_dist_calculator_distance(search_distance_, + search_batch_distance_); ctx->resize_results(count); ctx->check_need_adjuct_ctx(entity_->doc_cnt()); @@ -645,6 +661,8 @@ int VamanaStreamer::search_bf_impl(const void *query, } ctx->clear(); + ctx->update_dist_calculator_distance(search_distance_, + search_batch_distance_); ctx->resize_results(count); const auto &filter = static_cast(ctx)->filter(); @@ -686,6 +704,8 @@ int VamanaStreamer::search_bf_by_p_keys_impl( } ctx->clear(); + ctx->update_dist_calculator_distance(search_distance_, + search_batch_distance_); ctx->resize_results(count); auto &topk = ctx->topk_heap(); diff --git a/src/core/algorithm/vamana/vamana_streamer.h b/src/core/algorithm/vamana/vamana_streamer.h index a1217c9ac..3446c73ed 100644 --- a/src/core/algorithm/vamana/vamana_streamer.h +++ b/src/core/algorithm/vamana/vamana_streamer.h @@ -150,6 +150,11 @@ class VamanaStreamer : public IndexStreamer { IndexMeta meta_{}; IndexMetric::Pointer metric_{}; + IndexMetric::MatrixDistance add_distance_{}; + IndexMetric::MatrixDistance search_distance_{}; + IndexMetric::MatrixBatchDistance add_batch_distance_{}; + IndexMetric::MatrixBatchDistance search_batch_distance_{}; + Stats stats_{}; std::mutex mutex_{}; diff --git a/tests/core/algorithm/vamana/vamana_streamer_test.cc b/tests/core/algorithm/vamana/vamana_streamer_test.cc index 2454d64b3..12f833c3f 100644 --- a/tests/core/algorithm/vamana/vamana_streamer_test.cc +++ b/tests/core/algorithm/vamana/vamana_streamer_test.cc @@ -785,6 +785,82 @@ TEST_F(VamanaStreamerTest, TestConcurrentBuild) { ASSERT_GT(result.size(), 0UL); } +TEST_F(VamanaStreamerTest, TestAsymmetricQueryMetric) { + constexpr size_t kTestDimension = 2; + + ailego::Params metric_params; + metric_params.set("proxima.mips_euclidean.metric.injection_type", 0); + IndexMeta meta(IndexMeta::DataType::DT_FP32, kTestDimension); + meta.set_metric("MipsSquaredEuclidean", 0, metric_params); + + ailego::Params params; + params.set(PARAM_VAMANA_STREAMER_MAX_DEGREE, 8U); + params.set(PARAM_VAMANA_STREAMER_SEARCH_LIST_SIZE, 16U); + params.set(PARAM_VAMANA_STREAMER_ALPHA, 1.2f); + params.set(PARAM_VAMANA_STREAMER_EF, 16U); + params.set(PARAM_VAMANA_STREAMER_BRUTE_FORCE_THRESHOLD, 0U); + + auto streamer = IndexFactory::CreateStreamer("VamanaStreamer"); + ASSERT_TRUE(streamer); + ASSERT_EQ(0, streamer->init(meta, params)); + + auto storage = IndexFactory::CreateStorage("MMapFileStorage"); + ASSERT_TRUE(storage); + ASSERT_EQ(0, storage->init(ailego::Params())); + ASSERT_EQ(0, storage->open(dir_ + "TestAsymmetricQueryMetric.index", true)); + ASSERT_EQ(0, streamer->open(storage)); + + IndexQueryMeta query_meta(IndexMeta::DataType::DT_FP32, kTestDimension); + NumericalVector unit_record(kTestDimension); + unit_record[0] = 1.0f; + unit_record[1] = 0.0f; + NumericalVector scaled_record(kTestDimension); + scaled_record[0] = 2.0f; + scaled_record[1] = 0.0f; + + auto context = streamer->create_context(); + ASSERT_TRUE(context); + ASSERT_EQ(0, streamer->add_impl(10, unit_record.data(), query_meta, context)); + ASSERT_EQ(0, + streamer->add_impl(20, scaled_record.data(), query_meta, context)); + + auto *vamana_context = dynamic_cast(context.get()); + ASSERT_TRUE(vamana_context); + EXPECT_FLOAT_EQ(1.0f, vamana_context->dist_calculator().dist( + unit_record.data(), scaled_record.data())); + + context->set_topk(2); + ASSERT_EQ(0, + streamer->search_bf_impl(unit_record.data(), query_meta, context)); + ASSERT_EQ(2UL, context->result().size()); + EXPECT_EQ(20UL, context->result()[0].key()); + EXPECT_FLOAT_EQ(-2.0f, context->result()[0].score()); + EXPECT_EQ(10UL, context->result()[1].key()); + EXPECT_FLOAT_EQ(-1.0f, context->result()[1].score()); + EXPECT_FLOAT_EQ(-2.0f, vamana_context->dist_calculator().dist( + unit_record.data(), scaled_record.data())); + + auto graph_context = streamer->create_context(); + ASSERT_TRUE(graph_context); + graph_context->set_topk(1); + ASSERT_EQ( + 0, streamer->search_impl(unit_record.data(), query_meta, graph_context)); + ASSERT_EQ(1UL, graph_context->result().size()); + EXPECT_EQ(20UL, graph_context->result()[0].key()); + EXPECT_FLOAT_EQ(-2.0f, graph_context->result()[0].score()); + + auto primary_key_context = streamer->create_context(); + ASSERT_TRUE(primary_key_context); + primary_key_context->set_topk(2); + const std::vector> primary_keys{{10, 20}}; + ASSERT_EQ( + 0, streamer->search_bf_by_p_keys_impl(unit_record.data(), primary_keys, + query_meta, primary_key_context)); + ASSERT_EQ(2UL, primary_key_context->result().size()); + EXPECT_EQ(20UL, primary_key_context->result()[0].key()); + EXPECT_FLOAT_EQ(-2.0f, primary_key_context->result()[0].score()); +} + // Test Vamana + INT8 quantization + rotation end-to-end TEST_F(VamanaStreamerTest, TestInt8WithRotate) { constexpr size_t kTestDim = 128;