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
57 changes: 57 additions & 0 deletions include/rabitqlib/index/estimator.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -307,4 +307,61 @@ inline void split_single_estdist_direct(
low_dist = est_dist - (cur_bin.f_error() * g_error);
};

inline void xy_single_base_dist(
const char* base_data,
float (*base_ip_func)(const float*, const uint8_t*, size_t),
const SplitSingleQuery<float>& q_obj,
size_t padded_dim,
size_t base_bits,
float& ip_base,
float& est_dist,
float& low_dist,
float g_add = 0,
float g_error = 0
) {
ConstBaseDataMap<float> cur_base(base_data, padded_dim, base_bits);

ip_base = base_ip_func(q_obj.rotated_query(), cur_base.base_code(), padded_dim);

est_dist =
cur_base.f_add() + g_add + (cur_base.f_rescale() * (ip_base + q_obj.kbase_sumq()));

low_dist = est_dist - (cur_base.f_error() * g_error);
}

inline void xy_single_full_dist(
const char* base_data,
const char* ex_data,
float (*base_ip_func)(const float*, const uint8_t*, size_t),
float (*ex_ip_func)(const float*, const uint8_t*, size_t),
const SplitSingleQuery<float>& q_obj,
size_t padded_dim,
size_t base_bits,
size_t ex_bits,
float& est_dist,
float& low_dist,
float& ip_base,
float g_add = 0,
float g_error = 0
) {
ConstBaseDataMap<float> cur_base(base_data, padded_dim, base_bits);
ip_base = base_ip_func(q_obj.rotated_query(), cur_base.base_code(), padded_dim);

if (ex_bits == 0) {
est_dist = cur_base.f_add() + g_add +
(cur_base.f_rescale() * (ip_base + q_obj.kbase_sumq()));
low_dist = est_dist - (cur_base.f_error() * g_error);
return;
}

ConstExDataMap<float> cur_ex(ex_data, padded_dim, ex_bits);

float ip = (static_cast<float>(1 << ex_bits) * ip_base) +
ex_ip_func(q_obj.rotated_query(), cur_ex.ex_code(), padded_dim);

est_dist = cur_ex.f_add_ex() + g_add + (cur_ex.f_rescale_ex() * (ip + q_obj.kbxsumq()));

low_dist = est_dist - (cur_base.f_error() * g_error / static_cast<float>(1 << ex_bits));
}

} // namespace rabitqlib
22 changes: 17 additions & 5 deletions include/rabitqlib/index/query.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <cmath>
#include <cstdint>
#include <numeric>
#include <stdexcept>
#include <utility>

