Skip to content
Merged
Show file tree
Hide file tree
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
5 changes: 5 additions & 0 deletions src/core/algorithm/vamana/vamana_context.h
Original file line number Diff line number Diff line change
Expand Up @@ -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_;
}
Expand Down
7 changes: 7 additions & 0 deletions src/core/algorithm/vamana/vamana_dist_calculator.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
20 changes: 20 additions & 0 deletions src/core/algorithm/vamana/vamana_streamer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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()) {
Expand Down Expand Up @@ -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()) {
Expand Down Expand Up @@ -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());

Expand Down Expand Up @@ -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<IndexContext *>(ctx)->filter();
Expand Down Expand Up @@ -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();
Expand Down
5 changes: 5 additions & 0 deletions src/core/algorithm/vamana/vamana_streamer.h
Original file line number Diff line number Diff line change
Expand Up @@ -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_{};

Expand Down
76 changes: 76 additions & 0 deletions tests/core/algorithm/vamana/vamana_streamer_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<float> unit_record(kTestDimension);
unit_record[0] = 1.0f;
unit_record[1] = 0.0f;
NumericalVector<float> 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<VamanaContext *>(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<std::vector<uint64_t>> 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;
Expand Down
Loading