diff --git a/include/rabitqlib/index/estimator.hpp b/include/rabitqlib/index/estimator.hpp index ed121aa..6d93dff 100644 --- a/include/rabitqlib/index/estimator.hpp +++ b/include/rabitqlib/index/estimator.hpp @@ -162,13 +162,52 @@ template inline void qg_batch_estdist( const char* batch_data, const BatchQuery& q_obj, size_t padded_dim, T* est_distance ) { + // Each 4-dimensional codebook can contribute at most 255, so 1024 dimensions + // produce at most 255 * (1024 / 4) = 65280 in the uint16_t FastScan result. + constexpr size_t kSafeChunkDim = 1024; + ConstQGBatchDataMap cur_batch(batch_data, padded_dim); + + if (padded_dim <= kSafeChunkDim) { + std::array accu_res{}; + fastscan::accumulate( + cur_batch.bin_code(), q_obj.lut(), accu_res.data(), padded_dim + ); + + ConstRowMajorArrayMap ip_arr(accu_res.data(), 1, fastscan::kBatchSize); + ConstRowMajorArrayMap f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize); + ConstRowMajorArrayMap f_rescale_arr( + cur_batch.f_rescale(), 1, fastscan::kBatchSize + ); + RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); + + est_dist_arr = f_add_arr + q_obj.g_add() + + (f_rescale_arr * (q_obj.delta() * (ip_arr.template cast()) + + q_obj.sum_vl_lut() + q_obj.k1xsumq())); + return; + } + + std::array accu_values{}; std::array accu_res{}; + const auto* codes_ptr = cur_batch.bin_code(); + const auto* lut_ptr = q_obj.lut(); + size_t remaining_dim = padded_dim; - ConstQGBatchDataMap cur_batch(batch_data, padded_dim); + while (remaining_dim > kSafeChunkDim) { + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); + codes_ptr += kSafeChunkDim << 2; + lut_ptr += kSafeChunkDim << 2; + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_values[i] += accu_res[i]; + } + remaining_dim -= kSafeChunkDim; + } - fastscan::accumulate(cur_batch.bin_code(), q_obj.lut(), accu_res.data(), padded_dim); + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_values[i] += accu_res[i]; + } - ConstRowMajorArrayMap ip_arr(accu_res.data(), 1, fastscan::kBatchSize); + ConstRowMajorArrayMap ip_arr(accu_values.data(), 1, fastscan::kBatchSize); ConstRowMajorArrayMap f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize); ConstRowMajorArrayMap f_rescale_arr(cur_batch.f_rescale(), 1, fastscan::kBatchSize); diff --git a/include/rabitqlib/index/symqg/qg_builder.hpp b/include/rabitqlib/index/symqg/qg_builder.hpp index 8bedd43..8111904 100644 --- a/include/rabitqlib/index/symqg/qg_builder.hpp +++ b/include/rabitqlib/index/symqg/qg_builder.hpp @@ -75,8 +75,9 @@ class QGBuilder { std::vector centroid = compute_centroid(data, num_nodes_, dim_, num_threads_); - PID entry_point = - exact_nn(data, centroid.data(), num_nodes_, dim_, num_threads_, euclidean_sqr); + PID entry_point = exact_nn( + data, centroid.data(), num_nodes_, dim_, num_threads_, qg_.raw_dist_func_ + ); std::cout << "Setting entry_point to " << entry_point << '\n' << std::flush; diff --git a/tests/unit/rabitqlib/index/qg_test.cpp b/tests/unit/rabitqlib/index/qg_test.cpp index 02c2f69..8fb5d00 100644 --- a/tests/unit/rabitqlib/index/qg_test.cpp +++ b/tests/unit/rabitqlib/index/qg_test.cpp @@ -1,8 +1,12 @@ -#include "rabitqlib/index/symqg/qg.hpp" - #include +#include +#include +#include #include +#include + +#include "rabitqlib/index/symqg/qg_builder.hpp" namespace rabitqlib::symqg { namespace { @@ -26,5 +30,65 @@ TEST(QuantizedGraphLifecycleTest, DestroysConcreteRotatorThroughBasePointer) { EXPECT_EQ(graph.num_vertices(), 33U); } +TEST(QGBuilderMetricTest, UsesInnerProductDistanceToChooseEntryPoint) { + constexpr size_t kNumPoints = 33; + constexpr size_t kDim = 64; + constexpr size_t kDegree = 32; + std::vector data(kNumPoints * kDim, 0.0F); + data[0] = 100.0F; + + const std::vector centroid = compute_centroid(data.data(), kNumPoints, kDim, 1); + const PID expected = + exact_nn(data.data(), centroid.data(), kNumPoints, kDim, 1, dot_product_dis); + const PID euclidean_entry = + exact_nn(data.data(), centroid.data(), kNumPoints, kDim, 1, euclidean_sqr); + ASSERT_EQ(expected, 0U); + ASSERT_NE(euclidean_entry, expected); + + QuantizedGraph graph( + kNumPoints, kDim, kDegree, METRIC_IP, RotatorType::MatrixRotator + ); + QGBuilder builder(graph, kDegree, data.data(), 1); + + EXPECT_EQ(graph.entry_point(), expected); +} + +TEST(QGEstimatorTest, AccumulatesAcrossUint16SafeChunks) { + constexpr std::array kDimensions = {1024, 1088, 2048}; + + for (size_t padded_dim : kDimensions) { + SCOPED_TRACE(padded_dim); + + std::vector query(padded_dim, 1.0F); + BatchQuery q_obj(query.data(), padded_dim); + + std::vector batch_data(QGBatchDataMap::data_bytes(padded_dim)); + QGBatchDataMap batch_map(batch_data.data(), padded_dim); + std::fill( + batch_map.bin_code(), + batch_map.bin_code() + (padded_dim * fastscan::kBatchSize / 8), + uint8_t{0xff} + ); + std::fill_n(batch_map.f_add(), fastscan::kBatchSize, 0.0F); + std::fill_n(batch_map.f_rescale(), fastscan::kBatchSize, 1.0F); + + int64_t scalar_accumulator = 0; + for (size_t codebook = 0; codebook < padded_dim / 4; ++codebook) { + const uint8_t selected_value = q_obj.lut()[(codebook * 16) + 15]; + ASSERT_EQ(selected_value, uint8_t{0xff}); + scalar_accumulator += selected_value; + } + + const float expected = q_obj.delta() * static_cast(scalar_accumulator) + + q_obj.sum_vl_lut() + q_obj.k1xsumq(); + std::array estimated{}; + qg_batch_estdist(batch_data.data(), q_obj, padded_dim, estimated.data()); + + for (float distance : estimated) { + EXPECT_FLOAT_EQ(distance, expected); + } + } +} + } // namespace } // namespace rabitqlib::symqg