diff --git a/include/rabitqlib/index/hnsw/hnsw.hpp b/include/rabitqlib/index/hnsw/hnsw.hpp index 71b59e6..6e5b0af 100644 --- a/include/rabitqlib/index/hnsw/hnsw.hpp +++ b/include/rabitqlib/index/hnsw/hnsw.hpp @@ -1237,9 +1237,15 @@ inline void HierarchicalNSW::searchBaseLayerST_AdaptiveRerankOptDirect( float distk = 1e10; EstimateRecord start_estimate_record; - get_full_est_direct( - q_to_centroids, query_wrapper, ep_id, start_estimate_record - ); + if (ex_bits_ > 0) { + get_full_est_direct( + q_to_centroids, query_wrapper, ep_id, start_estimate_record + ); + } else { + get_bin_est_direct( + q_to_centroids, query_wrapper, ep_id, start_estimate_record + ); + } float est_dist = start_estimate_record.est_dist; float low_dist = start_estimate_record.low_dist; diff --git a/tests/unit/rabitqlib/index/hnsw_test.cpp b/tests/unit/rabitqlib/index/hnsw_test.cpp new file mode 100644 index 0000000..f512624 --- /dev/null +++ b/tests/unit/rabitqlib/index/hnsw_test.cpp @@ -0,0 +1,44 @@ +#include "rabitqlib/index/hnsw/hnsw.hpp" + +#include + +#include +#include + +namespace rabitqlib::hnsw { +namespace { + +TEST(HnswOneBitSearchTest, UsesBinaryEstimateForEntryPoint) { + constexpr size_t dim = 64; + constexpr size_t count = 8; + std::vector data(count * dim); + std::vector centroid(dim); + std::vector query(dim); + + for (size_t i = 0; i < dim; ++i) { + centroid[i] = static_cast(static_cast(i % 13) - 6) / 5.0F; + query[i] = centroid[i]; + } + for (size_t point = 0; point < count; ++point) { + for (size_t i = 0; i < dim; ++i) { + data[(point * dim) + i] = + centroid[i] + static_cast((point + 1) * ((i % 11) + 1)); + } + } + + std::vector cluster_ids(count, 0); + HierarchicalNSW index(count, dim, 1, 2, 10, 100, METRIC_L2); + index.construct(1, centroid.data(), count, data.data(), cluster_ids.data(), 1, false); + + const auto results = index.search(query.data(), 1, 1, 10, 1); + + ASSERT_EQ(results.size(), 1U); + ASSERT_EQ(results[0].size(), 1U); + ASSERT_EQ(results[0][0].second, 0U); + EXPECT_TRUE(std::isfinite(results[0][0].first)); + const float exact_distance = euclidean_sqr(query.data(), data.data(), dim); + EXPECT_NEAR(results[0][0].first, exact_distance, exact_distance * 0.1F); +} + +} // namespace +} // namespace rabitqlib::hnsw