fix(vamana): honor asymmetric query metrics (#635)

This commit is contained in:
luoxiaojian 2026-07-30 20:11:37 +08:00 committed by GitHub
parent d15a37e425
commit ad36b75966
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 113 additions and 0 deletions

View File

@ -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_;
}

View File

@ -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;

View File

@ -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<IndexContext *>(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();

View File

@ -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_{};

View File

@ -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;