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
45 changes: 42 additions & 3 deletions include/rabitqlib/index/estimator.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -162,13 +162,52 @@ template <typename T, typename TA = uint16_t>
inline void qg_batch_estdist(
const char* batch_data, const BatchQuery<T>& 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<T> cur_batch(batch_data, padded_dim);

if (padded_dim <= kSafeChunkDim) {
std::array<TA, fastscan::kBatchSize> accu_res{};
fastscan::accumulate(
cur_batch.bin_code(), q_obj.lut(), accu_res.data(), padded_dim
);

ConstRowMajorArrayMap<TA> ip_arr(accu_res.data(), 1, fastscan::kBatchSize);
ConstRowMajorArrayMap<T> f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize);
ConstRowMajorArrayMap<T> f_rescale_arr(
cur_batch.f_rescale(), 1, fastscan::kBatchSize
);
RowMajorArrayMap<T> 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<T>()) +
q_obj.sum_vl_lut() + q_obj.k1xsumq()));
return;
}

std::array<int32_t, fastscan::kBatchSize> accu_values{};
std::array<TA, fastscan::kBatchSize> accu_res{};
const auto* codes_ptr = cur_batch.bin_code();
const auto* lut_ptr = q_obj.lut();
size_t remaining_dim = padded_dim;

ConstQGBatchDataMap<T> 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<TA> ip_arr(accu_res.data(), 1, fastscan::kBatchSize);
ConstRowMajorArrayMap<int32_t> ip_arr(accu_values.data(), 1, fastscan::kBatchSize);
ConstRowMajorArrayMap<T> f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize);
ConstRowMajorArrayMap<T> f_rescale_arr(cur_batch.f_rescale(), 1, fastscan::kBatchSize);

Expand Down
5 changes: 3 additions & 2 deletions include/rabitqlib/index/symqg/qg_builder.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -75,8 +75,9 @@ class QGBuilder {
std::vector<float> centroid =
compute_centroid(data, num_nodes_, dim_, num_threads_);

PID entry_point =
exact_nn(data, centroid.data(), num_nodes_, dim_, num_threads_, euclidean_sqr<float>);
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;

Expand Down
68 changes: 66 additions & 2 deletions tests/unit/rabitqlib/index/qg_test.cpp
Original file line number Diff line number Diff line change
@@ -1,8 +1,12 @@
#include "rabitqlib/index/symqg/qg.hpp"

#include <gtest/gtest.h>

#include <algorithm>
#include <array>
#include <cstdint>
#include <stdexcept>
#include <vector>

#include "rabitqlib/index/symqg/qg_builder.hpp"

namespace rabitqlib::symqg {
namespace {
Expand All @@ -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<float> data(kNumPoints * kDim, 0.0F);
data[0] = 100.0F;

const std::vector<float> centroid = compute_centroid(data.data(), kNumPoints, kDim, 1);
const PID expected =
exact_nn(data.data(), centroid.data(), kNumPoints, kDim, 1, dot_product_dis<float>);
const PID euclidean_entry =
exact_nn(data.data(), centroid.data(), kNumPoints, kDim, 1, euclidean_sqr<float>);
ASSERT_EQ(expected, 0U);
ASSERT_NE(euclidean_entry, expected);

QuantizedGraph<float> 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<size_t, 3> kDimensions = {1024, 1088, 2048};

for (size_t padded_dim : kDimensions) {
SCOPED_TRACE(padded_dim);

std::vector<float> query(padded_dim, 1.0F);
BatchQuery<float> q_obj(query.data(), padded_dim);

std::vector<char> batch_data(QGBatchDataMap<float>::data_bytes(padded_dim));
QGBatchDataMap<float> 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<float>(scalar_accumulator) +
q_obj.sum_vl_lut() + q_obj.k1xsumq();
std::array<float, fastscan::kBatchSize> 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
Loading