#include "rabitqlib/defines.hpp"
Expand Down Expand Up @@ -116,6 +117,7 @@ class SplitSingleQuery {
std::vector<uint64_t> QueryBin_;
T G_add_;
T G_k1xSumq_;
T G_kbaseSumq_;
T G_kbxSumq_;
T G_error_;
T delta_;
Expand All @@ -129,15 +131,25 @@ class SplitSingleQuery {
size_t padded_dim,
size_t ex_bits,
quant::RabitqConfig config,
size_t metric_type = METRIC_L2
size_t metric_type = METRIC_L2,
size_t base_bits = 1
)
: rotated_query_(rotated_query), QueryBin_(padded_dim * kNumBits / 64, 0) {
if (base_bits < 1 || base_bits > 8 || ex_bits > 8 || base_bits + ex_bits > 9) {
throw std::invalid_argument(
"SplitSingleQuery requires base_bits in [1, 8], ex_bits in [0, 8], "
"and base_bits + ex_bits <= 9"
);
}

float c_1 = -static_cast<float>((1 << 1) - 1) / 2.F;
float c_b = -static_cast<float>((1 << (ex_bits + 1)) - 1) / 2.F;
float c_base = -static_cast<float>((1U << base_bits) - 1) / 2.F;
float c_b = -static_cast<float>((1U << (base_bits + ex_bits)) - 1) / 2.F;
T sumq =
std::accumulate(rotated_query, rotated_query + padded_dim, static_cast<T>(0));

G_k1xSumq_ = sumq * c_1;
G_kbaseSumq_ = sumq * c_base;
G_kbxSumq_ = sumq * c_b;

metric_type_ = (metric_type == METRIC_IP) ? METRIC_IP : METRIC_L2;
Expand All @@ -153,9 +165,6 @@ class SplitSingleQuery {
rabitqlib::new_transpose_bin_512(
quant_query.data(), QueryBin_.data(), padded_dim, kNumBits
);

// new_transpose_bin_512 already stores the query in the bit-plane/chunk
// layout consumed by warmup_ip_x0_q_512.
}

[[nodiscard]] size_t num_bits() const { return kNumBits; }
Expand All @@ -170,6 +179,9 @@ class SplitSingleQuery {

[[nodiscard]] T k1xsumq() const { return G_k1xSumq_; }

// Equals k1xsumq() when base_bits == 1.
[[nodiscard]] T kbase_sumq() const { return G_kbaseSumq_; }

[[nodiscard]] T kbxsumq() const { return G_kbxSumq_; }

[[nodiscard]] T g_add() const { return G_add_; }
Expand Down
47 changes: 47 additions & 0 deletions include/rabitqlib/quantization/data_layout.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,53 @@ struct ConstExDataMap {
const T& f_recale_ex_;
};

template <typename T>
struct BaseDataMap {
public:
explicit BaseDataMap(char* data, size_t padded_dim, size_t base_bits)
: base_code_(reinterpret_cast<uint8_t*>(data))
, f_add_(*reinterpret_cast<T*>(data + (padded_dim * base_bits / 8)))
, f_rescale_(*(reinterpret_cast<T*>(data + (padded_dim * base_bits / 8)) + 1))
, f_error_(*(reinterpret_cast<T*>(data + (padded_dim * base_bits / 8)) + 2)) {}

static size_t data_bytes(size_t padded_dim, size_t base_bits) {
return (padded_dim * base_bits / 8) + (sizeof(T) * 3);
}

[[nodiscard]] uint8_t* base_code() { return base_code_; }
[[nodiscard]] T& f_add() { return f_add_; }
[[nodiscard]] T& f_rescale() { return f_rescale_; }
[[nodiscard]] T& f_error() { return f_error_; }

private:
uint8_t* base_code_;
T& f_add_;
T& f_rescale_;
T& f_error_;
};

template <typename T>
struct ConstBaseDataMap {
public:
explicit ConstBaseDataMap(const char* data, size_t padded_dim, size_t base_bits)
: base_code_(reinterpret_cast<const uint8_t*>(data))
, f_add_(*reinterpret_cast<const T*>(data + (padded_dim * base_bits / 8)))
, f_rescale_(*(reinterpret_cast<const T*>(data + (padded_dim * base_bits / 8)) + 1))
, f_error_(*(reinterpret_cast<const T*>(data + (padded_dim * base_bits / 8)) + 2)) {
}

[[nodiscard]] const uint8_t* base_code() const { return base_code_; }
[[nodiscard]] const T& f_add() const { return f_add_; }
[[nodiscard]] const T& f_rescale() const { return f_rescale_; }
[[nodiscard]] const T& f_error() const { return f_error_; }

private:
const uint8_t* base_code_;
const T& f_add_;
const T& f_rescale_;
const T& f_error_;
};

template <typename T>
struct BinDataMap {
public:
Expand Down
24 changes: 24 additions & 0 deletions include/rabitqlib/quantization/rabitq.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -265,6 +265,30 @@ inline void quantize_split_single(
}
}

inline void quantize_xy_single(
const float* data,
const float* centroid,
size_t padded_dim,
size_t base_bits,
size_t ex_bits,
char* base_data,
char* ex_data,
MetricType metric_type = METRIC_L2,
RabitqConfig config = RabitqConfig()
) {
rabitq_impl::xy_bits::split_code_with_factor<float>(
data,
centroid,
padded_dim,
base_bits,
ex_bits,
base_data,
ex_data,
metric_type,
config.t_const
);
}

template <typename T, typename TP>
inline void quantize_full_single(
const T* data,
Expand Down
Loading
Loading