diff --git a/README.md b/README.md index 7f5626e..d53963f 100644 --- a/README.md +++ b/README.md @@ -135,14 +135,6 @@ RaBitQ is developed by the University, Singapore. A GPU implementation is also available in [cuvs_rabitq](https://github.com/Stardust-SJF/cuvs_rabitq/tree/cuvs_ivf_rabitq). -## Accuracy at a glance - -![RaBitQ estimation error benchmark across MSong, YouTube, OpenAI embeddings, Word2Vec, and GIST](docs/docs/assets/img/acc_bench.png) - -*Average and maximum relative estimation error across six datasets; lower is -better. Results from the -[SIGMOD camera-ready paper](https://doi.org/10.1145/3725413).* - ## RaBitQ across the vector-search ecosystem The projects below illustrate adoption of RaBitQ techniques across vector diff --git a/docs/docs/assets/img/acc_bench.png b/docs/docs/assets/img/acc_bench.png deleted file mode 100644 index 33180bb..0000000 Binary files a/docs/docs/assets/img/acc_bench.png and /dev/null differ diff --git a/docs/docs/index.md b/docs/docs/index.md index cf428c2..504bbe7 100644 --- a/docs/docs/index.md +++ b/docs/docs/index.md @@ -33,14 +33,6 @@ indexes backed by optimized AVX2 and AVX-512 kernels. -## Accuracy at a glance - -![RaBitQ estimation error benchmark across MSong, YouTube, OpenAI embeddings, Word2Vec, and GIST](assets/img/acc_bench.png) - -*Average and maximum relative estimation error across six datasets; lower is -better. Results from the -[SIGMOD camera-ready paper](https://doi.org/10.1145/3725413).* - ## Start with Python Install the latest release from PyPI: diff --git a/include/rabitqlib/index/estimator.hpp b/include/rabitqlib/index/estimator.hpp index 3f75c05..b6b00ab 100644 --- a/include/rabitqlib/index/estimator.hpp +++ b/include/rabitqlib/index/estimator.hpp @@ -77,11 +77,19 @@ inline void split_batch_estdist( } } - ConstRowMajorArrayMap f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize); + std::array f_add_values; + std::array f_rescale_values; + std::array f_error_values; + cur_batch.f_add().copy_to(f_add_values.data(), f_add_values.size()); + cur_batch.f_rescale().copy_to(f_rescale_values.data(), f_rescale_values.size()); + cur_batch.f_error().copy_to(f_error_values.data(), f_error_values.size()); + ConstRowMajorArrayMap f_add_arr(f_add_values.data(), 1, fastscan::kBatchSize); ConstRowMajorArrayMap f_rescale_arr( - cur_batch.f_rescale(), 1, fastscan::kBatchSize + f_rescale_values.data(), 1, fastscan::kBatchSize + ); + ConstRowMajorArrayMap f_error_arr( + f_error_values.data(), 1, fastscan::kBatchSize ); - ConstRowMajorArrayMap f_error_arr(cur_batch.f_error(), 1, fastscan::kBatchSize); RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); RowMajorArrayMap ip_x0_qr_arr(ip_x0_qr, 1, fastscan::kBatchSize); @@ -145,7 +153,13 @@ inline void split_single_fulldist_direct( ConstExDataMap cur_ex(ex_data, padded_dim, ex_bits); // [TODO: optimize this function] - ip_x0_qr = Kernel::mask_ip_x0_q(q_obj.rotated_query(), cur_bin.bin_code(), padded_dim); + // Direct HNSW kernels retain their legacy word-typed ABI. Generic kernels use the + // byte-oriented overload and do not manufacture a typed view of packed storage. + ip_x0_qr = Kernel::mask_ip_x0_q( + q_obj.rotated_query(), + reinterpret_cast(cur_bin.bin_code()), + padded_dim + ); est_dist = cur_ex.f_add_ex() + g_add + @@ -176,9 +190,13 @@ inline void qg_batch_estdist( ); ConstRowMajorArrayMap ip_arr(accu_res.data(), 1, fastscan::kBatchSize); - ConstRowMajorArrayMap f_add_arr(cur_batch.f_add(), 1, fastscan::kBatchSize); + std::array f_add_values; + std::array f_rescale_values; + cur_batch.f_add().copy_to(f_add_values.data(), f_add_values.size()); + cur_batch.f_rescale().copy_to(f_rescale_values.data(), f_rescale_values.size()); + ConstRowMajorArrayMap f_add_arr(f_add_values.data(), 1, fastscan::kBatchSize); ConstRowMajorArrayMap f_rescale_arr( - cur_batch.f_rescale(), 1, fastscan::kBatchSize + f_rescale_values.data(), 1, fastscan::kBatchSize ); RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); @@ -210,8 +228,14 @@ inline void qg_batch_estdist( } 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); + std::array f_add_values; + std::array f_rescale_values; + cur_batch.f_add().copy_to(f_add_values.data(), f_add_values.size()); + cur_batch.f_rescale().copy_to(f_rescale_values.data(), f_rescale_values.size()); + ConstRowMajorArrayMap f_add_arr(f_add_values.data(), 1, fastscan::kBatchSize); + ConstRowMajorArrayMap f_rescale_arr( + f_rescale_values.data(), 1, fastscan::kBatchSize + ); RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); @@ -294,8 +318,9 @@ inline void split_single_estdist_direct( ) { ConstBinDataMap cur_bin(bin_data, padded_dim); + // Direct HNSW kernels retain their legacy word-typed ABI; see the note above. ip_x0_qr = Kernel::warmup_ip_x0_q_512( - cur_bin.bin_code(), + reinterpret_cast(cur_bin.bin_code()), q_obj.query_bin(), q_obj.delta(), q_obj.vl(), diff --git a/include/rabitqlib/index/ivf/initializer.hpp b/include/rabitqlib/index/ivf/initializer.hpp index 8890692..d5458b3 100644 --- a/include/rabitqlib/index/ivf/initializer.hpp +++ b/include/rabitqlib/index/ivf/initializer.hpp @@ -21,9 +21,13 @@ namespace rabitqlib::ivf { template inline void parallel_for(size_t start, size_t end, size_t numThreads, Function fn) { - if (numThreads <= 0) { - numThreads = std::thread::hardware_concurrency(); + if (start >= end) { + return; } + if (numThreads == 0) { + numThreads = std::max(1, std::thread::hardware_concurrency()); + } + numThreads = std::min(numThreads, end - start); if (numThreads == 1) { for (size_t id = start; id < end; id++) { @@ -38,36 +42,49 @@ inline void parallel_for(size_t start, size_t end, size_t numThreads, Function f std::exception_ptr last_exception = nullptr; std::mutex last_except_mutex; - threads.reserve(numThreads); - for (size_t thread_id = 0; thread_id < numThreads; ++thread_id) { - threads.push_back(std::thread([&, thread_id] { - while (true) { - size_t id = current.fetch_add(1); - - if (id >= end) { - break; - } - - try { - fn(id, thread_id); - } catch (...) { - std::unique_lock last_except_lock(last_except_mutex); - last_exception = std::current_exception(); - /* - * This will work even when current is the largest value that - * size_t can fit, because fetch_add returns the previous value - * before the increment (what will result in overflow - * and produce 0 instead of current + 1). - */ - current = end; - break; - } + const auto join_threads = [&threads] { + for (auto& thread : threads) { + if (thread.joinable()) { + thread.join(); } - })); - } - for (auto& thread : threads) { - thread.join(); + } + }; + + try { + threads.reserve(numThreads); + for (size_t thread_id = 0; thread_id < numThreads; ++thread_id) { + threads.emplace_back([&, thread_id] { + while (true) { + size_t id = current.fetch_add(1); + + if (id >= end) { + break; + } + + try { + fn(id, thread_id); + } catch (...) { + std::unique_lock last_except_lock(last_except_mutex + ); + last_exception = std::current_exception(); + /* + * This will work even when current is the largest value that + * size_t can fit, because fetch_add returns the previous value + * before the increment (what will result in overflow + * and produce 0 instead of current + 1). + */ + current = end; + break; + } + } + }); + } + } catch (...) { + current = end; + join_threads(); + throw; } + join_threads(); if (last_exception) { std::rethrow_exception(last_exception); } @@ -97,10 +114,15 @@ inline Initializer::~Initializer() {} class FlatInitializer : public Initializer { private: std::vector centroids_; + MetricType metric_type_; public: - explicit FlatInitializer(size_t d, size_t k) - : Initializer(d, k), centroids_(num_cluster_ * dim_) {} + explicit FlatInitializer( + size_t d, size_t k, MetricType metric_type = MetricType::METRIC_L2 + ) + : Initializer(d, k), centroids_(num_cluster_ * dim_), metric_type_(metric_type) { + validate_metric_type(metric_type_); + } ~FlatInitializer() override = default; @@ -118,7 +140,10 @@ class FlatInitializer : public Initializer { std::vector> centroid_dist(this->num_cluster_); for (PID i = 0; i < num_cluster_; ++i) { centroid_dist[i].id = i; - centroid_dist[i].distance = std::sqrt(euclidean_sqr(query, centroid(i), dim_)); + centroid_dist[i].distance = + metric_type_ == METRIC_IP + ? dot_product_dis(query, centroid(i), dim_) + : std::sqrt(euclidean_sqr(query, centroid(i), dim_)); } std::partial_sort( centroid_dist.begin(), diff --git a/include/rabitqlib/index/ivf/ivf.hpp b/include/rabitqlib/index/ivf/ivf.hpp index 1cf6d70..80e814a 100644 --- a/include/rabitqlib/index/ivf/ivf.hpp +++ b/include/rabitqlib/index/ivf/ivf.hpp @@ -5,7 +5,6 @@ #include #include #include -#include #include #include #include @@ -32,16 +31,23 @@ namespace rabitqlib::ivf { class IVF { private: - std::unique_ptr initer_; // initializer for candidate clusters - char* batch_data_ = nullptr; // 1-bit code and factors - char* ex_data_ = nullptr; // extra-bit codes or original float32 vectors - PID* ids_ = nullptr; // PID of vectors (organized by clusters) - size_t num_ = 0; // num of data points - size_t dim_ = 0; // dimension of data points - size_t padded_dim_ = 0; // dimension after padding - size_t num_cluster_ = 0; // num of centroids (clusters) - bool raw_reranking_ = false; // raw vectors replace extra-bit codes - size_t ex_bits_ = 0; // total bits = ex_bits_ + 1 + using ByteStorage = + std::vector>; + using FloatStorage = std::vector>; + using IdStorage = std::vector>; + + std::unique_ptr initer_; // initializer for candidate clusters + ByteStorage batch_storage_; // 1-bit code and factors + ByteStorage ex_storage_; // extra-bit codes and packed factors + FloatStorage raw_storage_; // original vectors for raw reranking + IdStorage id_storage_; // point IDs organized by cluster + size_t num_ = 0; // num of data points + size_t dim_ = 0; // dimension of data points + size_t padded_dim_ = 0; // dimension after padding + size_t num_cluster_ = 0; // num of centroids (clusters) + bool raw_reranking_ = false; // raw vectors replace extra-bit codes + bool ready_ = false; // construction or loading completed + size_t ex_bits_ = 0; // total bits = ex_bits_ + 1 RotatorType type_ = RotatorType::FhtKacRotator; // type of rotator std::unique_ptr> rotator_; // Data Rotator std::vector cluster_lst_; // List of clusters in ivf @@ -73,15 +79,59 @@ class IVF { [[nodiscard]] size_t ex_data_bytes() const { return rerank_vector_bytes() * num_; } + [[nodiscard]] char* batch_data() { + return reinterpret_cast(batch_storage_.data()); + } + + [[nodiscard]] const char* batch_data() const { + return reinterpret_cast(batch_storage_.data()); + } + + [[nodiscard]] char* ex_data() { + return raw_reranking_ ? reinterpret_cast(raw_storage_.data()) + : reinterpret_cast(ex_storage_.data()); + } + + [[nodiscard]] const char* ex_data() const { + return raw_reranking_ ? reinterpret_cast(raw_storage_.data()) + : reinterpret_cast(ex_storage_.data()); + } + + [[nodiscard]] PID* ids() { return id_storage_.data(); } + + [[nodiscard]] const PID* ids() const { return id_storage_.data(); } + void allocate_memory(const std::vector&); void init_clusters(const std::vector&); void free_memory() { initer_.reset(); - std::free(std::exchange(batch_data_, nullptr)); - std::free(std::exchange(ex_data_, nullptr)); - std::free(std::exchange(ids_, nullptr)); + ByteStorage().swap(batch_storage_); + ByteStorage().swap(ex_storage_); + FloatStorage().swap(raw_storage_); + IdStorage().swap(id_storage_); + } + + void swap(IVF& other) noexcept { + using std::swap; + swap(initer_, other.initer_); + swap(batch_storage_, other.batch_storage_); + swap(ex_storage_, other.ex_storage_); + swap(raw_storage_, other.raw_storage_); + swap(id_storage_, other.id_storage_); + swap(num_, other.num_); + swap(dim_, other.dim_); + swap(padded_dim_, other.padded_dim_); + swap(num_cluster_, other.num_cluster_); + swap(raw_reranking_, other.raw_reranking_); + swap(ready_, other.ready_); + swap(ex_bits_, other.ex_bits_); + swap(type_, other.type_); + swap(rotator_, other.rotator_); + swap(cluster_lst_, other.cluster_lst_); + swap(metric_type_, other.metric_type_); + swap(ip_func_, other.ip_func_); } void @@ -149,13 +199,27 @@ inline IVF::IVF( if ((bits < 1 || bits > 9) && bits != 32) { throw std::invalid_argument("IVF bits must be in [1, 9] or 32 for raw reranking"); }; - if (raw_reranking_ && (n == 0 || n > buffer::kSearchBufferMaxPointCount || dim == 0 || - dim > std::numeric_limits::max() / sizeof(float) / n || - dim > (std::numeric_limits::max() / 32) - 64)) { - throw std::invalid_argument("Invalid raw IVF vector count or dimension"); + if (n == 0 || n > buffer::kSearchBufferMaxPointCount) { + throw std::invalid_argument("IVF point count exceeds the supported ID range"); } - rotator_.reset(choose_rotator(dim, type, round_up_to_multiple(dim_, 64))); - padded_dim_ = rotator_->size(); + if (cluster_num == 0 || cluster_num > buffer::kSearchBufferMaxPointCount) { + throw std::invalid_argument("IVF cluster count exceeds the supported ID range"); + } + if (dim == 0 || dim > (std::numeric_limits::max() / 32) - 64) { + throw std::invalid_argument("IVF dimension is invalid or too large"); + } + padded_dim_ = round_up_to_multiple(dim_, 64); + const size_t max_size = std::numeric_limits::max(); + const size_t batch_bytes = BatchDataMap::data_bytes(padded_dim_); + const size_t rerank_bytes = rerank_vector_bytes(); + if (n > max_size / sizeof(PID) || n > max_size / sizeof(float) / padded_dim_ || + n > max_size / batch_bytes || (rerank_bytes != 0 && n > max_size / rerank_bytes) || + cluster_num > max_size / sizeof(float) / padded_dim_ || + (type == RotatorType::MatrixRotator && dim > max_size / sizeof(float) / padded_dim_ + )) { + throw std::invalid_argument("IVF configuration exceeds addressable storage"); + } + rotator_.reset(choose_rotator(dim, type, padded_dim_)); /* check size */ assert(padded_dim_ % 64 == 0); assert(padded_dim_ >= dim_); @@ -177,6 +241,13 @@ inline void IVF::construct( bool faster = false, size_t num_threads = std::numeric_limits::max() ) { + if (num_ == 0 || dim_ == 0 || num_cluster_ == 0 || rotator_ == nullptr) { + throw std::logic_error("IVF must be configured before construction"); + } + if (data == nullptr || centroids == nullptr || cluster_ids == nullptr) { + throw std::invalid_argument("IVF construction inputs must not be null"); + } + // get id list for each cluster std::vector counts(num_cluster_, 0); std::vector> id_lists(num_cluster_); @@ -213,23 +284,29 @@ inline void IVF::construct( } this->initer_->add_vectors(rotated_centroids.data(), num_threads); + ready_ = true; } inline void IVF::allocate_memory(const std::vector& cluster_sizes) { + ready_ = false; free_memory(); cluster_lst_.clear(); if (num_cluster_ < 20000UL) { - this->initer_ = std::make_unique(padded_dim_, num_cluster_); + this->initer_ = + std::make_unique(padded_dim_, num_cluster_, metric_type_); } else { this->initer_ = std::make_unique(padded_dim_, num_cluster_); } - this->batch_data_ = - memory::align_allocate<64, char, true>(batch_data_bytes(cluster_sizes)); + batch_storage_ = ByteStorage(batch_data_bytes(cluster_sizes)); if (rerank_vector_bytes() > 0) { - this->ex_data_ = memory::align_allocate<64, char, true>(ex_data_bytes()); + if (raw_reranking_) { + raw_storage_ = FloatStorage(num_ * dim_); + } else { + ex_storage_ = ByteStorage(ex_data_bytes()); + } } - this->ids_ = memory::align_allocate<64, PID, true>(ids_bytes()); + id_storage_ = IdStorage(num_); this->ip_func_ = select_excode_ipfunc(ex_bits_); } @@ -247,13 +324,13 @@ inline void IVF::init_clusters(const std::vector& cluster_sizes) { size_t num_batches = div_round_up(num, fastscan::kBatchSize); char* current_batch_data = - batch_data_ + (BatchDataMap::data_bytes(padded_dim_) * added_batches); + batch_data() + (BatchDataMap::data_bytes(padded_dim_) * added_batches); char* current_ex_data = rerank_vector_bytes() > 0 - ? ex_data_ + (added_vectors * rerank_vector_bytes()) + ? ex_data() + (added_vectors * rerank_vector_bytes()) : nullptr; - PID* ids = ids_ + added_vectors; + PID* cluster_ids = ids() + added_vectors; - Cluster cur_cluster(num, current_batch_data, current_ex_data, ids); + Cluster cur_cluster(num, current_batch_data, current_ex_data, cluster_ids); this->cluster_lst_.push_back(std::move(cur_cluster)); added_vectors += num; @@ -317,9 +394,14 @@ inline void IVF::quantize_cluster( } inline void IVF::save(const char* filename) const { - if (cluster_lst_.size() == 0) { + if (!ready_ || initer_ == nullptr || rotator_ == nullptr || + cluster_lst_.size() != num_cluster_ || batch_storage_.empty() || + id_storage_.size() != num_) { throw std::logic_error("Cannot save an unconstructed IVF index"); } + if (filename == nullptr || filename[0] == '\0') { + throw std::invalid_argument("IVF save filename must not be empty"); + } std::ofstream output(filename, std::ios::binary); output.exceptions(std::ios::failbit | std::ios::badbit); @@ -356,32 +438,33 @@ inline void IVF::save(const char* filename) const { /* Save data */ this->initer_->save(output, filename); - output.write( - reinterpret_cast(batch_data_), - static_cast(batch_data_bytes(cluster_sizes)) - ); - output.write( - reinterpret_cast(ex_data_), static_cast(ex_data_bytes()) - ); - output.write(reinterpret_cast(ids_), static_cast(ids_bytes())); + output.write(batch_data(), static_cast(batch_data_bytes(cluster_sizes))); + if (ex_data_bytes() != 0) { + output.write(ex_data(), static_cast(ex_data_bytes())); + } + output.write(reinterpret_cast(ids()), static_cast(ids_bytes())); output.close(); } inline void IVF::load(const char* filename) { + if (filename == nullptr || filename[0] == '\0') { + throw std::invalid_argument("IVF load filename must not be empty"); + } std::ifstream input(filename, std::ios::binary); if (!input.is_open()) { throw std::runtime_error("Cannot open IVF index file"); } + IVF loaded; input.exceptions(std::ios::failbit | std::ios::badbit); input.seekg(0, std::ios::end); const auto file_bytes = static_cast(input.tellg()); input.seekg(0); uint64_t magic = 0; input.read(reinterpret_cast(&magic), sizeof(magic)); - raw_reranking_ = magic == kRawFormatMagic; - if (raw_reranking_) { + loaded.raw_reranking_ = magic == kRawFormatMagic; + if (loaded.raw_reranking_) { uint32_t version = 0; input.read(reinterpret_cast(&version), sizeof(version)); if (version != kRawFormatVersion) { @@ -392,21 +475,21 @@ inline void IVF::load(const char* filename) { } /* Load meta data */ - input.read(reinterpret_cast(&this->num_), sizeof(size_t)); - input.read(reinterpret_cast(&this->dim_), sizeof(size_t)); - input.read(reinterpret_cast(&this->num_cluster_), sizeof(size_t)); - input.read(reinterpret_cast(&this->ex_bits_), sizeof(size_t)); - input.read(reinterpret_cast(&type_), sizeof(type_)); - input.read(reinterpret_cast(&metric_type_), sizeof(metric_type_)); - validate_metric_type(metric_type_); - if ((raw_reranking_ && num_ == 0) || num_ > buffer::kSearchBufferMaxPointCount || - dim_ == 0 || num_cluster_ == 0 || - num_cluster_ > buffer::kSearchBufferMaxPointCount || ex_bits_ > 8 || - (raw_reranking_ && ex_bits_ != 0) || - dim_ > (std::numeric_limits::max() / 32) - 64) { + input.read(reinterpret_cast(&loaded.num_), sizeof(size_t)); + input.read(reinterpret_cast(&loaded.dim_), sizeof(size_t)); + input.read(reinterpret_cast(&loaded.num_cluster_), sizeof(size_t)); + input.read(reinterpret_cast(&loaded.ex_bits_), sizeof(size_t)); + input.read(reinterpret_cast(&loaded.type_), sizeof(loaded.type_)); + input.read(reinterpret_cast(&loaded.metric_type_), sizeof(loaded.metric_type_)); + validate_metric_type(loaded.metric_type_); + if (loaded.num_ == 0 || loaded.num_ > buffer::kSearchBufferMaxPointCount || + loaded.dim_ == 0 || loaded.num_cluster_ == 0 || + loaded.num_cluster_ > buffer::kSearchBufferMaxPointCount || loaded.ex_bits_ > 8 || + (loaded.raw_reranking_ && loaded.ex_bits_ != 0) || + loaded.dim_ > (std::numeric_limits::max() / 32) - 64) { throw std::runtime_error("Invalid IVF index metadata"); } - padded_dim_ = round_up_to_multiple(dim_, 64); + loaded.padded_dim_ = round_up_to_multiple(loaded.dim_, 64); // Bound every payload by the actual file before allocating or multiplying sizes. size_t remaining = file_bytes - static_cast(input.tellg()); const auto consume = [&remaining](size_t count, size_t bytes) { @@ -415,57 +498,73 @@ inline void IVF::load(const char* filename) { } remaining -= count * bytes; }; - consume(num_cluster_, sizeof(size_t)); - consume(num_, sizeof(PID)); - consume(num_, rerank_vector_bytes()); - if (type_ == RotatorType::MatrixRotator) { - consume(dim_, sizeof(float) * padded_dim_); - } else if (type_ == RotatorType::FhtKacRotator) { - consume(1, padded_dim_ / 2); + consume(loaded.num_cluster_, sizeof(size_t)); + consume(loaded.num_, sizeof(PID)); + consume(loaded.num_, loaded.rerank_vector_bytes()); + if (loaded.type_ == RotatorType::MatrixRotator) { + consume(loaded.dim_, sizeof(float) * loaded.padded_dim_); + } else if (loaded.type_ == RotatorType::FhtKacRotator) { + consume(1, loaded.padded_dim_ / 2); } else { throw std::runtime_error("Invalid IVF rotator type"); } - if (num_cluster_ < 20000UL) { - consume(num_cluster_, sizeof(float) * padded_dim_); + if (loaded.num_cluster_ < 20000UL) { + consume(loaded.num_cluster_, sizeof(float) * loaded.padded_dim_); } /* Load number of vectors of each cluster */ - std::vector cluster_sizes(num_cluster_, 0); + std::vector cluster_sizes(loaded.num_cluster_, 0); input.read( reinterpret_cast(cluster_sizes.data()), - static_cast(sizeof(size_t) * num_cluster_) + static_cast(sizeof(size_t) * loaded.num_cluster_) ); size_t total = 0; for (size_t size : cluster_sizes) { - if (size > num_ - total) { + if (size > loaded.num_ - total) { throw std::runtime_error("Invalid cluster counts in IVF index file"); } total += size; consume( div_round_up(size, fastscan::kBatchSize), - BatchDataMap::data_bytes(padded_dim_) + BatchDataMap::data_bytes(loaded.padded_dim_) ); } - if (total != num_) { + if (total != loaded.num_) { throw std::runtime_error("Invalid cluster counts in IVF index file"); } - rotator_.reset(choose_rotator(dim_, type_, padded_dim_)); + loaded.rotator_.reset( + choose_rotator(loaded.dim_, loaded.type_, loaded.padded_dim_) + ); /* Load rotator */ - this->rotator_->load(input); + loaded.rotator_->load(input); /* Load data */ - allocate_memory(cluster_sizes); - this->initer_->load(input, filename); - input.read(batch_data_, static_cast(batch_data_bytes(cluster_sizes))); - input.read(ex_data_, static_cast(ex_data_bytes())); - input.read(reinterpret_cast(ids_), static_cast(ids_bytes())); + loaded.allocate_memory(cluster_sizes); + loaded.initer_->load(input, filename); + input.read( + loaded.batch_data(), static_cast(loaded.batch_data_bytes(cluster_sizes)) + ); + if (loaded.ex_data_bytes() != 0) { + input.read(loaded.ex_data(), static_cast(loaded.ex_data_bytes())); + } + input.read( + reinterpret_cast(loaded.ids()), static_cast(loaded.ids_bytes()) + ); + for (size_t i = 0; i < loaded.num_; ++i) { + if (loaded.ids()[i] >= loaded.num_ || + loaded.ids()[i] >= buffer::kSearchBufferCheckedMask) { + throw std::runtime_error("Invalid point ID in IVF index file"); + } + } /* Init each cluster */ - init_clusters(cluster_sizes); + loaded.init_clusters(cluster_sizes); + loaded.ready_ = true; input.close(); + swap(loaded); } inline void IVF::search( @@ -492,6 +591,24 @@ inline void IVF::search( float* __restrict__ dists, bool use_hacc ) const { + if (!ready_ || initer_ == nullptr || rotator_ == nullptr || + cluster_lst_.size() != num_cluster_ || batch_storage_.empty() || + id_storage_.size() != num_) { + throw std::logic_error("IVF index must be constructed or loaded before search"); + } + if (query == nullptr) { + throw std::invalid_argument("IVF search query must not be null"); + } + if (results == nullptr) { + throw std::invalid_argument("IVF search results must not be null"); + } + if (k == 0 || k > num_ || k > buffer::kSearchBufferMaxPointCount) { + throw std::invalid_argument("IVF search k must be between 1 and the point count"); + } + if (nprobe == 0) { + throw std::invalid_argument("IVF search nprobe must be positive"); + } + nprobe = std::min(nprobe, num_cluster_); // corner case std::vector rotated_query(padded_dim_); this->rotator_->rotate(query, rotated_query.data()); @@ -514,10 +631,13 @@ inline void IVF::search( if (metric_type_ == METRIC_L2) { q_obj.set_g_add(dist); } else if (metric_type_ == METRIC_IP) { + const float residual_norm = std::sqrt( + euclidean_sqr(rotated_query.data(), initer_->centroid(cid), padded_dim_) + ); auto g_add_ip = dot_product( rotated_query.data(), initer_->centroid(cid), padded_dim_ ); - q_obj.set_g_add(dist, g_add_ip); + q_obj.set_g_add(residual_norm, g_add_ip); } else { // unsupported throw std::invalid_argument("Quantization only supports L2 and IP metrics"); @@ -526,11 +646,14 @@ inline void IVF::search( search_cluster(cur_cluster, q_obj, knns, use_hacc, query); } + const size_t found = knns.size(); if (dists != nullptr) { knns.copy_results(results, dists); + std::fill(dists + found, dists + k, std::numeric_limits::infinity()); } else { knns.copy_results(results); } + std::fill(results + found, results + k, kPidMax); } inline void IVF::search_cluster( diff --git a/include/rabitqlib/index/lut.hpp b/include/rabitqlib/index/lut.hpp index 7b135c6..242d941 100644 --- a/include/rabitqlib/index/lut.hpp +++ b/include/rabitqlib/index/lut.hpp @@ -2,8 +2,9 @@ #include #include +#include +#include #include -#include #include #include "rabitqlib/fastscan/fastscan.hpp" @@ -19,16 +20,37 @@ class Lut { static_assert(std::is_floating_point_v, "T must be a floating-point type in Lut"); private: + [[nodiscard]] static size_t checked_table_length(size_t padded_dim) { + if (padded_dim == 0 || padded_dim % 16 != 0) { + throw std::invalid_argument( + "FastScan dimension must be a positive multiple of 16" + ); + } + if (padded_dim > std::numeric_limits::max() / 4) { + throw std::length_error("FastScan lookup table is too large"); + } + return padded_dim * 4; + } + + [[nodiscard]] static size_t checked_storage_length(size_t padded_dim, bool use_hacc) { + const size_t table_length = checked_table_length(padded_dim); + const size_t copies = use_hacc ? 2 : 1; + if (table_length > std::numeric_limits::max() / copies) { + throw std::length_error("FastScan lookup table is too large"); + } + return table_length * copies; + } + size_t table_length_ = 0; std::vector lut_; - T delta_; - T sum_vl_lut_; + T delta_ = 0; + T sum_vl_lut_ = 0; public: explicit Lut() = default; explicit Lut(const T* rotated_query, size_t padded_dim, bool use_hacc = false) - : table_length_(padded_dim << 2) - , lut_(table_length_ * (static_cast(use_hacc) + 1)) { + : table_length_(checked_table_length(padded_dim)) + , lut_(checked_storage_length(padded_dim, use_hacc)) { // quantize float lut std::vector lut_float(table_length_); fastscan::pack_lut(padded_dim, rotated_query, lut_float.data()); @@ -53,13 +75,6 @@ class Lut { size_t num_table = table_length_ / 16; sum_vl_lut_ = vl_lut * static_cast(num_table); } - Lut& operator=(Lut&& other) noexcept { - lut_ = std::move(other.lut_); - delta_ = other.delta_; - sum_vl_lut_ = other.sum_vl_lut_; - return *this; - } - [[nodiscard]] const uint8_t* lut() const { return lut_.data(); }; [[nodiscard]] T delta() const { return delta_; }; [[nodiscard]] T sum_vl() const { return sum_vl_lut_; }; diff --git a/include/rabitqlib/index/query.hpp b/include/rabitqlib/index/query.hpp index c5c6a38..3440d76 100644 --- a/include/rabitqlib/index/query.hpp +++ b/include/rabitqlib/index/query.hpp @@ -4,7 +4,6 @@ #include #include #include -#include #include #include "rabitqlib/defines.hpp" @@ -28,9 +27,7 @@ class BatchQuery { explicit BatchQuery( const T* rotated_query, size_t padded_dim, MetricType metric_type = METRIC_L2 ) - : metric_type_(metric_type) { - lookup_table_ = std::move(Lut(rotated_query, padded_dim)); - + : lookup_table_(rotated_query, padded_dim), metric_type_(metric_type) { float c_1 = -((1 << 1) - 1) / 2.F; T sumq = @@ -74,11 +71,9 @@ class SplitBatchQuery { MetricType metric_type = METRIC_L2, bool use_hacc = true ) - : rotated_query_(rotated_query) { - lookup_table_ = std::move(Lut(rotated_query, padded_dim, use_hacc)); - - metric_type_ = (metric_type == METRIC_IP) ? METRIC_IP : METRIC_L2; - + : rotated_query_(rotated_query) + , lookup_table_(rotated_query, padded_dim, use_hacc) + , metric_type_((metric_type == METRIC_IP) ? METRIC_IP : METRIC_L2) { float c_1 = -static_cast((1 << 1) - 1) / 2.F; float c_b = -static_cast((1 << (ex_bits + 1)) - 1) / 2.F; T sumq = diff --git a/include/rabitqlib/index/symqg/detail/pipnn.hpp b/include/rabitqlib/index/symqg/detail/pipnn.hpp index 51a3fb1..eb93fe2 100644 --- a/include/rabitqlib/index/symqg/detail/pipnn.hpp +++ b/include/rabitqlib/index/symqg/detail/pipnn.hpp @@ -303,8 +303,6 @@ inline InitialGraph build_initial_graph( } using namespace pipnn_impl; const size_t threads = std::max(1, std::min(num_threads, total_threads())); - // Match QGBuilder's thread setting, including Eigen when threads == 1. - omp_set_num_threads(static_cast(threads)); ScratchPool scratch(count, dim, threads); auto leaves = cluster(data, count, dim, metric, threads, scratch); RowMajorMatrix projections(dim, kHashBits); diff --git a/include/rabitqlib/index/symqg/qg.hpp b/include/rabitqlib/index/symqg/qg.hpp index 73da6a0..2dcbdbd 100644 --- a/include/rabitqlib/index/symqg/qg.hpp +++ b/include/rabitqlib/index/symqg/qg.hpp @@ -12,6 +12,7 @@ #include #include #include +#include #include #include @@ -22,7 +23,6 @@ #include "rabitqlib/quantization/data_layout.hpp" #include "rabitqlib/quantization/pack_excode.hpp" #include "rabitqlib/quantization/rabitq.hpp" -#include "rabitqlib/utils/array.hpp" #include "rabitqlib/utils/buffer.hpp" #include "rabitqlib/utils/io.hpp" #include "rabitqlib/utils/memory.hpp" @@ -65,26 +65,24 @@ class QuantizedGraph { friend struct QGConstructionTestAccess; private: - size_t num_points_ = 0; // num points - size_t degree_bound_ = 0; // degree bound - size_t dim_ = 0; // dimension - size_t padded_dim_ = 0; // padded dimension - T (*raw_dist_func_)(const T*, const T*, size_t); // dist func for raw vector - PID entry_point_ = 0; // Entry point of graph + size_t num_points_ = 0; // num points + size_t degree_bound_ = 0; // degree bound + size_t dim_ = 0; // dimension + size_t padded_dim_ = 0; // padded dimension + T (*raw_dist_func_)(const T*, const T*, size_t) = nullptr; // raw-vector distance + PID entry_point_ = 0; // Entry point of graph MetricType metric_type_ = MetricType::METRIC_L2; RotatorType rotator_type_ = RotatorType::FhtKacRotator; size_t quantization_bits_ = 0; // 0: raw vectors, 4/8: packed RaBitQ vectors std::vector centroid_; // rotated global centroid for qg-quant ex_ipfunc quantized_ip_func_ = nullptr; - Array< - char, - std::vector, - memory::AlignedAllocator< - char, - 1 << 22, - true>> - data_; // vectors/codes + graph quantization data + edges + using RowStorage = std::vector>; + // Complete rows are contiguous in both raw and quantized modes: + // vector/code, neighbor quantization data, then packed neighbor IDs. + // Typed storage establishes T lifetimes for raw vectors; packed portions + // are accessed only as bytes, never as T values. + RowStorage data_; std::unique_ptr> rotator_; // data rotator // Position of row data (raw vector or packed qg-quant vector), neighbor @@ -94,29 +92,56 @@ class QuantizedGraph { size_t neighbor_offset_ = 0; // offset of neighbors size_t row_offset_ = 0; // length of entire row size_t ef_ = 0; + bool ready_ = false; + + [[nodiscard]] static size_t checked_add(size_t lhs, size_t rhs) { + if (lhs > std::numeric_limits::max() - rhs) { + throw std::length_error("QuantizedGraph storage size exceeds size_t"); + } + return lhs + rhs; + } + + [[nodiscard]] static size_t checked_multiply(size_t lhs, size_t rhs) { + if (lhs != 0 && rhs > std::numeric_limits::max() / lhs) { + throw std::length_error("QuantizedGraph storage size exceeds size_t"); + } + return lhs * rhs; + } + + [[nodiscard]] static size_t padded_dimension(size_t dim) { + return (checked_add(dim, 63) / 64) * 64; + } void validate_configuration() const; + void initialize_layout(); + void initialize(); - void copy_vectors(const T*); + void copy_vectors(const T*, size_t); void set_quantization_centroid(const T* centroid); + [[nodiscard]] char* get_row_data(PID data_id) { + return reinterpret_cast(get_vector(data_id)); + } + + [[nodiscard]] const char* get_row_data(PID data_id) const { + return reinterpret_cast(get_vector(data_id)); + } + [[nodiscard]] T* get_vector(PID data_id) { - return reinterpret_cast(&data_.at(row_offset_ * data_id)); + return data_.data() + ((row_offset_ / sizeof(T)) * data_id); } [[nodiscard]] const T* get_vector(PID data_id) const { - return reinterpret_cast(&data_.at(row_offset_ * data_id)); + return data_.data() + ((row_offset_ / sizeof(T)) * data_id); } - [[nodiscard]] char* get_quantized_vector(PID data_id) { - return &data_.at(row_offset_ * data_id); - } + [[nodiscard]] char* get_quantized_vector(PID data_id) { return get_row_data(data_id); } [[nodiscard]] const char* get_quantized_vector(PID data_id) const { - return &data_.at(row_offset_ * data_id); + return get_row_data(data_id); } void prepare_query(const T*, std::vector&, std::optional>&) const; @@ -131,21 +156,23 @@ class QuantizedGraph { void reconstruct_quantized_vector(PID, T*) const; [[nodiscard]] char* get_batch_data(PID data_id) { - return &data_.at((row_offset_ * data_id) + batch_data_offset_); + return get_row_data(data_id) + batch_data_offset_; } [[nodiscard]] const char* get_batch_data(PID data_id) const { - return &data_.at((row_offset_ * data_id) + batch_data_offset_); + return get_row_data(data_id) + batch_data_offset_; } - [[nodiscard]] PID* get_neighbors(PID data_id) { - return reinterpret_cast(&data_.at((row_offset_ * data_id) + neighbor_offset_) + [[nodiscard]] rabitqlib::detail::PackedArrayView get_neighbors(PID data_id) { + return rabitqlib::detail::PackedArrayView( + get_row_data(data_id) + neighbor_offset_ ); } - [[nodiscard]] const PID* get_neighbors(PID data_id) const { - return reinterpret_cast( - &data_.at((row_offset_ * data_id) + neighbor_offset_) + [[nodiscard]] rabitqlib::detail::ConstPackedArrayView get_neighbors(PID data_id + ) const { + return rabitqlib::detail::ConstPackedArrayView( + get_row_data(data_id) + neighbor_offset_ ); } @@ -176,6 +203,11 @@ class QuantizedGraph { ~QuantizedGraph() = default; + QuantizedGraph(const QuantizedGraph&) = delete; + QuantizedGraph& operator=(const QuantizedGraph&) = delete; + QuantizedGraph(QuantizedGraph&&) noexcept = default; + QuantizedGraph& operator=(QuantizedGraph&&) noexcept = default; + [[nodiscard]] auto num_vertices() const { return this->num_points_; } [[nodiscard]] auto dimension() const { return this->dim_; } @@ -190,7 +222,12 @@ class QuantizedGraph { [[nodiscard]] bool is_quantized() const { return quantization_bits_ != 0; } - void set_ep(PID entry) { this->entry_point_ = entry; }; + void set_ep(PID entry) { + if (entry >= num_points_) { + throw std::invalid_argument("QuantizedGraph entry point is out of range"); + } + entry_point_ = entry; + } void save(const char*) const; @@ -231,6 +268,13 @@ inline QuantizedGraph::QuantizedGraph( template inline void QuantizedGraph::validate_configuration() const { validate_metric_type(metric_type_); + if (dim_ == 0) { + throw std::invalid_argument("QuantizedGraph dimension must be positive"); + } + if (rotator_type_ != RotatorType::MatrixRotator && + rotator_type_ != RotatorType::FhtKacRotator) { + throw std::invalid_argument("QuantizedGraph rotator type is invalid"); + } if (degree_bound_ == 0 || degree_bound_ % fastscan::kBatchSize != 0) { throw std::invalid_argument( "QuantizedGraph degree bound must be a positive multiple of 32" @@ -246,6 +290,9 @@ inline void QuantizedGraph::validate_configuration() const { "QuantizedGraph point count exceeds the search-buffer ID limit" ); } + if (entry_point_ >= num_points_) { + throw std::invalid_argument("QuantizedGraph entry point is out of range"); + } if (quantization_bits_ != 0 && quantization_bits_ != 4 && quantization_bits_ != 8) { throw std::invalid_argument( "QuantizedGraph quantization bits must be 0 (vanilla), 4, or 8" @@ -258,7 +305,8 @@ inline void QuantizedGraph::validate_configuration() const { } template -inline void QuantizedGraph::copy_vectors(const T* data) { +inline void QuantizedGraph::copy_vectors(const T* data, size_t num_threads) { + const int thread_count = static_cast(num_threads); if (quantization_bits_ != 0) { if constexpr (!std::is_same_v) { throw std::logic_error("qg-quant currently requires float data"); @@ -268,7 +316,7 @@ inline void QuantizedGraph::copy_vectors(const T* data) { "qg-quant centroid must be set before copying vectors" ); } -#pragma omp parallel +#pragma omp parallel num_threads(thread_count) { std::vector rotated_data(padded_dim_); std::vector quantized_data(padded_dim_); @@ -278,6 +326,8 @@ inline void QuantizedGraph::copy_vectors(const T* data) { ExDataMap output( get_quantized_vector(i), padded_dim_, quantization_bits_ ); + T f_add; + T f_rescale; T unused_f_error = 0; quant::quantize_full_single( rotated_data.data(), @@ -285,11 +335,13 @@ inline void QuantizedGraph::copy_vectors(const T* data) { padded_dim_, quantization_bits_, quantized_data.data(), - output.f_add_ex(), - output.f_rescale_ex(), + f_add, + f_rescale, unused_f_error, metric_type_ ); + output.f_add_ex() = f_add; + output.f_rescale_ex() = f_rescale; quant::rabitq_impl::ex_bits::packing_rabitqplus_code( quantized_data.data(), output.ex_code(), @@ -301,7 +353,7 @@ inline void QuantizedGraph::copy_vectors(const T* data) { return; } } -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(thread_count) for (size_t i = 0; i < num_points_; ++i) { const T* src = data + (dim_ * i); T* dst = get_vector(i); @@ -320,10 +372,17 @@ inline void QuantizedGraph::set_quantization_centroid(const T* centroid) { template inline void QuantizedGraph::save(const char* filename) const { + if (!ready_ || rotator_ == nullptr) { + throw std::logic_error("QuantizedGraph must be built or loaded before save"); + } + if (filename == nullptr || filename[0] == '\0') { + throw std::invalid_argument("QuantizedGraph save filename must not be empty"); + } std::ofstream output(filename, std::ios::binary); if (!output.is_open()) { throw std::runtime_error("Cannot open quantized graph file for writing"); } + output.exceptions(std::ios::badbit | std::ios::failbit); constexpr uint64_t kFormatMagic = 0x5147524142495451ULL; // "QGRABITQ" constexpr uint32_t kFormatVersion = 1; @@ -352,16 +411,23 @@ inline void QuantizedGraph::save(const char* filename) const { } /* Data */ - data_.save(output); + output.write( + get_row_data(0), + static_cast(checked_multiply(num_points_, row_offset_)) + ); /* Rotator */ this->rotator_->save(output); + output.flush(); output.close(); } template inline void QuantizedGraph::load(const char* filename) { + if (filename == nullptr || filename[0] == '\0') { + throw std::invalid_argument("QuantizedGraph load filename must not be empty"); + } /* Check existence */ if (!file_exists(filename)) { throw std::runtime_error("Quantized graph file does not exist"); @@ -372,13 +438,30 @@ inline void QuantizedGraph::load(const char* filename) { throw std::runtime_error("Cannot open quantized graph file"); } + auto read_exact = [&](void* destination, size_t bytes, const char* field) { + if (bytes > static_cast(std::numeric_limits::max())) { + throw std::runtime_error("QuantizedGraph field is too large to read"); + } + input.read( + reinterpret_cast(destination), static_cast(bytes) + ); + if (!input) { + throw std::runtime_error( + std::string("Truncated QuantizedGraph file while reading ") + field + ); + } + }; + auto read_value = [&](auto& value, const char* field) { + read_exact(&value, sizeof(value), field); + }; + constexpr uint64_t kFormatMagic = 0x5147524142495451ULL; // "QGRABITQ" constexpr uint32_t kFormatVersion = 1; uint64_t magic = 0; - input.read(reinterpret_cast(&magic), sizeof(magic)); + read_value(magic, "format marker"); if (magic == kFormatMagic) { uint32_t version = 0; - input.read(reinterpret_cast(&version), sizeof(version)); + read_value(version, "format version"); if (version != kFormatVersion) { throw std::runtime_error("Unsupported QuantizedGraph file version"); } @@ -388,46 +471,88 @@ inline void QuantizedGraph::load(const char* filename) { input.seekg(0); } - /* Basic variants */ - input.read(reinterpret_cast(&num_points_), sizeof(size_t)); - input.read(reinterpret_cast(°ree_bound_), sizeof(size_t)); - input.read(reinterpret_cast(&dim_), sizeof(size_t)); - input.read(reinterpret_cast(&padded_dim_), sizeof(size_t)); - input.read(reinterpret_cast(&entry_point_), sizeof(PID)); - input.read(reinterpret_cast(&rotator_type_), sizeof(RotatorType)); - input.read(reinterpret_cast(&metric_type_), sizeof(MetricType)); + QuantizedGraph loaded; + size_t stored_padded_dim = 0; + read_value(loaded.num_points_, "point count"); + read_value(loaded.degree_bound_, "degree bound"); + read_value(loaded.dim_, "dimension"); + read_value(stored_padded_dim, "padded dimension"); + read_value(loaded.entry_point_, "entry point"); + read_value(loaded.rotator_type_, "rotator type"); + read_value(loaded.metric_type_, "metric type"); if (magic == kFormatMagic) { - input.read( - reinterpret_cast(&quantization_bits_), sizeof(quantization_bits_) - ); + read_value(loaded.quantization_bits_, "quantization bits"); } else { - quantization_bits_ = 0; + loaded.quantization_bits_ = 0; } - raw_dist_func_ = (metric_type_ == METRIC_IP) ? dot_product_dis : euclidean_sqr; + loaded.raw_dist_func_ = + (loaded.metric_type_ == METRIC_IP) ? dot_product_dis : euclidean_sqr; + loaded.validate_configuration(); + loaded.padded_dim_ = padded_dimension(loaded.dim_); + if (stored_padded_dim != loaded.padded_dim_) { + throw std::runtime_error("Invalid padded dimension in quantized graph file"); + } + loaded.initialize_layout(); - validate_configuration(); - initialize(); + const size_t centroid_bytes = loaded.quantization_bits_ == 0 + ? 0 + : checked_multiply(loaded.padded_dim_, sizeof(T)); + const size_t data_bytes = checked_multiply(loaded.num_points_, loaded.row_offset_); + const size_t rotator_bytes = + loaded.rotator_type_ == RotatorType::MatrixRotator + ? checked_multiply(checked_multiply(sizeof(T), loaded.dim_), loaded.padded_dim_) + : checked_multiply(loaded.padded_dim_, size_t{4}) / 8; + const size_t expected_payload_bytes = + checked_add(checked_add(centroid_bytes, data_bytes), rotator_bytes); - if (quantization_bits_ != 0) { - centroid_.resize(padded_dim_); - input.read(reinterpret_cast(centroid_.data()), padded_dim_ * sizeof(T)); + const auto payload_position = input.tellg(); + if (payload_position < 0) { + throw std::runtime_error("Cannot determine QuantizedGraph payload position"); + } + const size_t file_size = get_filesize(filename); + const size_t payload_offset = static_cast(payload_position); + if (payload_offset > file_size || + file_size - payload_offset != expected_payload_bytes) { + throw std::runtime_error("Invalid QuantizedGraph payload size"); } - /* Data */ - data_.load(input); + loaded.initialize(); - /* Rotator */ - this->rotator_->load(input); - if (rotator_->size() != padded_dim_) { - throw std::runtime_error("Invalid padded dimension in quantized graph file"); + if (loaded.quantization_bits_ != 0) { + loaded.centroid_.resize(loaded.padded_dim_); + read_exact(loaded.centroid_.data(), centroid_bytes, "quantization centroid"); + } + + read_exact(loaded.get_row_data(0), data_bytes, "graph data"); + + for (PID source = 0; source < loaded.num_points_; ++source) { + const auto neighbors = loaded.get_neighbors(source); + for (size_t i = 0; i < loaded.degree_bound_; ++i) { + if (neighbors[i] >= loaded.num_points_) { + throw std::runtime_error("Invalid QuantizedGraph neighbor ID"); + } + } + } + + loaded.rotator_->load(input); + if (!input) { + throw std::runtime_error("Truncated QuantizedGraph file while reading rotator"); } input.close(); + // ef is a runtime search setting rather than persisted index state. Preserve + // the target object's value, matching the previous in-place load behavior. + loaded.ef_ = ef_; + loaded.ready_ = true; + *this = std::move(loaded); } template inline void QuantizedGraph::set_ef(size_t cur_ef) { + if (cur_ef == 0) { + throw std::invalid_argument("QuantizedGraph ef must be positive"); + } this->ef_ = cur_ef; } @@ -438,6 +563,19 @@ inline void QuantizedGraph::search( uint32_t* __restrict__ results, T* __restrict__ dists ) { + if (!ready_ || rotator_ == nullptr) { + throw std::logic_error("QuantizedGraph must be built or loaded before search"); + } + if (query == nullptr || results == nullptr || dists == nullptr) { + throw std::invalid_argument("QuantizedGraph search buffers must not be null"); + } + if (k == 0 || k > num_points_) { + throw std::invalid_argument("QuantizedGraph k must be between 1 and num_points"); + } + if (ef_ < k) { + throw std::invalid_argument("QuantizedGraph ef must be at least k"); + } + std::vector rotated_query(padded_dim_); std::optional> quantized_query; prepare_query(query, rotated_query, quantized_query); @@ -477,6 +615,9 @@ inline void QuantizedGraph::search( } update_results(res_pool, *vis, query, quantized_query ? &*quantized_query : nullptr); + if (res_pool.size() != k) { + throw std::runtime_error("QuantizedGraph search could not produce k results"); + } res_pool.copy_results(results, dists); } @@ -522,7 +663,7 @@ void QuantizedGraph::scan_neighbors( batch_data += QGBatchDataMap::data_bytes(padded_dim_); } - const PID* ptr_nb = get_neighbors(data_id); + const auto neighbors = get_neighbors(data_id); for (size_t begin = 0; begin < cur_degree; begin += fastscan::kBatchSize) { const T threshold = search_pool.top_dist(); uint32_t candidate_mask = 0; @@ -541,16 +682,14 @@ void QuantizedGraph::scan_neighbors( const auto lane = static_cast(__builtin_ctz(candidate_mask)); candidate_mask &= candidate_mask - 1; const size_t i = begin + lane; - PID cur_neighbor = ptr_nb[i]; + PID cur_neighbor = neighbors[i]; T dist = est_dist[i]; if (search_pool.is_full(dist) || vis.get(cur_neighbor)) { continue; } search_pool.insert(cur_neighbor, dist); // update search buffer - memory::mem_prefetch_l2( - reinterpret_cast(get_vector(search_pool.next_id())), 10 - ); + memory::mem_prefetch_l2(get_row_data(search_pool.next_id()), 10); } } } @@ -566,11 +705,14 @@ inline void QuantizedGraph::update_results( return; } - auto data = result_pool.data(); - for (auto record : data) { - PID* ptr_nb = get_neighbors(record.id); + const auto& pool_data = result_pool.data(); + const std::vector> data( + pool_data.begin(), pool_data.begin() + static_cast(result_pool.size()) + ); + for (const auto& record : data) { + auto neighbors = get_neighbors(record.id); for (uint32_t i = 0; i < this->degree_bound_; ++i) { - PID cur_neighbor = ptr_nb[i]; + PID cur_neighbor = neighbors[i]; if (!vis.get(cur_neighbor)) { vis.set(cur_neighbor); result_pool.insert( @@ -586,27 +728,40 @@ inline void QuantizedGraph::update_results( // initialize const offsets & data array template -inline void QuantizedGraph::initialize() { - rotator_.reset( - choose_rotator(dim_, rotator_type_, round_up_to_multiple(dim_, 64)) +inline void QuantizedGraph::initialize_layout() { + if (quantization_bits_ == 0) { + batch_data_offset_ = checked_multiply(dim_, sizeof(T)); + } else { + const size_t code_bits = checked_multiply(padded_dim_, quantization_bits_); + batch_data_offset_ = checked_add(code_bits / 8, checked_multiply(sizeof(T), 2)); + } + + const size_t binary_batch_bytes = + checked_multiply(padded_dim_, fastscan::kBatchSize) / 8; + const size_t factor_bytes = + checked_multiply(checked_multiply(sizeof(T), fastscan::kBatchSize), size_t{2}); + const size_t batch_bytes = checked_add(binary_batch_bytes, factor_bytes); + neighbor_offset_ = checked_add( + batch_data_offset_, + checked_multiply(batch_bytes, degree_bound_ / fastscan::kBatchSize) ); - padded_dim_ = rotator_->size(); + row_offset_ = + checked_add(neighbor_offset_, checked_multiply(degree_bound_, sizeof(PID))); +} + +template +inline void QuantizedGraph::initialize() { + padded_dim_ = padded_dimension(dim_); + rotator_.reset(choose_rotator(dim_, rotator_type_, padded_dim_)); - /* check size */ assert(padded_dim_ % 64 == 0); assert(padded_dim_ >= dim_); - this->batch_data_offset_ = - quantization_bits_ == 0 ? dim_ * sizeof(T) - : ExDataMap::data_bytes(padded_dim_, quantization_bits_); - this->neighbor_offset_ = - batch_data_offset_ + (QGBatchDataMap::data_bytes(padded_dim_) * - (degree_bound_ / fastscan::kBatchSize)); - this->row_offset_ = neighbor_offset_ + (degree_bound_ * sizeof(PID)); + initialize_layout(); - data_ = Array, memory::AlignedAllocator>( - std::vector{num_points_, row_offset_} - ); + const size_t data_bytes = checked_multiply(num_points_, row_offset_); + assert(row_offset_ % sizeof(T) == 0); + data_ = RowStorage(data_bytes / sizeof(T)); if (quantization_bits_ != 0) { quantized_ip_func_ = select_excode_ipfunc(quantization_bits_); @@ -732,9 +887,9 @@ inline void QuantizedGraph::update_qg( return; } // copy neighbors - PID* neighbor_ptr = get_neighbors(cur_id); + auto neighbors = get_neighbors(cur_id); for (size_t i = 0; i < cur_degree; ++i) { - neighbor_ptr[i] = new_neighbors[i].id; + neighbors[i] = new_neighbors[i].id; } // rotated data @@ -758,18 +913,30 @@ inline void QuantizedGraph::update_qg( // quantize batches for current vertex auto* batch_data = get_batch_data(cur_id); - const auto* data = rotated_data.data(); for (size_t i = 0; i < cur_degree; i += fastscan::kBatchSize) { + const size_t batch_size = std::min(cur_degree - i, fastscan::kBatchSize); quant::quantize_qg_batch( - data, + rotated_data.data() + (i * padded_dim_), rotated_centroid.data(), - std::min(cur_degree - i, fastscan::kBatchSize), + batch_size, padded_dim_, batch_data, metric_type_ ); + if (batch_size < fastscan::kBatchSize) { + QGBatchDataMap batch(batch_data, padded_dim_); + std::fill( + batch.f_add() + batch_size, + batch.f_add() + fastscan::kBatchSize, + static_cast(0) + ); + std::fill( + batch.f_rescale() + batch_size, + batch.f_rescale() + fastscan::kBatchSize, + static_cast(0) + ); + } - data += fastscan::kBatchSize * padded_dim_; batch_data += QGBatchDataMap::data_bytes(padded_dim_); } } diff --git a/include/rabitqlib/index/symqg/qg_builder.hpp b/include/rabitqlib/index/symqg/qg_builder.hpp index c2e78d4..80f63cc 100644 --- a/include/rabitqlib/index/symqg/qg_builder.hpp +++ b/include/rabitqlib/index/symqg/qg_builder.hpp @@ -65,9 +65,7 @@ class QGBuilder { , num_threads_{std::max(1, std::min(num_threads, total_threads()))} , num_nodes_{qg_.num_vertices()} , dim_{qg_.dimension()} - , degree_bound_(qg_.degree_bound()) { - omp_set_num_threads(static_cast(num_threads_)); - } + , degree_bound_(qg_.degree_bound()) {} public: explicit QGBuilder( @@ -78,6 +76,12 @@ class QGBuilder { QGInitialization init = QGInitialization::PiPNN ) : QGBuilder(index, ef_build, num_threads) { + if (data == nullptr) { + throw std::invalid_argument("QGBuilder data must not be null"); + } + if (init != QGInitialization::PiPNN && init != QGInitialization::Random) { + throw std::invalid_argument("Unknown QG initialization"); + } if (init == QGInitialization::PiPNN) { auto seed = detail::build_initial_graph( data, num_nodes_, dim_, degree_bound_, qg_.metric_type_, num_threads_ @@ -87,8 +91,6 @@ class QGBuilder { seeded_ = false; initialize_storage(data); random_init(); - } else { - throw std::invalid_argument("Unknown QG initialization"); } } @@ -118,7 +120,7 @@ class QGBuilder { } } initialize_storage(data); -#pragma omp parallel +#pragma omp parallel num_threads(num_threads_) { CandidateList row; row.reserve(degree_bound_); @@ -147,7 +149,11 @@ class QGBuilder { public: // One complete search/prune/reverse-edge/degree-completion iteration. Call // before serving a seeded graph: search() requires full final rows. - void refine() { iter(true); } + void refine() { + qg_.ready_ = false; + iter(true); + qg_.ready_ = true; + } // One refinement for PiPNN initialization, three passes for random init. void build() { @@ -164,16 +170,18 @@ class QGBuilder { "The number of QG build iterations must be at least 2" ); } + qg_.ready_ = false; // for first iterations, we do not need to refine the graph structure for (size_t i = 0; i < num_iter - 1; ++i) { iter(false); } iter(true); + qg_.ready_ = true; } [[nodiscard]] bool check_dup() const { std::atomic flag(false); -#pragma omp parallel for +#pragma omp parallel for num_threads(num_threads_) for (size_t i = 0; i < num_nodes_; ++i) { std::unordered_set edges; for (auto nei : new_neighbors_[i]) { @@ -193,6 +201,9 @@ class QGBuilder { }; inline void QGBuilder::initialize_storage(const float* data) { + if (data == nullptr) { + throw std::invalid_argument("QGBuilder data must not be null"); + } // Allocate refinement scratch only after PiPNN's dense workspace is released. new_neighbors_.resize(num_nodes_); pruned_neighbors_.resize(num_nodes_); @@ -203,8 +214,9 @@ inline void QGBuilder::initialize_storage(const float* data) { degrees_.assign(num_nodes_, degree_bound_); std::vector centroid = compute_centroid(data, num_nodes_, dim_, num_threads_); + qg_.ready_ = false; qg_.set_quantization_centroid(centroid.data()); - qg_.copy_vectors(data); + qg_.copy_vectors(data, num_threads_); PID entry_point = 0; if (qg_.is_quantized()) { @@ -345,7 +357,7 @@ inline void QGBuilder::heuristic_prune( * @param refine refine = true means recording pruned candidates */ inline void QGBuilder::search_new_neighbors(bool refine) { -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (size_t i = 0; i < num_nodes_; ++i) { PID cur_id = i; auto tid = omp_get_thread_num(); @@ -361,7 +373,7 @@ inline void QGBuilder::search_new_neighbors(bool refine) { std::vector reconstructed; std::optional> prepared; const float* source = qg_.prepare_build_query(cur_id, reconstructed, prepared); - const PID* ids = qg_.get_neighbors(cur_id); + const auto ids = qg_.get_neighbors(cur_id); for (size_t j = 0; j < degrees_[cur_id]; ++j) { if (ids[j] != cur_id && !vis.get(ids[j])) { candidates.emplace_back( @@ -400,7 +412,7 @@ inline void QGBuilder::add_reverse_edges(bool refine) { // Keep new_neighbors_ read-only while reverse candidates are collected. Mutating a // destination row here races with another worker reading that row as its source. -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (PID data_id = 0; data_id < num_nodes_; ++data_id) { for (const auto& nei : new_neighbors_[data_id]) { const PID destination = nei.id; @@ -426,7 +438,7 @@ inline void QGBuilder::add_reverse_edges(bool refine) { } } -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (PID data_id = 0; data_id < num_nodes_; ++data_id) { CandidateList& tmp_pool = reverse_buffer[data_id]; if (qg_.is_quantized() && !tmp_pool.empty()) { @@ -450,7 +462,7 @@ inline void QGBuilder::add_reverse_edges(bool refine) { inline void QGBuilder::random_init() { const PID min_id = 0; const PID max_id = num_nodes_ - 1; -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (size_t i = 0; i < num_nodes_; ++i) { std::unordered_set neighbor_set; neighbor_set.reserve(degree_bound_); @@ -483,7 +495,7 @@ inline void QGBuilder::random_init() { * */ inline void QGBuilder::graph_refine() { -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (size_t i = 0; i < num_nodes_; ++i) { CandidateList& cur_neighbors = new_neighbors_[i]; size_t cur_degree = cur_neighbors.size(); @@ -560,7 +572,7 @@ inline void QGBuilder::iter(bool refine) { } // update qg -#pragma omp parallel for schedule(dynamic) +#pragma omp parallel for schedule(dynamic) num_threads(num_threads_) for (size_t i = 0; i < num_nodes_; ++i) { qg_.update_qg(i, new_neighbors_[i]); degrees_[i] = new_neighbors_[i].size(); diff --git a/include/rabitqlib/quantization/data_layout.hpp b/include/rabitqlib/quantization/data_layout.hpp index bc12905..15231b1 100644 --- a/include/rabitqlib/quantization/data_layout.hpp +++ b/include/rabitqlib/quantization/data_layout.hpp @@ -2,24 +2,169 @@ #include #include +#include +#include +#include #include "rabitqlib/fastscan/fastscan.hpp" namespace rabitqlib { +namespace detail { + +// Factor blocks are part of a packed byte layout and therefore need not be aligned for T. +// Access them through memcpy so the representation remains valid C++17 storage rather than +// manufacturing T objects inside a char buffer with reinterpret_cast. +template +class PackedValueRef { + static_assert( + std::is_trivially_copyable_v, "packed values must be trivially copyable" + ); + + public: + explicit PackedValueRef(char* data) : data_(data) {} + PackedValueRef(const PackedValueRef&) = default; + + PackedValueRef& operator=(T value) { + std::memcpy(data_, &value, sizeof(value)); + return *this; + } + + PackedValueRef& operator=(const PackedValueRef& other) { + return *this = static_cast(other); + } + + operator T() const { + T value; + std::memcpy(&value, data_, sizeof(value)); + return value; + } + + private: + char* data_; +}; + +template +class PackedArrayView { + static_assert( + std::is_trivially_copyable_v, "packed values must be trivially copyable" + ); + + public: + using difference_type = std::ptrdiff_t; + using value_type = T; + using pointer = void; + using reference = PackedValueRef; + using iterator_category = std::forward_iterator_tag; + + explicit PackedArrayView(char* data) : data_(data) {} + + [[nodiscard]] reference operator*() const { return reference(data_); } + [[nodiscard]] reference operator[](size_t index) const { + return reference(data_ + (index * sizeof(T))); + } + + PackedArrayView& operator++() { + data_ += sizeof(T); + return *this; + } + PackedArrayView operator++(int) { + PackedArrayView previous = *this; + ++*this; + return previous; + } + PackedArrayView& operator--() { + data_ -= sizeof(T); + return *this; + } + PackedArrayView& operator+=(difference_type count) { + data_ += count * static_cast(sizeof(T)); + return *this; + } + PackedArrayView& operator-=(difference_type count) { return *this += -count; } + + friend PackedArrayView operator+(PackedArrayView view, difference_type count) { + view += count; + return view; + } + friend PackedArrayView operator+(difference_type count, PackedArrayView view) { + return view + count; + } + friend PackedArrayView operator-(PackedArrayView view, difference_type count) { + view -= count; + return view; + } + friend difference_type operator-(PackedArrayView lhs, PackedArrayView rhs) { + return (lhs.data_ - rhs.data_) / static_cast(sizeof(T)); + } + friend bool operator==(PackedArrayView lhs, PackedArrayView rhs) { + return lhs.data_ == rhs.data_; + } + friend bool operator!=(PackedArrayView lhs, PackedArrayView rhs) { + return !(lhs == rhs); + } + + void copy_to(T* output, size_t count) const { + std::memcpy(output, data_, count * sizeof(T)); + } + + private: + char* data_; +}; + +template +class ConstPackedArrayView { + static_assert( + std::is_trivially_copyable_v, "packed values must be trivially copyable" + ); + + public: + explicit ConstPackedArrayView(const char* data) : data_(data) {} + + [[nodiscard]] T operator[](size_t index) const { + T value; + std::memcpy(&value, data_ + (index * sizeof(T)), sizeof(value)); + return value; + } + + void copy_to(T* output, size_t count) const { + std::memcpy(output, data_, count * sizeof(T)); + } + + private: + const char* data_; +}; + +template +[[nodiscard]] inline T load_packed_value(const char* data) { + static_assert( + std::is_trivially_copyable_v, "packed values must be trivially copyable" + ); + T value; + std::memcpy(&value, data, sizeof(value)); + return value; +} + +} // namespace detail + template struct BatchDataMap { public: explicit BatchDataMap(char* data, size_t padded_dim) : batch_bin_code_(reinterpret_cast(data)) - , f_add_(reinterpret_cast(data + (padded_dim * fastscan::kBatchSize / 8)) - ) // 1 bit code - , f_rescale_(f_add_ + fastscan::kBatchSize) - , f_error_(f_rescale_ + fastscan::kBatchSize) {} + , f_add_(data + (padded_dim * fastscan::kBatchSize / 8)) // 1 bit code + , f_rescale_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize) + ) + , f_error_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize * 2) + ) {} [[nodiscard]] uint8_t* bin_code() { return batch_bin_code_; } - [[nodiscard]] T* f_add() { return f_add_; } - [[nodiscard]] T* f_rescale() { return f_rescale_; } - [[nodiscard]] T* f_error() { return f_error_; } + [[nodiscard]] detail::PackedArrayView f_add() { return f_add_; } + [[nodiscard]] detail::PackedArrayView f_rescale() { return f_rescale_; } + [[nodiscard]] detail::PackedArrayView f_error() { return f_error_; } static size_t data_bytes(size_t padded_dim) { return (padded_dim * fastscan::kBatchSize / 8) + @@ -28,9 +173,9 @@ struct BatchDataMap { private: uint8_t* batch_bin_code_; - T* f_add_; - T* f_rescale_; - T* f_error_; + detail::PackedArrayView f_add_; + detail::PackedArrayView f_rescale_; + detail::PackedArrayView f_error_; }; template @@ -38,21 +183,26 @@ struct ConstBatchDataMap { public: explicit ConstBatchDataMap(const char* data, size_t padded_dim) : batch_bin_code_(reinterpret_cast(data)) - , f_add_(reinterpret_cast(data + (padded_dim * fastscan::kBatchSize / 8)) - ) // 1 bit code - , f_rescale_(f_add_ + fastscan::kBatchSize) - , f_error_(f_rescale_ + fastscan::kBatchSize) {} + , f_add_(data + (padded_dim * fastscan::kBatchSize / 8)) // 1 bit code + , f_rescale_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize) + ) + , f_error_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize * 2) + ) {} [[nodiscard]] const uint8_t* bin_code() const { return batch_bin_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_; } + [[nodiscard]] detail::ConstPackedArrayView f_add() const { return f_add_; } + [[nodiscard]] detail::ConstPackedArrayView f_rescale() const { return f_rescale_; } + [[nodiscard]] detail::ConstPackedArrayView f_error() const { return f_error_; } private: const uint8_t* batch_bin_code_; - const T* f_add_; - const T* f_rescale_; - const T* f_error_; + detail::ConstPackedArrayView f_add_; + detail::ConstPackedArrayView f_rescale_; + detail::ConstPackedArrayView f_error_; }; template @@ -60,13 +210,15 @@ struct QGBatchDataMap { public: explicit QGBatchDataMap(char* data, size_t padded_dim) : batch_bin_code_(reinterpret_cast(data)) - , f_add_(reinterpret_cast(data + (padded_dim * fastscan::kBatchSize / 8)) - ) // 1 bit code - , f_rescale_(f_add_ + fastscan::kBatchSize) {} + , f_add_(data + (padded_dim * fastscan::kBatchSize / 8)) // 1 bit code + , f_rescale_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize) + ) {} [[nodiscard]] uint8_t* bin_code() { return batch_bin_code_; } - [[nodiscard]] T* f_add() { return f_add_; } - [[nodiscard]] T* f_rescale() { return f_rescale_; } + [[nodiscard]] detail::PackedArrayView f_add() { return f_add_; } + [[nodiscard]] detail::PackedArrayView f_rescale() { return f_rescale_; } static size_t data_bytes(size_t padded_dim) { return (padded_dim * fastscan::kBatchSize / 8) + @@ -75,8 +227,8 @@ struct QGBatchDataMap { private: uint8_t* batch_bin_code_; - T* f_add_; - T* f_rescale_; + detail::PackedArrayView f_add_; + detail::PackedArrayView f_rescale_; }; template @@ -84,13 +236,15 @@ struct ConstQGBatchDataMap { public: explicit ConstQGBatchDataMap(const char* data, size_t padded_dim) : batch_bin_code_(reinterpret_cast(data)) - , f_add_(reinterpret_cast(data + (padded_dim * fastscan::kBatchSize / 8)) - ) // 1 bit code - , f_rescale_(f_add_ + fastscan::kBatchSize) {} + , f_add_(data + (padded_dim * fastscan::kBatchSize / 8)) // 1 bit code + , f_rescale_( + data + (padded_dim * fastscan::kBatchSize / 8) + + (sizeof(T) * fastscan::kBatchSize) + ) {} - [[nodiscard]] const uint8_t* bin_code() { return batch_bin_code_; } - [[nodiscard]] const T* f_add() { return f_add_; } - [[nodiscard]] const T* f_rescale() { return f_rescale_; } + [[nodiscard]] const uint8_t* bin_code() const { return batch_bin_code_; } + [[nodiscard]] detail::ConstPackedArrayView f_add() const { return f_add_; } + [[nodiscard]] detail::ConstPackedArrayView f_rescale() const { return f_rescale_; } static size_t data_bytes(size_t padded_dim) { return (padded_dim * fastscan::kBatchSize / 8) + @@ -99,8 +253,8 @@ struct ConstQGBatchDataMap { private: const uint8_t* batch_bin_code_; - const T* f_add_; - const T* f_rescale_; + detail::ConstPackedArrayView f_add_; + detail::ConstPackedArrayView f_rescale_; }; template @@ -108,21 +262,21 @@ struct ExDataMap { public: explicit ExDataMap(char* data, size_t padded_dim, size_t ex_bits) : ex_code_(reinterpret_cast(data)) - , f_add_ex_(*reinterpret_cast(data + (padded_dim * ex_bits / 8))) - , f_recale_ex_(*(reinterpret_cast(data + (padded_dim * ex_bits / 8)) + 1)) {} + , f_add_ex_(data + (padded_dim * ex_bits / 8)) + , f_recale_ex_(data + (padded_dim * ex_bits / 8) + sizeof(T)) {} static size_t data_bytes(size_t padded_dim, size_t ex_bits) { return ex_bits > 0 ? (padded_dim * ex_bits / 8) + (sizeof(T) * 2) : 0; } [[nodiscard]] uint8_t* ex_code() { return ex_code_; } - [[nodiscard]] T& f_add_ex() { return f_add_ex_; } - [[nodiscard]] T& f_rescale_ex() { return f_recale_ex_; } + [[nodiscard]] detail::PackedValueRef f_add_ex() { return f_add_ex_; } + [[nodiscard]] detail::PackedValueRef f_rescale_ex() { return f_recale_ex_; } private: uint8_t* ex_code_; - T& f_add_ex_; - T& f_recale_ex_; + detail::PackedValueRef f_add_ex_; + detail::PackedValueRef f_recale_ex_; }; template @@ -130,18 +284,19 @@ struct ConstExDataMap { public: explicit ConstExDataMap(const char* data, size_t padded_dim, size_t ex_bits) : ex_code_(reinterpret_cast(data)) - , f_add_ex_(*reinterpret_cast(data + (padded_dim * ex_bits / 8))) - , f_recale_ex_(*(reinterpret_cast(data + (padded_dim * ex_bits / 8)) + 1) - ) {} + , f_add_ex_(data + (padded_dim * ex_bits / 8)) + , f_recale_ex_(data + (padded_dim * ex_bits / 8) + sizeof(T)) {} [[nodiscard]] const uint8_t* ex_code() const { return ex_code_; } - [[nodiscard]] const T& f_add_ex() const { return f_add_ex_; } - [[nodiscard]] const T& f_rescale_ex() const { return f_recale_ex_; } + [[nodiscard]] T f_add_ex() const { return detail::load_packed_value(f_add_ex_); } + [[nodiscard]] T f_rescale_ex() const { + return detail::load_packed_value(f_recale_ex_); + } private: const uint8_t* ex_code_; - const T& f_add_ex_; - const T& f_recale_ex_; + const char* f_add_ex_; + const char* f_recale_ex_; }; template @@ -149,24 +304,24 @@ struct BaseDataMap { public: explicit BaseDataMap(char* data, size_t padded_dim, size_t base_bits) : base_code_(reinterpret_cast(data)) - , f_add_(*reinterpret_cast(data + (padded_dim * base_bits / 8))) - , f_rescale_(*(reinterpret_cast(data + (padded_dim * base_bits / 8)) + 1)) - , f_error_(*(reinterpret_cast(data + (padded_dim * base_bits / 8)) + 2)) {} + , f_add_(data + (padded_dim * base_bits / 8)) + , f_rescale_(data + (padded_dim * base_bits / 8) + sizeof(T)) + , f_error_(data + (padded_dim * base_bits / 8) + (sizeof(T) * 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_; } + [[nodiscard]] detail::PackedValueRef f_add() { return f_add_; } + [[nodiscard]] detail::PackedValueRef f_rescale() { return f_rescale_; } + [[nodiscard]] detail::PackedValueRef f_error() { return f_error_; } private: uint8_t* base_code_; - T& f_add_; - T& f_rescale_; - T& f_error_; + detail::PackedValueRef f_add_; + detail::PackedValueRef f_rescale_; + detail::PackedValueRef f_error_; }; template @@ -174,70 +329,69 @@ struct ConstBaseDataMap { public: explicit ConstBaseDataMap(const char* data, size_t padded_dim, size_t base_bits) : base_code_(reinterpret_cast(data)) - , f_add_(*reinterpret_cast(data + (padded_dim * base_bits / 8))) - , f_rescale_(*(reinterpret_cast(data + (padded_dim * base_bits / 8)) + 1)) - , f_error_(*(reinterpret_cast(data + (padded_dim * base_bits / 8)) + 2)) { - } + , f_add_(data + (padded_dim * base_bits / 8)) + , f_rescale_(data + (padded_dim * base_bits / 8) + sizeof(T)) + , f_error_(data + (padded_dim * base_bits / 8) + (sizeof(T) * 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_; } + [[nodiscard]] T f_add() const { return detail::load_packed_value(f_add_); } + [[nodiscard]] T f_rescale() const { return detail::load_packed_value(f_rescale_); } + [[nodiscard]] T f_error() const { return detail::load_packed_value(f_error_); } private: const uint8_t* base_code_; - const T& f_add_; - const T& f_rescale_; - const T& f_error_; + const char* f_add_; + const char* f_rescale_; + const char* f_error_; }; template struct BinDataMap { public: explicit BinDataMap(char* data, size_t padded_dim) - : bin_code_(reinterpret_cast(data)) - , f_add_(*reinterpret_cast(data + (padded_dim / 8))) - , f_rescale_(*(reinterpret_cast(data + (padded_dim / 8)) + 1)) - , f_error_(*(reinterpret_cast(data + (padded_dim / 8)) + 2)) {} + : bin_code_(reinterpret_cast(data)) + , f_add_(data + (padded_dim / 8)) + , f_rescale_(data + (padded_dim / 8) + sizeof(T)) + , f_error_(data + (padded_dim / 8) + (sizeof(T) * 2)) {} - [[nodiscard]] uint64_t* bin_code() { return bin_code_; } - [[nodiscard]] T& f_add() { return f_add_; } - [[nodiscard]] T& f_rescale() { return f_rescale_; } - [[nodiscard]] T& f_error() { return f_error_; } + [[nodiscard]] uint8_t* bin_code() { return bin_code_; } + [[nodiscard]] detail::PackedValueRef f_add() { return f_add_; } + [[nodiscard]] detail::PackedValueRef f_rescale() { return f_rescale_; } + [[nodiscard]] detail::PackedValueRef f_error() { return f_error_; } static size_t data_bytes(size_t padded_dim) { return (padded_dim / 8) + (sizeof(T) * 3); } private: - uint64_t* bin_code_; - T& f_add_; - T& f_rescale_; - T& f_error_; + uint8_t* bin_code_; + detail::PackedValueRef f_add_; + detail::PackedValueRef f_rescale_; + detail::PackedValueRef f_error_; }; template struct ConstBinDataMap { public: explicit ConstBinDataMap(const char* data, size_t padded_dim) - : bin_code_(reinterpret_cast(data)) - , f_add_(*reinterpret_cast(data + (padded_dim / 8))) - , f_rescale_(*(reinterpret_cast(data + (padded_dim / 8)) + 1)) - , f_error_(*(reinterpret_cast(data + (padded_dim / 8)) + 2)) {} + : bin_code_(reinterpret_cast(data)) + , f_add_(data + (padded_dim / 8)) + , f_rescale_(data + (padded_dim / 8) + sizeof(T)) + , f_error_(data + (padded_dim / 8) + (sizeof(T) * 2)) {} - [[nodiscard]] const uint64_t* bin_code() { return bin_code_; } - [[nodiscard]] const T& f_add() { return f_add_; } - [[nodiscard]] const T& f_rescale() { return f_rescale_; } - [[nodiscard]] const T& f_error() { return f_error_; } + [[nodiscard]] const uint8_t* bin_code() const { return bin_code_; } + [[nodiscard]] T f_add() const { return detail::load_packed_value(f_add_); } + [[nodiscard]] T f_rescale() const { return detail::load_packed_value(f_rescale_); } + [[nodiscard]] T f_error() const { return detail::load_packed_value(f_error_); } static size_t data_bytes(size_t padded_dim) { return (padded_dim / 8) + (sizeof(T) * 3); } private: - const uint64_t* bin_code_; - const T& f_add_; - const T& f_rescale_; - const T& f_error_; + const uint8_t* bin_code_; + const char* f_add_; + const char* f_rescale_; + const char* f_error_; }; -} // namespace rabitqlib \ No newline at end of file +} // namespace rabitqlib diff --git a/include/rabitqlib/quantization/rabitq.hpp b/include/rabitqlib/quantization/rabitq.hpp index f3ea9a6..5b4176f 100644 --- a/include/rabitqlib/quantization/rabitq.hpp +++ b/include/rabitqlib/quantization/rabitq.hpp @@ -137,17 +137,23 @@ inline void quantize_compact_one_bit( MetricType metric_type = METRIC_L2 ) { BinDataMap cur_bin_data(bin_data, padded_dim); + T f_add; + T f_rescale; + T f_error; - rabitq_impl::one_bit::one_bit_compact_code( + rabitq_impl::one_bit::one_bit_compact_code_to_bytes( data, centroid, padded_dim, cur_bin_data.bin_code(), - cur_bin_data.f_add(), - cur_bin_data.f_rescale(), - cur_bin_data.f_error(), + f_add, + f_rescale, + f_error, metric_type ); + cur_bin_data.f_add() = f_add; + cur_bin_data.f_rescale() = f_rescale; + cur_bin_data.f_error() = f_error; } template @@ -202,6 +208,8 @@ inline void quantize_compact_ex_bits( ExDataMap cur_ex_data(ex_data, padded_dim, ex_bits); // we do not use this error factor here + T f_add_ex; + T f_rescale_ex; T ex_error; rabitq_impl::ex_bits::ex_bits_compact_code( @@ -210,12 +218,14 @@ inline void quantize_compact_ex_bits( padded_dim, ex_bits, cur_ex_data.ex_code(), - cur_ex_data.f_add_ex(), - cur_ex_data.f_rescale_ex(), + f_add_ex, + f_rescale_ex, ex_error, metric_type, config.t_const ); + cur_ex_data.f_add_ex() = f_add_ex; + cur_ex_data.f_rescale_ex() = f_rescale_ex; } inline void quantize_split_batch( @@ -427,4 +437,4 @@ inline TF full_est_dist( return est_dist; } -} // namespace rabitqlib::quant \ No newline at end of file +} // namespace rabitqlib::quant diff --git a/include/rabitqlib/quantization/rabitq_impl.hpp b/include/rabitqlib/quantization/rabitq_impl.hpp index b0447da..2e2c1de 100644 --- a/include/rabitqlib/quantization/rabitq_impl.hpp +++ b/include/rabitqlib/quantization/rabitq_impl.hpp @@ -179,11 +179,11 @@ inline void one_bit_code_with_factor( * approximate distances between the original vectors. */ template -inline void one_bit_compact_code( +inline void one_bit_compact_code_to_bytes( const T* data, const T* centroid, size_t padded_dim, - TC* compact_code, + uint8_t* compact_code, T& f_add, T& f_recale, T& f_error, @@ -204,51 +204,91 @@ inline void one_bit_compact_code( metric_type ); - pack_binary(binary_code.data(), compact_code, padded_dim); + pack_binary_to_bytes(binary_code.data(), compact_code, padded_dim); +} + +template +inline void one_bit_compact_code( + const T* data, + const T* centroid, + size_t padded_dim, + TC* compact_code, + T& f_add, + T& f_recale, + T& f_error, + MetricType metric_type = METRIC_L2 +) { + one_bit_compact_code_to_bytes( + data, + centroid, + padded_dim, + reinterpret_cast(compact_code), + f_add, + f_recale, + f_error, + metric_type + ); } // Requires a positive padded_dim divisible by sizeof(TC) * 8. -template +template < + typename T, + typename TC, + bool Parallel = false, + typename FAdd, + typename FRescale, + typename FError> inline void one_bit_compact_codes( const T* data, const T* centroid, size_t num, size_t padded_dim, TC* compact_code, - T* f_add, - T* f_rescale, - T* f_error, + FAdd f_add, + FRescale f_rescale, + FError f_error, MetricType metric_type = METRIC_L2 ) { constexpr size_t kTypeBits = sizeof(TC) * 8; #pragma omp parallel for if (Parallel) for (size_t i = 0; i < num; ++i) { + T add; + T rescale; + T error; one_bit_compact_code( data + (padded_dim * i), centroid, padded_dim, compact_code + (padded_dim / kTypeBits * i), - f_add[i], - f_rescale[i], - f_error[i], + add, + rescale, + error, metric_type ); + f_add[i] = add; + f_rescale[i] = rescale; + f_error[i] = error; } } // Encoding requires a positive padded_dim divisible by 8; index pipelines use 64. // packed_code needs ceil(num / 32) * 32 * (padded_dim / 8) bytes, including tail padding. -template +template < + typename T, + bool Parallel = false, + typename FAdd, + typename FRescale, + typename FError> inline void one_bit_batch_code( const T* data, const T* centroid, size_t num, size_t padded_dim, uint8_t* packed_code, - T* f_add, - T* f_recale, - T* f_error, + FAdd f_add, + FRescale f_recale, + FError f_error, MetricType metric_type = METRIC_L2 ) { std::vector compact_codes(num * padded_dim / 8); @@ -873,17 +913,23 @@ inline void split_code_with_factor( // Base layer first: its factors describe the filter code on its own. BaseDataMap base_map(base_data, dim, base_bits); + T f_add_base; + T f_rescale_base; + T f_error_base; code_factors( residual, centroid, dim, base_code_int.data(), -(static_cast((1U << base_bits) - 1) / 2.F), - base_map.f_add(), - base_map.f_rescale(), - base_map.f_error(), + f_add_base, + f_rescale_base, + f_error_base, metric_type ); + base_map.f_add() = f_add_base; + base_map.f_rescale() = f_rescale_base; + base_map.f_error() = f_error_base; ex_bits::packing_rabitqplus_code(base_raw.data(), base_map.base_code(), dim, base_bits); if (ex_bits == 0) { @@ -892,6 +938,8 @@ inline void split_code_with_factor( // Refine layer: the same derivation over the combined code. ExDataMap ex_map(ex_data, dim, ex_bits); + T f_add_ex; + T f_rescale_ex; T f_error_ex = 0; code_factors( residual, @@ -899,11 +947,13 @@ inline void split_code_with_factor( dim, total_code.data(), -(static_cast((1U << total_bits) - 1) / 2.F), - ex_map.f_add_ex(), - ex_map.f_rescale_ex(), + f_add_ex, + f_rescale_ex, f_error_ex, metric_type ); + ex_map.f_add_ex() = f_add_ex; + ex_map.f_rescale_ex() = f_rescale_ex; ex_bits::packing_rabitqplus_code(ex_raw.data(), ex_map.ex_code(), dim, ex_bits); } diff --git a/include/rabitqlib/simd/space_dispatch.hpp b/include/rabitqlib/simd/space_dispatch.hpp index 67d5e64..1403e1f 100644 --- a/include/rabitqlib/simd/space_dispatch.hpp +++ b/include/rabitqlib/simd/space_dispatch.hpp @@ -85,6 +85,7 @@ void new_transpose_bin_avx2( void new_transpose_bin_512_avx2( const uint8_t* q, uint64_t* tq, size_t padded_dim, size_t b_query ); +float mask_ip_x0_q_avx2(const float* query, const uint8_t* data, size_t padded_dim); float mask_ip_x0_q_avx2(const float* query, const uint64_t* data, size_t padded_dim); void scalar_quantize_uint8_avx2( uint8_t* result, const float* vec0, size_t dim, float lo, float delta @@ -99,6 +100,7 @@ void new_transpose_bin_avx512( void new_transpose_bin_512_avx512( const uint8_t* q, uint64_t* tq, size_t padded_dim, size_t b_query ); +float mask_ip_x0_q_avx512(const float* query, const uint8_t* data, size_t padded_dim); float mask_ip_x0_q_avx512(const float* query, const uint64_t* data, size_t padded_dim); void scalar_quantize_uint8_avx512( uint8_t* result, const float* vec0, size_t dim, float lo, float delta diff --git a/include/rabitqlib/simd/warmup_dispatch.hpp b/include/rabitqlib/simd/warmup_dispatch.hpp index f7899cb..bae9358 100644 --- a/include/rabitqlib/simd/warmup_dispatch.hpp +++ b/include/rabitqlib/simd/warmup_dispatch.hpp @@ -5,6 +5,15 @@ namespace rabitqlib::simd { +float warmup_ip_x0_q_512_avx2( + const uint8_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +); + float warmup_ip_x0_q_512_avx2( const uint64_t* data, const uint64_t* query, @@ -14,6 +23,15 @@ float warmup_ip_x0_q_512_avx2( size_t b_query ); +float warmup_ip_x0_q_512_avx512( + const uint8_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +); + float warmup_ip_x0_q_512_avx512( const uint64_t* data, const uint64_t* query, diff --git a/include/rabitqlib/utils/array.hpp b/include/rabitqlib/utils/array.hpp index ee23b4c..03ea19f 100644 --- a/include/rabitqlib/utils/array.hpp +++ b/include/rabitqlib/utils/array.hpp @@ -19,10 +19,11 @@ #pragma once -#include #include #include +#include #include +#include #include #include #include @@ -39,7 +40,12 @@ template static_assert(std::is_same_v); size_t res = 1; - std::for_each(dims.begin(), dims.end(), [&](auto cur_d) { res *= cur_d; }); + for (const size_t cur_d : dims) { + if (cur_d != 0 && res > std::numeric_limits::max() / cur_d) { + throw std::bad_array_new_length(); + } + res *= cur_d; + } return res; } } // namespace array_impl @@ -56,7 +62,13 @@ class Array { [[nodiscard]] constexpr auto size() const -> size_t { return array_impl::size(dims_); } /// @brief num of bytes for all data objects - [[nodiscard]] constexpr auto bytes() const -> size_t { return sizeof(T) * size(); } + [[nodiscard]] constexpr auto bytes() const -> size_t { + const size_t num_elements = size(); + if (num_elements > std::numeric_limits::max() / sizeof(T)) { + throw std::bad_array_new_length(); + } + return sizeof(T) * num_elements; + } void destroy() { size_t num_elements = size(); @@ -131,4 +143,4 @@ class Array { [[no_unique_address]] Dims dims_; [[no_unique_address]] Alloc allocator_; }; -} // namespace rabitqlib \ No newline at end of file +} // namespace rabitqlib diff --git a/include/rabitqlib/utils/buffer.hpp b/include/rabitqlib/utils/buffer.hpp index f06da97..2da7570 100644 --- a/include/rabitqlib/utils/buffer.hpp +++ b/include/rabitqlib/utils/buffer.hpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include "rabitqlib/defines.hpp" @@ -23,7 +24,14 @@ template class SearchBuffer { private: std::vector, memory::AlignedAllocator>> data_; - size_t size_ = 0, cur_ = 0, capacity_; + size_t size_ = 0, cur_ = 0, capacity_ = 0; + + [[nodiscard]] static size_t storage_size(size_t capacity) { + if (capacity == std::numeric_limits::max()) { + throw std::length_error("SearchBuffer capacity is too large"); + } + return capacity + 1; + } [[nodiscard]] auto binary_search(T dist) const { size_t lo = 0; @@ -47,7 +55,8 @@ class SearchBuffer { public: SearchBuffer() = default; - explicit SearchBuffer(size_t capacity) : data_(capacity + 1), capacity_(capacity) {} + explicit SearchBuffer(size_t capacity) + : data_(storage_size(capacity)), capacity_(capacity) {} // insert a data point into buffer void insert(PID data_id, T dist) { @@ -84,12 +93,15 @@ class SearchBuffer { [[nodiscard]] auto has_next() const -> bool { return cur_ < size_; } void resize(size_t new_size) { - this->capacity_ = new_size; data_ = std::vector, memory::AlignedAllocator>>( - capacity_ + 1 + storage_size(new_size) ); + capacity_ = new_size; + clear(); } + [[nodiscard]] size_t size() const { return size_; } + void copy_results(PID* knn) const { for (size_t i = 0; i < size_; ++i) { knn[i] = data_[i].id; diff --git a/include/rabitqlib/utils/cpu_features.hpp b/include/rabitqlib/utils/cpu_features.hpp index 57a6827..fafe7d1 100644 --- a/include/rabitqlib/utils/cpu_features.hpp +++ b/include/rabitqlib/utils/cpu_features.hpp @@ -1,5 +1,7 @@ #pragma once +#include + namespace rabitqlib::cpu { struct Features { @@ -16,4 +18,12 @@ bool has_avx2(); bool has_avx512_core(); bool has_avx512_popcnt(); +namespace detail { + +Features filter_usable_features( + const Features& hardware, bool avx, bool osxsave, uint64_t xcr0 +); + +} // namespace detail + } // namespace rabitqlib::cpu diff --git a/include/rabitqlib/utils/io.hpp b/include/rabitqlib/utils/io.hpp index bf70e91..3bf2ff4 100644 --- a/include/rabitqlib/utils/io.hpp +++ b/include/rabitqlib/utils/io.hpp @@ -1,19 +1,61 @@ #pragma once -#include #include #include #include #include #include +#include #include #include #include +#include namespace rabitqlib { +namespace io_impl { + +inline size_t checked_add(size_t lhs, size_t rhs) { + if (lhs > std::numeric_limits::max() - rhs) { + throw std::length_error("File layout exceeds size_t"); + } + return lhs + rhs; +} + +inline size_t checked_multiply(size_t lhs, size_t rhs) { + if (lhs != 0 && rhs > std::numeric_limits::max() / lhs) { + throw std::length_error("File layout exceeds size_t"); + } + return lhs * rhs; +} + +inline void read_exact( + std::ifstream& input, void* destination, size_t bytes, const char* description +) { + if (bytes > static_cast(std::numeric_limits::max())) { + throw std::length_error("File field is too large to read"); + } + input.read(reinterpret_cast(destination), static_cast(bytes)); + if (!input) { + throw std::runtime_error( + std::string("Unexpected end of file while reading ") + description + ); + } +} + +template +void read_value(std::ifstream& input, T& value, const char* description) { + read_exact(input, &value, sizeof(value), description); +} + +} // namespace io_impl + // get num of bytes inline size_t get_filesize(const char* filename) { - return std::filesystem::file_size(filename); + const auto file_size = std::filesystem::file_size(filename); + if (file_size > std::numeric_limits::max()) { + throw std::length_error("File is too large to address"); + } + return static_cast(file_size); } inline bool file_exists(const char* filename) { return std::filesystem::exists(filename); } @@ -25,26 +67,46 @@ void load_vecs(const char* filename, M& row_mat) { throw std::runtime_error("File does not exist: " + std::string(filename)); } - assert((std::is_same_v> == true)); + static_assert(std::is_same_v>); - uint32_t tmp; - size_t file_size = get_filesize(filename); + const size_t file_size = get_filesize(filename); + if (file_size < sizeof(uint32_t)) { + throw std::runtime_error("Vector file is too small to contain a dimension"); + } std::ifstream input(filename, std::ios::binary); + if (!input.is_open()) { + throw std::runtime_error("Cannot open vector file: " + std::string(filename)); + } - input.read(reinterpret_cast(&tmp), sizeof(uint32_t)); + uint32_t stored_cols = 0; + io_impl::read_value(input, stored_cols, "vector dimension"); + if (stored_cols == 0) { + throw std::runtime_error("Vector dimension must be positive"); + } - size_t cols = tmp; - size_t rows = file_size / (cols * sizeof(T) + sizeof(uint32_t)); - row_mat = M(rows, cols); + const size_t cols = stored_cols; + const size_t record_bytes = + io_impl::checked_add(sizeof(uint32_t), io_impl::checked_multiply(cols, sizeof(T))); + if (file_size % record_bytes != 0) { + throw std::runtime_error("Vector file size is not a whole number of records"); + } + const size_t rows = file_size / record_bytes; + M loaded(rows, cols); input.seekg(0, std::ifstream::beg); for (size_t i = 0; i < rows; i++) { - input.read(reinterpret_cast(&tmp), sizeof(uint32_t)); - input.read(reinterpret_cast(&row_mat(i, 0)), sizeof(T) * cols); + uint32_t row_cols = 0; + io_impl::read_value(input, row_cols, "row dimension"); + if (row_cols != stored_cols) { + throw std::runtime_error("Vector file contains inconsistent row dimensions"); + } + io_impl::read_exact( + input, &loaded(i, 0), io_impl::checked_multiply(sizeof(T), cols), "vector data" + ); } - input.close(); + row_mat = std::move(loaded); } // load .*bin file to a matrix (e.g., RowMajorFloatMat) @@ -54,21 +116,40 @@ void load_bin(const char* filename, M& row_mat) { throw std::runtime_error("File does not exist: " + std::string(filename)); } - assert((std::is_same_v> == true)); + static_assert(std::is_same_v>); - uint32_t rows; - uint32_t cols; + const size_t file_size = get_filesize(filename); + if (file_size < 2 * sizeof(uint32_t)) { + throw std::runtime_error("Binary matrix file is too small to contain its header"); + } std::ifstream input(filename, std::ios::binary); + if (!input.is_open()) { + throw std::runtime_error( + "Cannot open binary matrix file: " + std::string(filename) + ); + } - input.read(reinterpret_cast(&rows), sizeof(uint32_t)); - input.read(reinterpret_cast(&cols), sizeof(uint32_t)); - - row_mat = M(rows, cols); + uint32_t stored_rows = 0; + uint32_t stored_cols = 0; + io_impl::read_value(input, stored_rows, "matrix row count"); + io_impl::read_value(input, stored_cols, "matrix column count"); + if (stored_rows == 0 || stored_cols == 0) { + throw std::runtime_error("Binary matrix dimensions must be positive"); + } - for (size_t i = 0; i < rows; i++) { - input.read(reinterpret_cast(&row_mat(i, 0)), sizeof(T) * cols); + const size_t rows = stored_rows; + const size_t cols = stored_cols; + const size_t payload_bytes = + io_impl::checked_multiply(io_impl::checked_multiply(rows, cols), sizeof(T)); + const size_t expected_file_size = + io_impl::checked_add(2 * sizeof(uint32_t), payload_bytes); + if (file_size != expected_file_size) { + throw std::runtime_error("Binary matrix payload size does not match its header"); } - input.close(); + M loaded(rows, cols); + io_impl::read_exact(input, loaded.data(), payload_bytes, "matrix data"); + + row_mat = std::move(loaded); } } // namespace rabitqlib diff --git a/include/rabitqlib/utils/memory.hpp b/include/rabitqlib/utils/memory.hpp index 93cd490..af11ee6 100644 --- a/include/rabitqlib/utils/memory.hpp +++ b/include/rabitqlib/utils/memory.hpp @@ -11,8 +11,6 @@ #include #include -#include "rabitqlib/utils/tools.hpp" - namespace rabitqlib::memory { #define PORTABLE_ALIGN32 __attribute__((aligned(32))) #define PORTABLE_ALIGN64 __attribute__((aligned(64))) @@ -21,13 +19,30 @@ template class AlignedAllocator { private: static_assert(Alignment >= alignof(T)); + static_assert((Alignment & (Alignment - 1)) == 0, "Alignment must be a power of two"); + + template + using ReboundAllocator = AlignedAllocator; + + [[nodiscard]] static constexpr size_t aligned_size(size_t nbytes) { + const size_t remainder = nbytes % Alignment; + if (remainder == 0) { + return nbytes; + } + const size_t padding = Alignment - remainder; + if (nbytes > std::numeric_limits::max() - padding) { + throw std::bad_array_new_length(); + } + return nbytes + padding; + } public: using value_type = T; + using is_always_equal = std::true_type; template struct rebind { - using other = AlignedAllocator; + using other = ReboundAllocator; }; constexpr AlignedAllocator() noexcept = default; @@ -35,15 +50,32 @@ class AlignedAllocator { constexpr AlignedAllocator(const AlignedAllocator&) noexcept = default; template - constexpr explicit AlignedAllocator(AlignedAllocator const&) noexcept {} + constexpr explicit AlignedAllocator(const ReboundAllocator&) noexcept {} + + friend constexpr bool + operator==(const AlignedAllocator&, const AlignedAllocator&) noexcept { + return true; + } + + friend constexpr bool + operator!=(const AlignedAllocator&, const AlignedAllocator&) noexcept { + return false; + } [[nodiscard]] T* allocate(std::size_t n) { if (n > std::numeric_limits::max() / sizeof(T)) { throw std::bad_array_new_length(); } - auto nbytes = round_up_to_multiple_of(n * sizeof(T), Alignment); + if (n == 0) { + return nullptr; + } + + const auto nbytes = aligned_size(n * sizeof(T)); auto* ptr = std::aligned_alloc(Alignment, nbytes); + if (ptr == nullptr) { + throw std::bad_alloc(); + } if (HugePage) { madvise(ptr, nbytes, MADV_HUGEPAGE); } @@ -78,8 +110,23 @@ struct Allocator { template inline T* align_allocate(size_t nbytes) { - auto size = round_up_to_multiple_of(nbytes, Alignment); + static_assert(Alignment >= alignof(T)); + static_assert((Alignment & (Alignment - 1)) == 0, "Alignment must be a power of two"); + + if (nbytes == 0) { + return nullptr; + } + + const size_t remainder = nbytes % Alignment; + const size_t padding = remainder == 0 ? 0 : Alignment - remainder; + if (nbytes > std::numeric_limits::max() - padding) { + throw std::bad_array_new_length(); + } + const size_t size = nbytes + padding; void* ptr = std::aligned_alloc(Alignment, size); + if (ptr == nullptr) { + throw std::bad_alloc(); + } if (HugePage) { madvise(ptr, size, MADV_HUGEPAGE); } diff --git a/include/rabitqlib/utils/rotator.hpp b/include/rabitqlib/utils/rotator.hpp index f0d6b27..4c401a2 100644 --- a/include/rabitqlib/utils/rotator.hpp +++ b/include/rabitqlib/utils/rotator.hpp @@ -7,6 +7,7 @@ #include #include #include +#include #include #include #include @@ -51,6 +52,9 @@ inline size_t padding_requirement(size_t dim, RotatorType type) { return dim; } if (type == RotatorType::FhtKacRotator) { + if (dim > std::numeric_limits::max() - 63) { + throw std::invalid_argument("Rotator dimension is too large to pad"); + } return round_up_to_multiple(dim, 64); } throw std::invalid_argument("Invalid rotator type in padding_requirement()"); @@ -273,10 +277,19 @@ template Rotator* choose_rotator( size_t dim, RotatorType type = RotatorType::FhtKacRotator, size_t padded_dim = 0 ) { + if (dim == 0) { + throw std::invalid_argument("Rotator dimension must be positive"); + } if (padded_dim == 0) { padded_dim = rotator_impl::padding_requirement(dim, type); } + if (padded_dim < dim) { + throw std::invalid_argument( + "Padded rotator dimension must not be smaller than input" + ); + } + if (padded_dim != rotator_impl::padding_requirement(padded_dim, type)) { throw std::invalid_argument("Invalid padded dimension for the rotator type"); } diff --git a/include/rabitqlib/utils/space.hpp b/include/rabitqlib/utils/space.hpp index 1891fec..9c8549b 100644 --- a/include/rabitqlib/utils/space.hpp +++ b/include/rabitqlib/utils/space.hpp @@ -144,24 +144,33 @@ inline T normalize_vec( return static_cast(dim) * value; } -// pack 0/1 data to usigned integer +// Pack 0/1 data into words without requiring the destination byte storage to +// contain aligned, lifetime-started T objects. template -inline void pack_binary( - const int* __restrict__ binary_code, T* __restrict__ compact_code, size_t length +inline void pack_binary_to_bytes( + const int* __restrict__ binary_code, uint8_t* __restrict__ compact_code, size_t length ) { constexpr size_t kTypeBits = sizeof(T) * 8; - auto* output = reinterpret_cast(compact_code); for (size_t i = 0; i < length; i += kTypeBits) { T cur = 0; for (size_t j = 0; j < kTypeBits; ++j) { cur |= (static_cast(binary_code[i + j]) << (kTypeBits - 1 - j)); } - std::memcpy(output, &cur, sizeof(cur)); - output += sizeof(cur); + std::memcpy(compact_code, &cur, sizeof(cur)); + compact_code += sizeof(cur); } } +// Backward-compatible typed-output overload. The implementation remains +// byte-oriented, so it does not dereference compact_code as T. +template +inline void pack_binary( + const int* __restrict__ binary_code, T* __restrict__ compact_code, size_t length +) { + pack_binary_to_bytes(binary_code, reinterpret_cast(compact_code), length); +} + template inline void data_range(const T* __restrict__ vec0, size_t dim, T& lo, T& hi) { ConstRowMajorArrayMap v0(vec0, 1, dim); @@ -324,6 +333,7 @@ void new_transpose_bin_512( const uint8_t* q, uint64_t* tq, size_t padded_dim, size_t b_query ); +float mask_ip_x0_q(const float* query, const uint8_t* data, size_t padded_dim); float mask_ip_x0_q(const float* query, const uint64_t* data, size_t padded_dim); inline float mask_ip_x0_q_old(const float* query, const uint64_t* data, size_t padded_dim) { diff --git a/include/rabitqlib/utils/warmup_space.hpp b/include/rabitqlib/utils/warmup_space.hpp index 2782e7b..6f65bea 100644 --- a/include/rabitqlib/utils/warmup_space.hpp +++ b/include/rabitqlib/utils/warmup_space.hpp @@ -5,6 +5,15 @@ namespace rabitqlib { +float warmup_ip_x0_q_512( + const uint8_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +); + float warmup_ip_x0_q_512( const uint64_t* data, const uint64_t* query, diff --git a/python_bindings/ivf_bindings.cpp b/python_bindings/ivf_bindings.cpp index 03b0517..9ab8f91 100644 --- a/python_bindings/ivf_bindings.cpp +++ b/python_bindings/ivf_bindings.cpp @@ -1,14 +1,24 @@ -#include +#include +#include +#include +#include +#include // IWYU pragma: keep; registers std::optional casters +#include #include +#include #include #include #include +#include #include #include #include "bindings_common.hpp" +#include "rabitqlib/defines.hpp" +#include "rabitqlib/index/ivf/initializer.hpp" #include "rabitqlib/index/ivf/ivf.hpp" +#include "rabitqlib/utils/rotator.hpp" namespace py = pybind11; diff --git a/python_bindings/symqg_bindings.cpp b/python_bindings/symqg_bindings.cpp index 68afbfa..9a54bf0 100644 --- a/python_bindings/symqg_bindings.cpp +++ b/python_bindings/symqg_bindings.cpp @@ -1,17 +1,27 @@ -#include +#include +#include +#include +#include +#include #include #include +#include +#include #include #include #include +#include #include #include #include #include "bindings_common.hpp" +#include "rabitqlib/defines.hpp" +#include "rabitqlib/fastscan/fastscan.hpp" #include "rabitqlib/index/symqg/qg.hpp" #include "rabitqlib/index/symqg/qg_builder.hpp" +#include "rabitqlib/utils/rotator.hpp" namespace py = pybind11; @@ -89,6 +99,9 @@ class SymqgIndex { if (ef == 0) { throw std::invalid_argument("ef must be positive"); } + if (ef < k) { + throw std::invalid_argument("ef must be at least k"); + } if (static_cast(query_array.shape(1)) != dim_) { throw std::invalid_argument("query dimension does not match index dim"); } @@ -102,8 +115,6 @@ class SymqgIndex { auto dists = py::array_t(shape); auto* ids_data = ids.mutable_data(); auto* dists_data = dists.mutable_data(); - std::fill(ids_data, ids_data + ids.size(), 0); - std::fill(dists_data, dists_data + dists.size(), 0.0F); const auto* queries_data = query_array.data(); const size_t requested_threads = diff --git a/scripts/check-tidy.sh b/scripts/check-tidy.sh index 15b7c85..cc7125d 100755 --- a/scripts/check-tidy.sh +++ b/scripts/check-tidy.sh @@ -50,12 +50,14 @@ if [[ -n "${CLANG_RESOURCE_DIR:-}" ]]; then extra_args+=("-extra-arg-before=-resource-dir=$CLANG_RESOURCE_DIR") fi for include_dir in "${system_include_dirs[@]}"; do + # Append compiler defaults so dependency paths from the compilation database + # take precedence over system installations (for example, pybind11). if [[ "$include_dir" == */lib/gcc/*/include ]]; then # Keep Clang's intrinsic headers ahead of GCC's, while still making # compiler-provided headers such as omp.h available. - extra_args+=("-extra-arg-before=-idirafter$include_dir") + extra_args+=("-extra-arg=-idirafter$include_dir") else - extra_args+=("-extra-arg-before=-isystem$include_dir") + extra_args+=("-extra-arg=-isystem$include_dir") fi done diff --git a/src/simd/dispatch.cpp b/src/simd/dispatch.cpp index 5abcef1..10f9e14 100644 --- a/src/simd/dispatch.cpp +++ b/src/simd/dispatch.cpp @@ -130,7 +130,7 @@ static void missing_new_transpose_bin_512(const uint8_t*, uint64_t*, size_t, siz missing_feature("new_transpose_bin_512"); } -static float missing_mask_ip_x0_q(const float*, const uint64_t*, size_t) { +static float missing_mask_ip_x0_q(const float*, const uint8_t*, size_t) { missing_feature("mask ip x0 q"); } @@ -149,7 +149,7 @@ static void missing_fastscan_accumulate_hacc( } static float missing_warmup_ip_x0_q_512( - const uint64_t*, const uint64_t*, float, float, size_t, size_t + const uint8_t*, const uint64_t*, float, float, size_t, size_t ) { missing_feature("warmup_ip_x0_q_512"); } @@ -341,12 +341,12 @@ const NewTransposeBin512Fn kNewTransposeBin512Fn = [] { } }(); -using MaskIpX0QFn = float (*)(const float*, const uint64_t*, size_t); +using MaskIpX0QFn = float (*)(const float*, const uint8_t*, size_t); const MaskIpX0QFn kMaskIpX0QFn = [] { if (cpu::has_avx512_core()) { - return simd::mask_ip_x0_q_avx512; + return static_cast(simd::mask_ip_x0_q_avx512); } else if (cpu::has_avx2()) { - return simd::mask_ip_x0_q_avx2; + return static_cast(simd::mask_ip_x0_q_avx2); } else { return simd::missing_mask_ip_x0_q; } @@ -412,10 +412,14 @@ void new_transpose_bin_512( kNewTransposeBin512Fn(q, tq, padded_dim, b_query); } -float mask_ip_x0_q(const float* query, const uint64_t* data, size_t padded_dim) { +float mask_ip_x0_q(const float* query, const uint8_t* data, size_t padded_dim) { return kMaskIpX0QFn(query, data, padded_dim); } +float mask_ip_x0_q(const float* query, const uint64_t* data, size_t padded_dim) { + return mask_ip_x0_q(query, reinterpret_cast(data), padded_dim); +} + } // namespace rabitqlib namespace rabitqlib::fastscan { @@ -480,10 +484,18 @@ void accumulate( uint16_t* __restrict__ result, size_t dim ) { + if (dim == 0 || dim % 16 != 0) { + throw std::invalid_argument("FastScan dimension must be a positive multiple of 16"); + } kAccumulateFn(codes, lp_table, result, dim); } void transfer_lut_hacc(const uint16_t* lut, size_t dim, uint8_t* hc_lut) { + if (dim == 0 || dim % 16 != 0) { + throw std::invalid_argument( + "high-accuracy FastScan dimension must be a positive multiple of 16" + ); + } kTransferLutHaccFn(lut, dim, hc_lut); } @@ -493,6 +505,11 @@ void accumulate_hacc( int32_t* accu_res, size_t dim ) { + if (dim == 0 || dim % 16 != 0) { + throw std::invalid_argument( + "high-accuracy FastScan dimension must be a positive multiple of 16" + ); + } kAccumulateHaccFn(codes, hc_lut, accu_res, dim); } @@ -501,19 +518,19 @@ void accumulate_hacc( namespace rabitqlib { using WarmupIpX0Q512Fn = - float (*)(const uint64_t*, const uint64_t*, float, float, size_t, size_t); + float (*)(const uint8_t*, const uint64_t*, float, float, size_t, size_t); const WarmupIpX0Q512Fn kWarmupIpX0Q512Fn = [] { if (rabitqlib::cpu::has_avx512_popcnt()) { - return rabitqlib::simd::warmup_ip_x0_q_512_avx512; + return static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx512); } else if (rabitqlib::cpu::has_avx2()) { - return rabitqlib::simd::warmup_ip_x0_q_512_avx2; + return static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx2); } else { return rabitqlib::simd::missing_warmup_ip_x0_q_512; } }(); float warmup_ip_x0_q_512( - const uint64_t* data, + const uint8_t* data, const uint64_t* query, float delta, float vl, @@ -523,4 +540,17 @@ float warmup_ip_x0_q_512( return kWarmupIpX0Q512Fn(data, query, delta, vl, padded_dim, b_query); } +float warmup_ip_x0_q_512( + const uint64_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +) { + return warmup_ip_x0_q_512( + reinterpret_cast(data), query, delta, vl, padded_dim, b_query + ); +} + } // namespace rabitqlib diff --git a/src/simd/fastscan_avx512.cpp b/src/simd/fastscan_avx512.cpp index bf01939..e019b78 100644 --- a/src/simd/fastscan_avx512.cpp +++ b/src/simd/fastscan_avx512.cpp @@ -110,8 +110,8 @@ void transfer_lut_hacc_avx512(const uint16_t* lut, size_t dim, uint8_t* hc_lut) ); __m128i lo = _mm512_cvtepi32_epi8(tmp); __m128i hi = _mm512_cvtepi32_epi8(_mm512_srli_epi32(tmp, 8)); - _mm_store_si128(reinterpret_cast<__m128i*>(fill_lo), lo); - _mm_store_si128(reinterpret_cast<__m128i*>(fill_hi), hi); + _mm_storeu_si128(reinterpret_cast<__m128i*>(fill_lo), lo); + _mm_storeu_si128(reinterpret_cast<__m128i*>(fill_hi), hi); lut += 16; } } diff --git a/src/simd/space_avx2.cpp b/src/simd/space_avx2.cpp index 45595b9..2a30a5b 100644 --- a/src/simd/space_avx2.cpp +++ b/src/simd/space_avx2.cpp @@ -171,9 +171,9 @@ void new_transpose_bin_512_avx2( } } -float mask_ip_x0_q_avx2(const float* query, const uint64_t* data, size_t padded_dim) { +float mask_ip_x0_q_avx2(const float* query, const uint8_t* data, size_t padded_dim) { const size_t num_blk = padded_dim / 64; - const auto* it_data = reinterpret_cast(data); + const auto* it_data = data; const float* it_query = query; const __m256i shifts0 = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); const __m256i shifts1 = _mm256_setr_epi32(8, 9, 10, 11, 12, 13, 14, 15); @@ -214,4 +214,8 @@ float mask_ip_x0_q_avx2(const float* query, const uint64_t* data, size_t padded_ return _mm_cvtss_f32(_mm_add_ss(lanes, _mm_movehdup_ps(lanes))); } +float mask_ip_x0_q_avx2(const float* query, const uint64_t* data, size_t padded_dim) { + return mask_ip_x0_q_avx2(query, reinterpret_cast(data), padded_dim); +} + } // namespace rabitqlib::simd diff --git a/src/simd/space_avx512.cpp b/src/simd/space_avx512.cpp index abdb78e..e466fde 100644 --- a/src/simd/space_avx512.cpp +++ b/src/simd/space_avx512.cpp @@ -127,9 +127,9 @@ void new_transpose_bin_512_avx512( } } -float mask_ip_x0_q_avx512(const float* query, const uint64_t* data, size_t padded_dim) { +float mask_ip_x0_q_avx512(const float* query, const uint8_t* data, size_t padded_dim) { const size_t num_blk = padded_dim / 64; - const uint8_t* it_data = reinterpret_cast(data); + const uint8_t* it_data = data; const float* it_query = query; // __m512 sum0 = _mm512_setzero_ps(); @@ -166,4 +166,8 @@ float mask_ip_x0_q_avx512(const float* query, const uint64_t* data, size_t padde return _mm512_reduce_add_ps(sum); } +float mask_ip_x0_q_avx512(const float* query, const uint64_t* data, size_t padded_dim) { + return mask_ip_x0_q_avx512(query, reinterpret_cast(data), padded_dim); +} + } // namespace rabitqlib::simd diff --git a/src/simd/warmup_avx2.cpp b/src/simd/warmup_avx2.cpp index b9d5fd8..e187744 100644 --- a/src/simd/warmup_avx2.cpp +++ b/src/simd/warmup_avx2.cpp @@ -2,6 +2,7 @@ #include #include +#include #include "rabitqlib/simd/warmup_dispatch.hpp" @@ -62,7 +63,7 @@ static inline __m256i popcount_avx2(__m256i v) { } float warmup_ip_x0_q_512_avx2( - const uint64_t* data, + const uint8_t* data, const uint64_t* query, float delta, float vl, @@ -88,8 +89,8 @@ float warmup_ip_x0_q_512_avx2( // Load 64 bytes of data using paired 32-byte loads __m256i data_vec_lo = _mm256_loadu_si256(reinterpret_cast(data)); __m256i data_vec_hi = - _mm256_loadu_si256(reinterpret_cast(data + 4)); - data += 8; // Advance 8 x 64-bit ints (64 bytes) + _mm256_loadu_si256(reinterpret_cast(data + 32)); + data += 64; acc_ppc = _mm256_add_epi64(acc_ppc, popcount_avx2(data_vec_lo)); acc_ppc = _mm256_add_epi64(acc_ppc, popcount_avx2(data_vec_hi)); @@ -114,37 +115,25 @@ float warmup_ip_x0_q_512_avx2( size_t remaining_dim = padded_dim - i; if (remaining_dim > 0) { size_t num_chunks_64 = remaining_dim / 64; - size_t num_chunks_32 = remaining_dim / 32; - - size_t chunks_lo = (num_chunks_32 > 8) ? 8 : num_chunks_32; - size_t chunks_hi = (num_chunks_32 > 8) ? (num_chunks_32 - 8) : 0; - - // 1. Create a baseline sequence register - __m256i sequence = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); - - // 2. Generate masks in-register using Greater-Than comparisons - // If chunks_lo is 3, limit will be [3,3,3,3,3,3,3,3]. - // 3 > seq results in [-1, -1, -1, 0, 0, 0, 0, 0], which is the exact mask needed. - __m256i limit_lo = _mm256_set1_epi32(static_cast(chunks_lo)); - __m256i mask_lo = _mm256_cmpgt_epi32(limit_lo, sequence); + size_t remaining_bytes = num_chunks_64 * sizeof(uint64_t); - __m256i limit_hi = _mm256_set1_epi32(static_cast(chunks_hi)); - __m256i mask_hi = _mm256_cmpgt_epi32(limit_hi, sequence); - - // 3. Vectorized execution continues with zero memory latency + alignas(32) uint8_t data_tail[64]{}; + std::memcpy(data_tail, data, remaining_bytes); __m256i data_vec_lo = - _mm256_maskload_epi32(reinterpret_cast(data), mask_lo); + _mm256_load_si256(reinterpret_cast(data_tail)); __m256i data_vec_hi = - _mm256_maskload_epi32(reinterpret_cast(data + 4), mask_hi); + _mm256_load_si256(reinterpret_cast(data_tail + 32)); acc_ppc = _mm256_add_epi64(acc_ppc, popcount_avx2(data_vec_lo)); acc_ppc = _mm256_add_epi64(acc_ppc, popcount_avx2(data_vec_hi)); for (size_t j = 0; j < b_query; ++j) { + alignas(32) uint8_t query_tail[64]{}; + std::memcpy(query_tail, query, remaining_bytes); __m256i query_vec_lo = - _mm256_maskload_epi32(reinterpret_cast(query), mask_lo); + _mm256_load_si256(reinterpret_cast(query_tail)); __m256i query_vec_hi = - _mm256_maskload_epi32(reinterpret_cast(query + 4), mask_hi); + _mm256_load_si256(reinterpret_cast(query_tail + 32)); query += num_chunks_64; __m256i pop_lo = popcount_avx2(_mm256_and_si256(data_vec_lo, query_vec_lo)); @@ -174,4 +163,17 @@ float warmup_ip_x0_q_512_avx2( return (delta * static_cast(ip_scalar)) + (vl * static_cast(ppc_scalar)); } +float warmup_ip_x0_q_512_avx2( + const uint64_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +) { + return warmup_ip_x0_q_512_avx2( + reinterpret_cast(data), query, delta, vl, padded_dim, b_query + ); +} + } // namespace rabitqlib::simd diff --git a/src/simd/warmup_avx512.cpp b/src/simd/warmup_avx512.cpp index 4333def..980cda5 100644 --- a/src/simd/warmup_avx512.cpp +++ b/src/simd/warmup_avx512.cpp @@ -8,7 +8,7 @@ namespace rabitqlib::simd { float warmup_ip_x0_q_512_avx512( - const uint64_t* data, + const uint8_t* data, const uint64_t* query, float delta, float vl, @@ -31,7 +31,7 @@ float warmup_ip_x0_q_512_avx512( for (; i < dim_end_512; i += 512) { __m512i data_vec = _mm512_loadu_si512(data); - data += 8; + data += 64; acc_ppc = _mm512_add_epi64(acc_ppc, _mm512_popcnt_epi64(data_vec)); @@ -72,4 +72,17 @@ float warmup_ip_x0_q_512_avx512( return (delta * static_cast(ip_scalar)) + (vl * static_cast(ppc_scalar)); } +float warmup_ip_x0_q_512_avx512( + const uint64_t* data, + const uint64_t* query, + float delta, + float vl, + size_t padded_dim, + size_t b_query +) { + return warmup_ip_x0_q_512_avx512( + reinterpret_cast(data), query, delta, vl, padded_dim, b_query + ); +} + } // namespace rabitqlib::simd diff --git a/src/utils/cpu_features.cpp b/src/utils/cpu_features.cpp index b4cf53a..920d10a 100644 --- a/src/utils/cpu_features.cpp +++ b/src/utils/cpu_features.cpp @@ -13,6 +13,12 @@ namespace rabitqlib::cpu { namespace { +#if defined(_M_X64) || defined(_M_IX86) || defined(__x86_64__) || defined(__i386__) +constexpr bool kIsX86 = true; +#else +constexpr bool kIsX86 = false; +#endif + #if defined(_MSC_VER) && (defined(_M_X64) || defined(_M_IX86)) void cpuid( uint32_t leaf, uint32_t subleaf, uint32_t* a, uint32_t* b, uint32_t* c, uint32_t* d @@ -34,32 +40,90 @@ void cpuid( void cpuid(uint32_t, uint32_t, uint32_t*, uint32_t*, uint32_t*, uint32_t*) {} #endif -Features detect_features() { - Features detected{}; +#if defined(_MSC_VER) && (defined(_M_X64) || defined(_M_IX86)) +uint64_t xgetbv(uint32_t index) { return _xgetbv(index); } +#elif defined(__x86_64__) || defined(__i386__) +uint64_t xgetbv(uint32_t index) { + uint32_t eax; + uint32_t edx; + __asm__ volatile("xgetbv" : "=a"(eax), "=d"(edx) : "c"(index)); + return (static_cast(edx) << 32) | eax; +} +#else +uint64_t xgetbv(uint32_t) { return 0; } +#endif -#if defined(__x86_64__) || defined(__i386__) - // leaf 1: ECX[28]=AVX, ECX[12]=FMA, ECX[23]=POPCNT - uint32_t eax, ebx, ecx, edx; +Features detect_features() { + Features hardware{}; + if constexpr (!kIsX86) { + return hardware; + } + + uint32_t eax = 0; + uint32_t ebx = 0; + uint32_t ecx = 0; + uint32_t edx = 0; + cpuid(0, 0, &eax, &ebx, &ecx, &edx); + const uint32_t max_leaf = eax; + if (max_leaf < 1) { + return hardware; + } + + // leaf 1: ECX[12]=FMA, ECX[27]=OSXSAVE, ECX[28]=AVX. cpuid(1, 0, &eax, &ebx, &ecx, &edx); + const bool fma = ((ecx >> 12) & 1U) != 0; + const bool osxsave = ((ecx >> 27) & 1U) != 0; + const bool avx = ((ecx >> 28) & 1U) != 0; + hardware.fma = fma; + + if (max_leaf >= 7) { + // leaf 7 (subleaf 0): EBX[5]=AVX2, EBX[16]=AVX512F, + // EBX[17]=AVX512DQ, EBX[30]=AVX512BW, + // ECX[14]=AVX512_VPOPCNTDQ. + uint32_t l7_eax = 0; + uint32_t l7_ebx = 0; + uint32_t l7_ecx = 0; + uint32_t l7_edx = 0; + cpuid(7, 0, &l7_eax, &l7_ebx, &l7_ecx, &l7_edx); + + hardware.avx2 = ((l7_ebx >> 5) & 1U) != 0; + hardware.avx512f = ((l7_ebx >> 16) & 1U) != 0; + hardware.avx512dq = ((l7_ebx >> 17) & 1U) != 0; + hardware.avx512bw = ((l7_ebx >> 30) & 1U) != 0; + hardware.avx512vpopcntdq = ((l7_ecx >> 14) & 1U) != 0; + } + + const uint64_t xcr0 = avx && osxsave ? xgetbv(0) : 0; + return detail::filter_usable_features(hardware, avx, osxsave, xcr0); +} - // leaf 7 (subleaf 0): EBX[5]=AVX2, EBX[16]=AVX512F, - // EBX[17]=AVX512DQ, EBX[30]=AVX512BW, - // ECX[14]=AVX512_VPOPCNTDQ - uint32_t l7_eax, l7_ebx, l7_ecx, l7_edx; - cpuid(7, 0, &l7_eax, &l7_ebx, &l7_ecx, &l7_edx); - - detected.fma = (ecx >> 12) & 1; - detected.avx2 = (l7_ebx >> 5) & 1; - detected.avx512f = (l7_ebx >> 16) & 1; - detected.avx512dq = (l7_ebx >> 17) & 1; - detected.avx512bw = (l7_ebx >> 30) & 1; - detected.avx512vpopcntdq = (l7_ecx >> 14) & 1; -#endif +} // namespace - return detected; +namespace detail { + +Features filter_usable_features( + const Features& hardware, bool avx, bool osxsave, uint64_t xcr0 +) { + constexpr uint64_t kAvxStateMask = (uint64_t{1} << 1) | (uint64_t{1} << 2); + constexpr uint64_t kAvx512StateMask = + kAvxStateMask | (uint64_t{1} << 5) | (uint64_t{1} << 6) | (uint64_t{1} << 7); + + const bool avx_state_enabled = + avx && osxsave && (xcr0 & kAvxStateMask) == kAvxStateMask; + const bool avx512_state_enabled = + avx_state_enabled && (xcr0 & kAvx512StateMask) == kAvx512StateMask; + + Features usable{}; + usable.fma = hardware.fma && avx_state_enabled; + usable.avx2 = hardware.avx2 && avx_state_enabled; + usable.avx512f = hardware.avx512f && avx512_state_enabled; + usable.avx512bw = hardware.avx512bw && avx512_state_enabled; + usable.avx512dq = hardware.avx512dq && avx512_state_enabled; + usable.avx512vpopcntdq = hardware.avx512vpopcntdq && avx512_state_enabled; + return usable; } -} // namespace +} // namespace detail const Features& features() { static const Features detected = detect_features(); @@ -73,7 +137,7 @@ bool has_avx2() { bool has_avx512_core() { const Features& detected = features(); - return detected.avx512f && detected.avx512bw && detected.avx512dq; + return detected.fma && detected.avx512f && detected.avx512bw && detected.avx512dq; } bool has_avx512_popcnt() { return has_avx512_core() && features().avx512vpopcntdq; } diff --git a/tests/python/test_ivf.py b/tests/python/test_ivf.py index 0922d74..e541935 100644 --- a/tests/python/test_ivf.py +++ b/tests/python/test_ivf.py @@ -115,6 +115,23 @@ def test_recall_vs_brute_force(built_ivf, base_data, query_data): assert r >= 0.5, f"Recall {r:.3f} too low" +def test_inner_product_routing_uses_inner_product(): + data = np.zeros((2, DIM), dtype=np.float32) + data[0, 0] = 0.9 + data[1, 0] = 100.0 + centroids = data.copy() + cluster_ids = np.array([0, 1], dtype=np.uint32) + query = np.zeros((1, DIM), dtype=np.float32) + query[0, 0] = 1.0 + + idx = IvfIndex(DIM, 2, 2, nbits=32, metric="ip") + idx.build(data, centroids, cluster_ids, num_threads=1) + ids, distances = idx.search(query, k=1, nprobe=1, num_threads=1) + + assert ids[0, 0] == 1 + np.testing.assert_allclose(distances[0, 0], -99.0, rtol=0, atol=1e-5) + + # ── optional parameters ─────────────────────────────────────────────────────── diff --git a/tests/python/test_symqg.py b/tests/python/test_symqg.py index 20ef6a5..d406202 100644 --- a/tests/python/test_symqg.py +++ b/tests/python/test_symqg.py @@ -1,5 +1,6 @@ """Tests for SymqgIndex: construction, search, properties, error handling, save/load.""" +import os import struct import numpy as np @@ -111,6 +112,11 @@ def test_search_before_build_raises(): idx.search(queries, k=1, ef=_EF) +def test_search_rejects_ef_smaller_than_k(built_symqg, query_data): + with pytest.raises(ValueError, match="ef must be at least k"): + built_symqg.search(query_data, k=10, ef=9) + + def test_invalid_degree_raises(): with pytest.raises(Exception): SymqgIndex(DIM, max_degree=16) @@ -121,6 +127,13 @@ def test_invalid_quantization_bits_raises(): SymqgIndex(DIM, max_degree=_MAX_DEGREE, quantization_bits=6) +def test_save_propagates_write_failure(built_symqg): + if not os.path.exists("/dev/full"): + pytest.skip("requires a failing output device") + with pytest.raises(RuntimeError): + built_symqg.save("/dev/full") + + @pytest.mark.parametrize("bits", [4, 8]) def test_qg_quant_build_search_and_roundtrip(bits, base_data, query_data, tmp_path): idx = SymqgIndex(DIM, max_degree=_MAX_DEGREE, quantization_bits=bits) diff --git a/tests/unit/rabitqlib/fastscan/fastscan_test.cpp b/tests/unit/rabitqlib/fastscan/fastscan_test.cpp index e176276..3ba14d6 100644 --- a/tests/unit/rabitqlib/fastscan/fastscan_test.cpp +++ b/tests/unit/rabitqlib/fastscan/fastscan_test.cpp @@ -2,12 +2,21 @@ #include +#include #include +#include #include #include #include +#include #include +#include "rabitqlib/defines.hpp" +#include "rabitqlib/fastscan/highacc_fastscan.hpp" +#include "rabitqlib/index/estimator.hpp" +#include "rabitqlib/index/lut.hpp" +#include "rabitqlib/index/query.hpp" +#include "rabitqlib/quantization/data_layout.hpp" #include "rabitqlib/simd/fastscan_dispatch.hpp" #include "rabitqlib/utils/cpu_features.hpp" @@ -166,5 +175,106 @@ TEST(FastScanPackingTest, AccumulatesReferenceLutValuesOnEverySupportedBackend) } } +TEST(FastScanHighAccuracyTest, RejectsDimensionsThatCannotFillASimdBlock) { + std::array lut{}; + std::array packed_lut{}; + std::array codes{}; + std::array low_result{}; + std::array result{}; + + for (size_t dim : {0U, 4U, 8U, 12U, 20U}) { + EXPECT_THROW( + accumulate(codes.data(), packed_lut.data(), low_result.data(), dim), + std::invalid_argument + ); + EXPECT_THROW( + transfer_lut_hacc(lut.data(), dim, packed_lut.data()), std::invalid_argument + ); + EXPECT_THROW( + accumulate_hacc(codes.data(), packed_lut.data(), result.data(), dim), + std::invalid_argument + ); + std::vector query(dim); + EXPECT_THROW((Lut(query.data(), dim, true)), std::invalid_argument); + EXPECT_THROW((Lut(query.data(), dim, false)), std::invalid_argument); + } +} + +TEST(FastScanHighAccuracyTest, Avx512TransferAcceptsUnalignedOutput) { + if (!cpu::has_avx512_core()) { + GTEST_SKIP() << "AVX512 is not supported on this CPU"; + } + + constexpr size_t kDim = 64; + std::array lut{}; + for (size_t i = 0; i < lut.size(); ++i) { + lut[i] = static_cast((i * 977U) & 0xffffU); + } + + std::array expected{}; + constexpr size_t kCodebooksPerRegister = 4; + constexpr size_t kBytesPerRegisterPair = 128; + constexpr size_t kBytesPerCodebook = 16; + constexpr size_t kHighByteOffset = 64; + for (size_t codebook = 0; codebook < kDim / 4; ++codebook) { + const size_t low_offset = + (codebook / kCodebooksPerRegister * kBytesPerRegisterPair) + + (codebook % kCodebooksPerRegister * kBytesPerCodebook); + for (size_t entry = 0; entry < 16; ++entry) { + const uint16_t value = lut[codebook * 16 + entry]; + expected[low_offset + entry] = static_cast(value); + expected[low_offset + kHighByteOffset + entry] = + static_cast(value >> 8); + } + } + + alignas(16) std::array storage{}; + ASSERT_NE(reinterpret_cast(storage.data() + 1) % 16, uintptr_t{0}); + simd::transfer_lut_hacc_avx512(lut.data(), kDim, storage.data() + 1); + EXPECT_TRUE(std::equal(expected.begin(), expected.end(), storage.begin() + 1)); +} + +TEST(FastScanHighAccuracyTest, AccumulatesAcrossChunks) { + if (!cpu::has_avx2()) { + GTEST_SKIP() << "AVX2 is not supported on this CPU"; + } + + constexpr size_t kDim = 4096; + std::vector query(kDim, 1.0F); + SplitBatchQuery q_obj(query.data(), kDim, 3, METRIC_L2, true); + + std::vector batch_data(BatchDataMap::data_bytes(kDim)); + BatchDataMap batch(batch_data.data(), kDim); + std::fill_n(batch.bin_code(), kDim * fastscan::kBatchSize / 8, uint8_t{0xff}); + std::fill_n(batch.f_add(), fastscan::kBatchSize, 0.0F); + std::fill_n(batch.f_rescale(), fastscan::kBatchSize, 1.0F); + std::fill_n(batch.f_error(), fastscan::kBatchSize, 0.0F); + + const int32_t scalar_accumulator = int32_t{65535} * static_cast(kDim / 4); + const float expected_ip = + q_obj.delta() * static_cast(scalar_accumulator) + q_obj.sum_vl_lut(); + const float expected_distance = expected_ip + q_obj.k1xsumq(); + + std::array estimated{}; + std::array lower{}; + std::array inner_products{}; + split_batch_estdist( + batch_data.data(), + q_obj, + kDim, + estimated.data(), + lower.data(), + inner_products.data(), + true + ); + + for (size_t lane = 0; lane < fastscan::kBatchSize; ++lane) { + EXPECT_TRUE(std::isfinite(estimated[lane])); + EXPECT_FLOAT_EQ(inner_products[lane], expected_ip); + EXPECT_FLOAT_EQ(estimated[lane], expected_distance); + EXPECT_FLOAT_EQ(lower[lane], expected_distance); + } +} + } // namespace } // namespace rabitqlib::fastscan diff --git a/tests/unit/rabitqlib/index/initializer_test.cpp b/tests/unit/rabitqlib/index/initializer_test.cpp new file mode 100644 index 0000000..8ce0b34 --- /dev/null +++ b/tests/unit/rabitqlib/index/initializer_test.cpp @@ -0,0 +1,45 @@ +#include "rabitqlib/index/ivf/initializer.hpp" + +#include + +#include +#include +#include + +namespace rabitqlib::ivf { +namespace { + +TEST(ParallelForTest, AutomaticThreadCountProcessesEveryItem) { + std::atomic calls{0}; + parallel_for(0, 100, 0, [&](size_t, size_t) { ++calls; }); + EXPECT_EQ(calls, 100U); +} + +TEST(ParallelForTest, EmptyRangeDoesNotInvokeFunction) { + size_t calls = 0; + parallel_for(5, 5, 0, [&](size_t, size_t) { ++calls; }); + EXPECT_EQ(calls, 0U); +} + +TEST(ParallelForTest, JoinsWorkersBeforeRethrowingFunctionException) { + std::atomic active{0}; + EXPECT_THROW( + parallel_for( + 0, + 100, + 4, + [&](size_t id, size_t) { + ++active; + --active; + if (id == 0) { + throw std::runtime_error("parallel failure"); + } + } + ), + std::runtime_error + ); + EXPECT_EQ(active, 0U); +} + +} // namespace +} // namespace rabitqlib::ivf diff --git a/tests/unit/rabitqlib/index/ivf_test.cpp b/tests/unit/rabitqlib/index/ivf_test.cpp index 6cf3ffd..d98dc63 100644 --- a/tests/unit/rabitqlib/index/ivf_test.cpp +++ b/tests/unit/rabitqlib/index/ivf_test.cpp @@ -2,9 +2,22 @@ #include +#include +#include +#include +#include #include +#include +#include +#include +#include #include #include +#include + +#include "rabitqlib/defines.hpp" +#include "rabitqlib/utils/buffer.hpp" +#include "rabitqlib/utils/rotator.hpp" namespace rabitqlib::ivf { namespace { @@ -16,6 +29,40 @@ TEST(IvfConfigurationTest, RejectsUnsupportedMetric) { ); } +TEST(IvfConfigurationTest, RejectsUnsupportedCountsAndNullConstructionInputs) { + constexpr size_t kDim = 64; + EXPECT_THROW((IVF(0, kDim, 1, 1)), std::invalid_argument); + EXPECT_THROW( + (IVF(buffer::kSearchBufferMaxPointCount + 1, kDim, 1, 1)), std::invalid_argument + ); + EXPECT_THROW((IVF(1, kDim, 0, 1)), std::invalid_argument); + EXPECT_THROW( + (IVF(1, kDim, buffer::kSearchBufferMaxPointCount + 1, 1)), std::invalid_argument + ); + + IVF index(1, kDim, 1, 1); + std::array vector{}; + const PID cluster = 0; + EXPECT_THROW( + index.construct(nullptr, vector.data(), &cluster, false, 1), std::invalid_argument + ); + EXPECT_THROW( + index.construct(vector.data(), nullptr, &cluster, false, 1), std::invalid_argument + ); + EXPECT_THROW( + index.construct(vector.data(), vector.data(), nullptr, false, 1), + std::invalid_argument + ); + + IVF unconfigured; + EXPECT_THROW( + unconfigured.construct(vector.data(), vector.data(), &cluster, false, 1), + std::logic_error + ); + const std::string path = ::testing::TempDir() + "rabitq_unconfigured_ivf.index"; + EXPECT_THROW(unconfigured.save(path.c_str()), std::logic_error); +} + TEST(IvfSearchTest, RawRerankingUsesOriginalCoordinates) { constexpr size_t kNum = 33; constexpr size_t kDim = 65; @@ -88,6 +135,108 @@ TEST(IvfSearchTest, AutomaticHighAccuracyMatchesBitPolicy) { } } +TEST(IvfSearchTest, RejectsUnbuiltStateAndInvalidArguments) { + constexpr size_t kDim = 64; + std::array query{}; + PID result = 0; + + IVF empty; + EXPECT_THROW(empty.search(query.data(), 1, 1, &result), std::logic_error); + + IVF index(1, kDim, 1, 1); + EXPECT_THROW(index.search(query.data(), 1, 1, &result), std::logic_error); + + std::array data; + data.fill(1.0F); + std::array centroid{}; + const PID cluster = 0; + index.construct(data.data(), centroid.data(), &cluster, false, 1); + + EXPECT_THROW(index.search(query.data(), 0, 1, &result), std::invalid_argument); + EXPECT_THROW(index.search(query.data(), 2, 1, &result), std::invalid_argument); + EXPECT_THROW( + index.search(query.data(), std::numeric_limits::max(), 1, &result), + std::invalid_argument + ); + EXPECT_THROW(index.search(query.data(), 1, 0, &result), std::invalid_argument); + EXPECT_THROW(index.search(nullptr, 1, 1, &result), std::invalid_argument); + EXPECT_THROW(index.search(query.data(), 1, 1, nullptr), std::invalid_argument); +} + +TEST(IvfSearchTest, FillsSentinelsWhenProbedClustersCannotFillK) { + constexpr size_t kNum = 2; + constexpr size_t kDim = 64; + std::array data{}; + std::array centroids{}; + std::array query{}; + std::array clusters{1, 1}; + std::fill(data.begin(), data.end(), 1.0F); + std::fill(centroids.begin() + static_cast(kDim), centroids.end(), 1.0F); + + IVF index(kNum, kDim, kNum, 32, MetricType::METRIC_L2, RotatorType::MatrixRotator); + index.construct(data.data(), centroids.data(), clusters.data(), false, 1); + + PID result = kPidMax; + float distance = -123.0F; + index.search(query.data(), 1, 1, &result, &distance); + EXPECT_EQ(result, kPidMax); + EXPECT_EQ(distance, std::numeric_limits::infinity()); +} + +TEST(IvfSearchTest, RoutesInnerProductQueriesByInnerProduct) { + constexpr size_t kNum = 2; + constexpr size_t kDim = 64; + std::vector data(kNum * kDim, 0.0F); + std::vector centroids(kNum * kDim, 0.0F); + std::array clusters{0, 1}; + std::array query{}; + data[0] = centroids[0] = 0.9F; + data[kDim] = centroids[kDim] = 100.0F; + query[0] = 1.0F; + + IVF index(kNum, kDim, kNum, 32, METRIC_IP); + index.construct(data.data(), centroids.data(), clusters.data(), false, 1); + + PID result = kPidMax; + float distance = 0.0F; + index.search(query.data(), 1, 1, &result, &distance); + EXPECT_EQ(result, 1U); + EXPECT_NEAR(distance, -99.0F, 1e-5F); +} + +TEST(IvfSearchTest, InnerProductRawRerankingUsesResidualNormForPruning) { + constexpr size_t kNum = 65; + constexpr size_t kDim = 64; + constexpr size_t kTopK = 1; + std::vector data(kNum * kDim); + std::vector centroids(2 * kDim, 0.0F); + std::vector clusters(kNum, 1); + std::array query{}; + query[0] = 1.0F; + centroids[0] = 100.0F; + centroids[kDim] = 50.0F; + + for (size_t id = 0; id < kNum; ++id) { + clusters[id] = id < 32 ? 0 : 1; + data[id * kDim] = id < 32 ? 10.0F + static_cast(id) * 0.01F : 0.0F; + for (size_t d = 1; d < kDim; ++d) { + data[id * kDim + d] = + std::sin(static_cast((id + 1) * (d + 3)) * 0.17F) * 20.0F; + } + } + data[32 * kDim] = 1000.0F; + + IVF index(kNum, kDim, 2, 32, METRIC_IP); + index.construct(data.data(), centroids.data(), clusters.data(), false, 1); + for (bool use_hacc : {false, true}) { + std::array ids{}; + std::array distances{}; + index.search(query.data(), kTopK, 2, ids.data(), distances.data(), use_hacc); + EXPECT_EQ(ids[0], 32U); + EXPECT_FLOAT_EQ(distances[0], -999.0F); + } +} + TEST(IvfPersistenceTest, ReloadsRawAndQuantizedStorage) { constexpr size_t kNum = 33; constexpr size_t kDim = 65; @@ -119,5 +268,94 @@ TEST(IvfPersistenceTest, ReloadsRawAndQuantizedStorage) { std::remove(path.c_str()); } +TEST(IvfPersistenceTest, FailedLoadPreservesExistingIndex) { + constexpr size_t kTargetNum = 2; + constexpr size_t kTargetDim = 64; + std::vector target_data(kTargetNum * kTargetDim, 1.0F); + std::vector target_centroid(kTargetDim, 0.0F); + std::array target_clusters{0, 0}; + std::fill( + target_data.begin() + static_cast(kTargetDim), target_data.end(), 2.0F + ); + const std::string path = ::testing::TempDir() + "rabitq_ivf_truncated.index"; + + IVF target(kTargetNum, kTargetDim, 1, 32, METRIC_IP, RotatorType::MatrixRotator); + target.construct( + target_data.data(), target_centroid.data(), target_clusters.data(), false, 1 + ); + std::array expected_ids{}; + std::array expected_distances{}; + target.search( + target_centroid.data(), + kTargetNum, + 1, + expected_ids.data(), + expected_distances.data() + ); + + constexpr size_t kSourceNum = 33; + constexpr size_t kSourceDim = 65; + std::vector source_data(kSourceNum * kSourceDim, 0.25F); + std::vector source_centroid(kSourceDim, 0.5F); + std::vector source_clusters(kSourceNum, 0); + IVF source(kSourceNum, kSourceDim, 1, 4, METRIC_L2, RotatorType::FhtKacRotator); + source.construct( + source_data.data(), source_centroid.data(), source_clusters.data(), false, 1 + ); + source.save(path.c_str()); + + std::ifstream input(path, std::ios::binary); + const std::vector bytes( + (std::istreambuf_iterator(input)), std::istreambuf_iterator() + ); + ASSERT_GT(bytes.size(), 1U); + input.close(); + std::ofstream output(path, std::ios::binary | std::ios::trunc); + output.write(bytes.data(), static_cast(bytes.size() - 1)); + output.close(); + + EXPECT_ANY_THROW(target.load(path.c_str())); + EXPECT_EQ(target.max_elements(), kTargetNum); + EXPECT_EQ(target.dimension(), kTargetDim); + EXPECT_EQ(target.num_clusters(), 1U); + EXPECT_EQ(target.nbits(), 32U); + EXPECT_EQ(target.metric_type(), METRIC_IP); + EXPECT_EQ(target.rotator_type(), RotatorType::MatrixRotator); + + std::array actual_ids{}; + std::array actual_distances{}; + target.search( + target_centroid.data(), kTargetNum, 1, actual_ids.data(), actual_distances.data() + ); + EXPECT_EQ(actual_ids, expected_ids); + EXPECT_EQ(actual_distances, expected_distances); + + std::remove(path.c_str()); +} + +TEST(IvfPersistenceTest, RejectsPointIdsUsingTheSearchBufferMarker) { + constexpr size_t kNum = 33; + constexpr size_t kDim = 64; + std::vector data(kNum * kDim, 0.25F); + std::vector centroid(kDim, 0.0F); + std::vector clusters(kNum, 0); + const std::string path = ::testing::TempDir() + "rabitq_ivf_bad_id.index"; + + IVF source(kNum, kDim, 1, 4); + source.construct(data.data(), centroid.data(), clusters.data(), false, 1); + source.save(path.c_str()); + + std::fstream file(path, std::ios::binary | std::ios::in | std::ios::out); + ASSERT_TRUE(file.is_open()); + file.seekp(-static_cast(sizeof(PID)), std::ios::end); + const PID marked_id = buffer::kSearchBufferCheckedMask; + file.write(reinterpret_cast(&marked_id), sizeof(marked_id)); + file.close(); + + IVF loaded; + EXPECT_THROW(loaded.load(path.c_str()), std::runtime_error); + std::remove(path.c_str()); +} + } // namespace } // namespace rabitqlib::ivf diff --git a/tests/unit/rabitqlib/index/qg_test.cpp b/tests/unit/rabitqlib/index/qg_test.cpp index 2fd012c..c47abe3 100644 --- a/tests/unit/rabitqlib/index/qg_test.cpp +++ b/tests/unit/rabitqlib/index/qg_test.cpp @@ -1,21 +1,74 @@ +#include "rabitqlib/index/symqg/qg.hpp" + #include #include #include #include +#include #include #include +#include +#include +#include +#include #include +#include #include #include #include +#include #include +#include "rabitqlib/defines.hpp" +#include "rabitqlib/fastscan/fastscan.hpp" +#include "rabitqlib/index/estimator.hpp" +#include "rabitqlib/index/lut.hpp" +#include "rabitqlib/index/query.hpp" +#include "rabitqlib/index/symqg/detail/pipnn.hpp" #include "rabitqlib/index/symqg/qg_builder.hpp" +#include "rabitqlib/quantization/data_layout.hpp" +#include "rabitqlib/quantization/rabitq.hpp" #include "rabitqlib/utils/cpu_features.hpp" +#include "rabitqlib/utils/rotator.hpp" +#include "rabitqlib/utils/space.hpp" namespace rabitqlib::symqg { struct QGConstructionTestAccess { + static void check_contiguous_rows(QuantizedGraph& graph) { + const size_t vector_bytes = + graph.is_quantized() + ? ExDataMap::data_bytes(graph.padded_dim_, graph.quantization_bits_) + : graph.dim_ * sizeof(float); + const size_t batch_bytes = QGBatchDataMap::data_bytes(graph.padded_dim_) * + (graph.degree_bound_ / fastscan::kBatchSize); + const size_t row_bytes = + vector_bytes + batch_bytes + (graph.degree_bound_ * sizeof(PID)); + const char* base = graph.get_row_data(0); + for (PID id = 0; id < graph.num_points_; ++id) { + const char* row = graph.get_row_data(id); + EXPECT_EQ(row, base + (id * row_bytes)); + EXPECT_EQ(graph.get_batch_data(id), row + vector_bytes); + if (graph.is_quantized()) { + EXPECT_EQ(graph.get_quantized_vector(id), row); + } else { + EXPECT_EQ(reinterpret_cast(graph.get_vector(id)), row); + graph.get_vector(id)[graph.dim_ - 1] = static_cast(id); + } + graph.get_neighbors(id)[0] = id; + PID stored_id = kPidMax; + std::memcpy(&stored_id, row + vector_bytes + batch_bytes, sizeof(stored_id)); + EXPECT_EQ(stored_id, id); + } + if (!graph.is_quantized()) { + for (PID id = 0; id < graph.num_points_; ++id) { + EXPECT_FLOAT_EQ( + graph.get_vector(id)[graph.dim_ - 1], static_cast(id) + ); + } + } + } + static QGBuilder from_graph( QuantizedGraph& graph, uint32_t ef, @@ -31,7 +84,9 @@ struct QGConstructionTestAccess { static auto codes(const QuantizedGraph& graph) { std::vector result; for (PID id = 0; id < graph.num_points_; ++id) { - const char* code = graph.get_quantized_vector(id); + const char* code = graph.is_quantized() + ? graph.get_quantized_vector(id) + : reinterpret_cast(graph.get_vector(id)); result.insert(result.end(), code, code + graph.batch_data_offset_); } return result; @@ -55,9 +110,37 @@ struct QGConstructionTestAccess { } } static void search(QGBuilder& builder) { builder.search_new_neighbors(false); } - static const PID* encoded_neighbors(const QuantizedGraph& graph, PID id) { + static auto encoded_neighbors(const QuantizedGraph& graph, PID id) { return graph.get_neighbors(id); } + static void set_encoded_neighbor( + QuantizedGraph& graph, PID source, size_t lane, PID target + ) { + graph.get_neighbors(source)[lane] = target; + } + static void copy_vectors( + QuantizedGraph& graph, const float* data, size_t threads + ) { + graph.copy_vectors(data, threads); + } + static void update( + QuantizedGraph& graph, + PID source, + const std::vector>& neighbors + ) { + graph.update_qg(source, neighbors); + } + static void fill_batch_factors(QuantizedGraph& graph, PID source, float value) { + QGBatchDataMap batch(graph.get_batch_data(source), graph.padded_dim_); + std::fill_n(batch.f_add(), fastscan::kBatchSize, value); + std::fill_n(batch.f_rescale(), fastscan::kBatchSize, value); + } + static std::pair batch_factors( + const QuantizedGraph& graph, PID source, size_t lane + ) { + ConstQGBatchDataMap batch(graph.get_batch_data(source), graph.padded_dim_); + return {batch.f_add()[lane], batch.f_rescale()[lane]}; + } static std::array estimate( const QuantizedGraph& graph, PID source, PID target ) { @@ -94,6 +177,19 @@ struct QGConstructionTestAccess { namespace { +TEST(QuantizedGraphLayoutTest, KeepsVectorsCodesAndNeighborsInContiguousRows) { + for (size_t bits : {0U, 4U, 8U}) { + SCOPED_TRACE(bits); + QuantizedGraph graph( + 33, 65, 32, METRIC_L2, RotatorType::FhtKacRotator, bits + ); + QGConstructionTestAccess::check_contiguous_rows(graph); + } +} + +static_assert(std::is_copy_constructible_v>); +static_assert(std::is_move_constructible_v>); + static_assert( !std::is_constructible_v< QGBuilder, @@ -264,8 +360,10 @@ TEST(QGConstructionTest, DefaultsToPipnnInitialization) { std::sort(expected.begin(), expected.end()); ASSERT_EQ(QGConstructionTestAccess::degrees(builder)[i], expected.size()); EXPECT_TRUE(QGConstructionTestAccess::neighbors(builder)[i].empty()); - const PID* actual = QGConstructionTestAccess::encoded_neighbors(graph, i); - EXPECT_TRUE(std::equal(expected.begin(), expected.end(), actual)); + const auto actual = QGConstructionTestAccess::encoded_neighbors(graph, i); + for (size_t j = 0; j < expected.size(); ++j) { + EXPECT_EQ(actual[j], expected[j]); + } } builder.build(); EXPECT_FLOAT_EQ(builder.avg_degree(), kDegree); @@ -295,6 +393,64 @@ TEST(QGConstructionTest, SupportsExplicitRandomInitialization) { ); } +TEST(QGConstructionTest, RejectsNullDataForEveryInitializationMode) { + QuantizedGraph graph(33, 64, 32); + EXPECT_THROW( + (QGBuilder(graph, 32, nullptr, 1, QGInitialization::PiPNN)), std::invalid_argument + ); + EXPECT_THROW( + (QGBuilder(graph, 32, nullptr, 1, QGInitialization::Random)), std::invalid_argument + ); +} + +TEST(QGConstructionTest, UsesExplicitThreadCountsWithoutChangingCallerState) { + if (!cpu::has_avx2()) { + GTEST_SKIP() << "FastScan requires AVX2/FMA"; + } + constexpr size_t kCount = 33, kDim = 64; + std::vector data(kCount * kDim, 0.25F); + const int caller_threads = omp_get_max_threads(); + + QuantizedGraph first(kCount, kDim, 32); + QuantizedGraph second(kCount, kDim, 32); + QGBuilder first_builder(first, 32, data.data(), 1, QGInitialization::Random); + EXPECT_EQ(omp_get_max_threads(), caller_threads); + QGBuilder second_builder(second, 32, data.data(), 2, QGInitialization::PiPNN); + EXPECT_EQ(omp_get_max_threads(), caller_threads); + + first_builder.build(2); + second_builder.build(2); + EXPECT_EQ(omp_get_max_threads(), caller_threads); +} + +TEST(QGConstructionTest, InitializesUnusedPartialBatchFactors) { + if (!cpu::has_avx2()) { + GTEST_SKIP() << "FastScan requires AVX2/FMA"; + } + constexpr size_t kCount = 33, kDim = 64; + std::vector data(kCount * kDim); + for (size_t i = 0; i < data.size(); ++i) { + data[i] = std::sin(static_cast(i) * 0.17F); + } + QuantizedGraph graph(kCount, kDim, 32); + QGConstructionTestAccess::copy_vectors(graph, data.data(), 1); + QGConstructionTestAccess::fill_batch_factors( + graph, 0, std::numeric_limits::quiet_NaN() + ); + std::vector> neighbors; + neighbors.emplace_back(1, 0.0F); + QGConstructionTestAccess::update(graph, 0, neighbors); + + const auto active = QGConstructionTestAccess::batch_factors(graph, 0, 0); + EXPECT_TRUE(std::isfinite(active.first)); + EXPECT_TRUE(std::isfinite(active.second)); + for (size_t lane = 1; lane < fastscan::kBatchSize; ++lane) { + const auto factors = QGConstructionTestAccess::batch_factors(graph, 0, lane); + EXPECT_FLOAT_EQ(factors.first, 0.0F); + EXPECT_FLOAT_EQ(factors.second, 0.0F); + } +} + TEST(QGConstructionTest, RefinesPartialSeedOnceAfterReleasingInputs) { if (!cpu::has_avx2()) { GTEST_SKIP() << "FastScan requires AVX2/FMA"; @@ -463,11 +619,145 @@ TEST(QuantizedGraphConfigurationTest, RejectsUnsupportedMetric) { ); } +TEST(QuantizedGraphConfigurationTest, RejectsZeroDimension) { + EXPECT_THROW( + (QuantizedGraph(33, 0, 32, METRIC_L2, RotatorType::MatrixRotator)), + std::invalid_argument + ); +} + +TEST(QuantizedGraphConfigurationTest, RejectsOutOfRangeEntryPoint) { + QuantizedGraph graph(33, 64, 32); + EXPECT_NO_THROW(graph.set_ep(32)); + EXPECT_THROW(graph.set_ep(33), std::invalid_argument); + EXPECT_EQ(graph.entry_point(), 32U); +} + +TEST(QuantizedGraphPersistenceTest, RejectsMalformedPayloadWithoutChangingTarget) { + constexpr size_t kNumPoints = 33; + constexpr size_t kDim = 64; + constexpr size_t kDegree = 32; + std::vector data(kNumPoints * kDim, 0.25F); + QuantizedGraph source( + kNumPoints, kDim, kDegree, METRIC_L2, RotatorType::MatrixRotator, 4 + ); + QGBuilder builder(source, kDegree, data.data(), 1); + builder.build(2); + + const std::string path = ::testing::TempDir() + "rabitq_qg_malformed.index"; + QuantizedGraph target(65, kDim, kDegree, METRIC_IP); + + source.save(path.c_str()); + std::filesystem::resize_file(path, std::filesystem::file_size(path) - 1); + EXPECT_THROW(target.load(path.c_str()), std::runtime_error); + EXPECT_EQ(target.num_vertices(), 65U); + EXPECT_EQ(target.metric_type(), METRIC_IP); + + source.save(path.c_str()); + { + std::fstream file(path, std::ios::binary | std::ios::in | std::ios::out); + ASSERT_TRUE(file.is_open()); + constexpr std::streamoff kPaddedDimensionOffset = + sizeof(uint64_t) + sizeof(uint32_t) + (3 * sizeof(size_t)); + file.seekp(kPaddedDimensionOffset); + const size_t invalid_padded_dim = kDim * 2; + file.write( + reinterpret_cast(&invalid_padded_dim), sizeof(invalid_padded_dim) + ); + ASSERT_TRUE(file.good()); + } + EXPECT_THROW(target.load(path.c_str()), std::runtime_error); + EXPECT_EQ(target.num_vertices(), 65U); + EXPECT_EQ(target.metric_type(), METRIC_IP); + + source.save(path.c_str()); + { + std::fstream file(path, std::ios::binary | std::ios::in | std::ios::out); + ASSERT_TRUE(file.is_open()); + constexpr std::streamoff kPointCountOffset = sizeof(uint64_t) + sizeof(uint32_t); + file.seekp(kPointCountOffset); + const size_t invalid_point_count = std::numeric_limits::max(); + file.write( + reinterpret_cast(&invalid_point_count), sizeof(invalid_point_count) + ); + ASSERT_TRUE(file.good()); + } + EXPECT_THROW(target.load(path.c_str()), std::invalid_argument); + EXPECT_EQ(target.num_vertices(), 65U); + EXPECT_EQ(target.metric_type(), METRIC_IP); + + if (std::filesystem::exists("/dev/full")) { + EXPECT_THROW(source.save("/dev/full"), std::ios_base::failure); + } + + QGConstructionTestAccess::set_encoded_neighbor(source, 0, 0, kNumPoints); + source.save(path.c_str()); + EXPECT_THROW(target.load(path.c_str()), std::runtime_error); + EXPECT_EQ(target.num_vertices(), 65U); + EXPECT_EQ(target.metric_type(), METRIC_IP); + + std::remove(path.c_str()); +} + TEST(QuantizedGraphLifecycleTest, DestroysConcreteRotatorThroughBasePointer) { QuantizedGraph graph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator); EXPECT_EQ(graph.num_vertices(), 33U); } +TEST(QuantizedGraphLifecycleTest, RejectsSearchAndSaveBeforeBuild) { + const std::string path = ::testing::TempDir() + "rabitq_qg_unbuilt.index"; + std::remove(path.c_str()); + std::array query{}; + std::array ids{}; + std::array distances{}; + + QuantizedGraph empty; + EXPECT_THROW(empty.save(path.c_str()), std::logic_error); + EXPECT_THROW( + empty.search(query.data(), 1, ids.data(), distances.data()), std::logic_error + ); + + QuantizedGraph configured(33, 64, 32); + configured.set_ef(1); + EXPECT_THROW(configured.save(path.c_str()), std::logic_error); + EXPECT_THROW( + configured.search(query.data(), 1, ids.data(), distances.data()), std::logic_error + ); + EXPECT_FALSE(std::filesystem::exists(path)); + std::remove(path.c_str()); +} + +TEST(QuantizedGraphLifecycleTest, BuilderDoesNotPublishPartialGraph) { + if (!cpu::has_avx2()) { + GTEST_SKIP() << "FastScan requires AVX2/FMA"; + } + constexpr size_t kNumPoints = 33; + constexpr size_t kDim = 64; + constexpr size_t kDegree = 32; + std::vector data(kNumPoints * kDim, 0.25F); + std::array query{}; + PID result = 0; + float distance = 0; + + for (const auto init : {QGInitialization::Random, QGInitialization::PiPNN}) { + SCOPED_TRACE(static_cast(init)); + const std::string path = ::testing::TempDir() + "rabitq_qg_partial.index"; + std::remove(path.c_str()); + QuantizedGraph graph(kNumPoints, kDim, kDegree); + QGBuilder builder(graph, kDegree, data.data(), 1, init); + graph.set_ef(1); + EXPECT_THROW(graph.save(path.c_str()), std::logic_error); + EXPECT_THROW(graph.search(query.data(), 1, &result, &distance), std::logic_error); + EXPECT_FALSE(std::filesystem::exists(path)); + + builder.build(2); + EXPECT_NO_THROW(graph.search(query.data(), 1, &result, &distance)); + EXPECT_NO_THROW(graph.save(path.c_str())); + EXPECT_TRUE(std::filesystem::exists(path)); + std::remove(path.c_str()); + } +} + TEST(QGBuilderMetricTest, UsesInnerProductDistanceToChooseEntryPoint) { constexpr size_t kNumPoints = 33; constexpr size_t kDim = 64; @@ -561,8 +851,8 @@ TEST(QGQuantTest, SearchesAndRoundTripsFourAndEightBitIndexes) { ::testing::TempDir() + "rabitq_qg_quant_" + std::to_string(bits) + ".index"; graph.save(path.c_str()); QuantizedGraph loaded; - loaded.load(path.c_str()); loaded.set_ef(kNumPoints); + loaded.load(path.c_str()); EXPECT_TRUE(loaded.is_quantized()); EXPECT_EQ(loaded.quantization_bits(), bits); @@ -577,6 +867,69 @@ TEST(QGQuantTest, SearchesAndRoundTripsFourAndEightBitIndexes) { } } +TEST(QGSearchTest, RejectsInvalidKAndEfInsteadOfReturningPartialResults) { + if (!cpu::has_avx2()) { + GTEST_SKIP() << "FastScan requires AVX2/FMA"; + } + constexpr size_t kNumPoints = 33, kDim = 64, kDegree = 32; + std::vector data(kNumPoints * kDim); + for (size_t i = 0; i < data.size(); ++i) { + data[i] = std::sin(static_cast(i) * 0.11F); + } + QuantizedGraph graph(kNumPoints, kDim, kDegree); + QGBuilder builder(graph, kDegree, data.data(), 1, QGInitialization::Random); + builder.build(2); + + std::array ids{}; + std::array distances{}; + EXPECT_THROW(graph.set_ef(0), std::invalid_argument); + graph.set_ef(4); + EXPECT_THROW( + graph.search(data.data(), 5, ids.data(), distances.data()), std::invalid_argument + ); + graph.set_ef(kNumPoints); + EXPECT_THROW( + graph.search(data.data(), 0, ids.data(), distances.data()), std::invalid_argument + ); + EXPECT_THROW( + graph.search( + data.data(), static_cast(kNumPoints + 1), ids.data(), distances.data() + ), + std::invalid_argument + ); + EXPECT_THROW( + graph.search(nullptr, 1, ids.data(), distances.data()), std::invalid_argument + ); + EXPECT_THROW( + graph.search(data.data(), 1, nullptr, distances.data()), std::invalid_argument + ); + EXPECT_THROW(graph.search(data.data(), 1, ids.data(), nullptr), std::invalid_argument); + + graph.search( + data.data(), static_cast(kNumPoints), ids.data(), distances.data() + ); + for (size_t i = 0; i < kNumPoints; ++i) { + EXPECT_LT(ids[i], kNumPoints); + EXPECT_TRUE(std::isfinite(distances[i])); + } + + for (PID source = 0; source < kNumPoints; ++source) { + for (size_t lane = 0; lane < kDegree; ++lane) { + QGConstructionTestAccess::set_encoded_neighbor(graph, source, lane, 0); + } + } + graph.set_ef(3); + ids.fill(kPidMax); + distances.fill(-123.0F); + EXPECT_THROW( + graph.search(data.data(), 3, ids.data(), distances.data()), std::runtime_error + ); + EXPECT_TRUE(std::all_of(ids.begin(), ids.end(), [](PID id) { return id == kPidMax; })); + EXPECT_TRUE(std::all_of(distances.begin(), distances.end(), [](float distance) { + return distance == -123.0F; + })); +} + TEST(QGSearchTest, ParallelQueriesMatchSerialAcrossIndexesAndSettings) { constexpr size_t kNumQueries = 9; constexpr size_t kTopK = 5; diff --git a/tests/unit/rabitqlib/quantization/data_layout_test.cpp b/tests/unit/rabitqlib/quantization/data_layout_test.cpp new file mode 100644 index 0000000..c7b6b04 --- /dev/null +++ b/tests/unit/rabitqlib/quantization/data_layout_test.cpp @@ -0,0 +1,184 @@ +#include "rabitqlib/quantization/data_layout.hpp" + +#include + +#include +#include +#include +#include +#include +#include + +#include "rabitqlib/fastscan/fastscan.hpp" +#include "rabitqlib/quantization/rabitq.hpp" + +namespace rabitqlib { +namespace { + +float load_float(const char* data) { + float value; + std::memcpy(&value, data, sizeof(value)); + return value; +} + +TEST(DataLayoutTest, BatchFactorsKeepPackedOffsetsOnUnalignedStorage) { + constexpr size_t kDim = 64; + constexpr size_t kCodeBytes = kDim * fastscan::kBatchSize / 8; + constexpr size_t kDataBytes = kCodeBytes + (3 * fastscan::kBatchSize * sizeof(float)); + const size_t data_bytes = BatchDataMap::data_bytes(kDim); + ASSERT_EQ(data_bytes, kCodeBytes + (3 * fastscan::kBatchSize * sizeof(float))); + + alignas(float) std::array storage; + storage.fill(char{0x5a}); + char* data = storage.data() + 1; + ASSERT_NE( + reinterpret_cast(data + kCodeBytes) % alignof(float), uintptr_t{0} + ); + BatchDataMap batch(data, kDim); + batch.f_add()[0] = 1.25F; + batch.f_rescale()[0] = -2.5F; + batch.f_error()[fastscan::kBatchSize - 1] = 3.75F; + + ConstBatchDataMap stored(data, kDim); + EXPECT_FLOAT_EQ(stored.f_add()[0], 1.25F); + EXPECT_FLOAT_EQ(stored.f_rescale()[0], -2.5F); + EXPECT_FLOAT_EQ(stored.f_error()[fastscan::kBatchSize - 1], 3.75F); + EXPECT_FLOAT_EQ(load_float(data + kCodeBytes), 1.25F); + EXPECT_FLOAT_EQ( + load_float(data + kCodeBytes + (fastscan::kBatchSize * sizeof(float))), -2.5F + ); + EXPECT_FLOAT_EQ( + load_float(data + kCodeBytes + ((3 * fastscan::kBatchSize - 1) * sizeof(float))), + 3.75F + ); + EXPECT_EQ(storage.front(), char{0x5a}); + EXPECT_EQ(storage.back(), char{0x5a}); +} + +TEST(DataLayoutTest, QgBatchFactorsKeepPackedOffsetsOnUnalignedStorage) { + constexpr size_t kDim = 64; + constexpr size_t kCodeBytes = kDim * fastscan::kBatchSize / 8; + constexpr size_t kDataBytes = kCodeBytes + (2 * fastscan::kBatchSize * sizeof(float)); + const size_t data_bytes = QGBatchDataMap::data_bytes(kDim); + ASSERT_EQ(data_bytes, kCodeBytes + (2 * fastscan::kBatchSize * sizeof(float))); + + alignas(float) std::array storage{}; + char* data = storage.data() + 1; + ASSERT_NE( + reinterpret_cast(data + kCodeBytes) % alignof(float), uintptr_t{0} + ); + QGBatchDataMap batch(data, kDim); + std::fill_n(batch.f_add(), fastscan::kBatchSize, 4.25F); + std::fill_n(batch.f_rescale(), fastscan::kBatchSize, -5.5F); + + ConstQGBatchDataMap stored(data, kDim); + EXPECT_FLOAT_EQ(stored.f_add()[fastscan::kBatchSize - 1], 4.25F); + EXPECT_FLOAT_EQ(stored.f_rescale()[fastscan::kBatchSize - 1], -5.5F); + EXPECT_FLOAT_EQ(load_float(data + kCodeBytes), 4.25F); + EXPECT_FLOAT_EQ( + load_float(data + kCodeBytes + (fastscan::kBatchSize * sizeof(float))), -5.5F + ); +} + +TEST(DataLayoutTest, SingleCodeFactorsKeepPackedOffsetsOnUnalignedStorage) { + constexpr size_t kDim = 64; + constexpr size_t kBits = 3; + constexpr size_t kCodeBytes = kDim * kBits / 8; + + constexpr size_t kExBytes = kCodeBytes + (2 * sizeof(float)); + alignas(float) std::array ex_storage{}; + char* ex_data = ex_storage.data() + 1; + ASSERT_NE( + reinterpret_cast(ex_data + kCodeBytes) % alignof(float), uintptr_t{0} + ); + ExDataMap ex(ex_data, kDim, kBits); + ex.f_add_ex() = 6.25F; + ex.f_rescale_ex() = -7.5F; + ConstExDataMap stored_ex(ex_data, kDim, kBits); + EXPECT_FLOAT_EQ(stored_ex.f_add_ex(), 6.25F); + EXPECT_FLOAT_EQ(stored_ex.f_rescale_ex(), -7.5F); + EXPECT_FLOAT_EQ(load_float(ex_data + kCodeBytes), 6.25F); + EXPECT_FLOAT_EQ(load_float(ex_data + kCodeBytes + sizeof(float)), -7.5F); + + constexpr size_t kBaseBytes = kCodeBytes + (3 * sizeof(float)); + alignas(float) std::array base_storage{}; + char* base_data = base_storage.data() + 1; + ASSERT_NE( + reinterpret_cast(base_data + kCodeBytes) % alignof(float), uintptr_t{0} + ); + BaseDataMap base(base_data, kDim, kBits); + base.f_add() = 8.25F; + base.f_rescale() = -9.5F; + base.f_error() = 10.75F; + ConstBaseDataMap stored_base(base_data, kDim, kBits); + EXPECT_FLOAT_EQ(stored_base.f_add(), 8.25F); + EXPECT_FLOAT_EQ(stored_base.f_rescale(), -9.5F); + EXPECT_FLOAT_EQ(stored_base.f_error(), 10.75F); + EXPECT_FLOAT_EQ(load_float(base_data + kCodeBytes), 8.25F); + EXPECT_FLOAT_EQ(load_float(base_data + kCodeBytes + sizeof(float)), -9.5F); + EXPECT_FLOAT_EQ(load_float(base_data + kCodeBytes + (2 * sizeof(float))), 10.75F); + + constexpr size_t kBinCodeBytes = kDim / 8; + constexpr size_t kBinBytes = kBinCodeBytes + (3 * sizeof(float)); + alignas(uint64_t) std::array bin_storage{}; + char* bin_data = bin_storage.data() + 1; + ASSERT_NE(reinterpret_cast(bin_data) % alignof(uint64_t), uintptr_t{0}); + ASSERT_NE( + reinterpret_cast(bin_data + kBinCodeBytes) % alignof(float), uintptr_t{0} + ); + BinDataMap bin(bin_data, kDim); + static_assert(std::is_same_v); + bin.bin_code()[0] = 0xa5; + bin.f_add() = 11.25F; + bin.f_rescale() = -12.5F; + bin.f_error() = 13.75F; + ConstBinDataMap stored_bin(bin_data, kDim); + static_assert(std::is_same_v); + EXPECT_EQ(stored_bin.bin_code()[0], uint8_t{0xa5}); + EXPECT_FLOAT_EQ(stored_bin.f_add(), 11.25F); + EXPECT_FLOAT_EQ(stored_bin.f_rescale(), -12.5F); + EXPECT_FLOAT_EQ(stored_bin.f_error(), 13.75F); + EXPECT_FLOAT_EQ(load_float(bin_data + kBinCodeBytes), 11.25F); + EXPECT_FLOAT_EQ(load_float(bin_data + kBinCodeBytes + sizeof(float)), -12.5F); + EXPECT_FLOAT_EQ(load_float(bin_data + kBinCodeBytes + (2 * sizeof(float))), 13.75F); +} + +TEST(DataLayoutTest, UnalignedSingleCodePreservesUint64PackedRepresentation) { + constexpr size_t kDim = 64; + constexpr size_t kCodeBytes = kDim / 8; + constexpr size_t kDataBytes = kCodeBytes + (3 * sizeof(float)); + std::array data{}; + std::array centroid{}; + for (size_t i = 0; i < kDim; ++i) { + data[i] = static_cast(static_cast(i % 7) - 3); + centroid[i] = static_cast(static_cast(i % 5) - 2) * 0.25F; + } + + uint64_t expected_code = 0; + float expected_add = 0; + float expected_rescale = 0; + float expected_error = 0; + quant::quantize_compact_one_bit( + data.data(), + centroid.data(), + kDim, + &expected_code, + expected_add, + expected_rescale, + expected_error + ); + + alignas(uint64_t) std::array storage{}; + char* packed = storage.data() + 1; + ASSERT_NE(reinterpret_cast(packed) % alignof(uint64_t), uintptr_t{0}); + quant::quantize_compact_one_bit(data.data(), centroid.data(), kDim, packed); + + ConstBinDataMap stored(packed, kDim); + EXPECT_EQ(std::memcmp(stored.bin_code(), &expected_code, sizeof(expected_code)), 0); + EXPECT_FLOAT_EQ(stored.f_add(), expected_add); + EXPECT_FLOAT_EQ(stored.f_rescale(), expected_rescale); + EXPECT_FLOAT_EQ(stored.f_error(), expected_error); +} + +} // namespace +} // namespace rabitqlib diff --git a/tests/unit/rabitqlib/utils/buffer_test.cpp b/tests/unit/rabitqlib/utils/buffer_test.cpp new file mode 100644 index 0000000..0222a5f --- /dev/null +++ b/tests/unit/rabitqlib/utils/buffer_test.cpp @@ -0,0 +1,39 @@ +#include "rabitqlib/utils/buffer.hpp" + +#include + +#include +#include +#include + +#include "rabitqlib/defines.hpp" + +namespace rabitqlib::buffer { +namespace { + +TEST(SearchBufferTest, RejectsCapacityWhoseSentinelWouldOverflow) { + EXPECT_THROW( + (SearchBuffer(std::numeric_limits::max())), std::length_error + ); + + SearchBuffer buffer; + EXPECT_TRUE(buffer.is_full()); + EXPECT_FALSE(buffer.has_next()); + EXPECT_EQ(buffer.size(), 0U); + buffer.insert(7, 1.0F); + EXPECT_EQ(buffer.size(), 0U); + EXPECT_THROW(buffer.resize(std::numeric_limits::max()), std::length_error); + + SearchBuffer populated(2); + populated.insert(7, 1.0F); + EXPECT_THROW(populated.resize(std::numeric_limits::max()), std::length_error); + EXPECT_EQ(populated.size(), 1U); + PID result = 0; + float distance = 0; + populated.copy_results(&result, &distance); + EXPECT_EQ(result, 7U); + EXPECT_FLOAT_EQ(distance, 1.0F); +} + +} // namespace +} // namespace rabitqlib::buffer diff --git a/tests/unit/rabitqlib/utils/cpu_features_test.cpp b/tests/unit/rabitqlib/utils/cpu_features_test.cpp index 974ce60..37829c0 100644 --- a/tests/unit/rabitqlib/utils/cpu_features_test.cpp +++ b/tests/unit/rabitqlib/utils/cpu_features_test.cpp @@ -2,7 +2,22 @@ #include -TEST(CpuFeatures, returns_stable_result) { +#include + +namespace { + +rabitqlib::cpu::Features all_hardware_features() { + return { + true, + true, + true, + true, + true, + true, + }; +} + +TEST(CpuFeatures, ReturnsStableResult) { const auto& first = rabitqlib::cpu::features(); const auto& second = rabitqlib::cpu::features(); @@ -12,3 +27,68 @@ TEST(CpuFeatures, returns_stable_result) { rabitqlib::cpu::has_avx512_core() && first.avx512vpopcntdq ); } + +TEST(CpuFeatures, RequiresOsManagedAvxState) { + constexpr uint64_t kAvxState = 0x6; + const auto hardware = all_hardware_features(); + + auto usable = + rabitqlib::cpu::detail::filter_usable_features(hardware, true, false, kAvxState); + EXPECT_FALSE(usable.avx2); + EXPECT_FALSE(usable.fma); + EXPECT_FALSE(usable.avx512f); + + usable = + rabitqlib::cpu::detail::filter_usable_features(hardware, false, true, kAvxState); + EXPECT_FALSE(usable.avx2); + EXPECT_FALSE(usable.fma); + EXPECT_FALSE(usable.avx512f); + + usable = rabitqlib::cpu::detail::filter_usable_features(hardware, true, true, 0); + EXPECT_FALSE(usable.avx2); + EXPECT_FALSE(usable.fma); + EXPECT_FALSE(usable.avx512f); +} + +TEST(CpuFeatures, EnablesAvx2WithoutAvx512State) { + constexpr uint64_t kAvxState = 0x6; + const auto usable = rabitqlib::cpu::detail::filter_usable_features( + all_hardware_features(), true, true, kAvxState + ); + + EXPECT_TRUE(usable.avx2); + EXPECT_TRUE(usable.fma); + EXPECT_FALSE(usable.avx512f); + EXPECT_FALSE(usable.avx512bw); + EXPECT_FALSE(usable.avx512dq); + EXPECT_FALSE(usable.avx512vpopcntdq); +} + +TEST(CpuFeatures, EnablesAvx512OnlyWithCompleteAvx512State) { + constexpr uint64_t kAvx512State = 0xe6; + const auto usable = rabitqlib::cpu::detail::filter_usable_features( + all_hardware_features(), true, true, kAvx512State + ); + + EXPECT_TRUE(usable.avx2); + EXPECT_TRUE(usable.fma); + EXPECT_TRUE(usable.avx512f); + EXPECT_TRUE(usable.avx512bw); + EXPECT_TRUE(usable.avx512dq); + EXPECT_TRUE(usable.avx512vpopcntdq); + + for (const uint64_t state_bit : + {uint64_t{1} << 5, uint64_t{1} << 6, uint64_t{1} << 7}) { + const auto incomplete = rabitqlib::cpu::detail::filter_usable_features( + all_hardware_features(), true, true, kAvx512State & ~state_bit + ); + EXPECT_TRUE(incomplete.avx2); + EXPECT_TRUE(incomplete.fma); + EXPECT_FALSE(incomplete.avx512f); + EXPECT_FALSE(incomplete.avx512bw); + EXPECT_FALSE(incomplete.avx512dq); + EXPECT_FALSE(incomplete.avx512vpopcntdq); + } +} + +} // namespace diff --git a/tests/unit/rabitqlib/utils/io_test.cpp b/tests/unit/rabitqlib/utils/io_test.cpp new file mode 100644 index 0000000..1aa3a9d --- /dev/null +++ b/tests/unit/rabitqlib/utils/io_test.cpp @@ -0,0 +1,110 @@ +#include "rabitqlib/utils/io.hpp" + +#include + +#include +#include +#include +#include +#include +#include + +#include "rabitqlib/defines.hpp" + +namespace rabitqlib { +namespace { + +template +void write_value(std::ofstream& output, const T& value) { + output.write(reinterpret_cast(&value), sizeof(value)); +} + +TEST(VectorIoTest, LoadsCompleteConsistentRecords) { + const std::string path = ::testing::TempDir() + "rabitq_valid.fvecs"; + { + std::ofstream output(path, std::ios::binary); + const uint32_t cols = 2; + const float values[] = {1.0F, 2.0F, 3.0F, 4.0F}; + write_value(output, cols); + output.write(reinterpret_cast(values), 2 * sizeof(float)); + write_value(output, cols); + output.write(reinterpret_cast(values + 2), 2 * sizeof(float)); + } + + RowMajorMatrix matrix; + load_vecs(path.c_str(), matrix); + ASSERT_EQ(matrix.rows(), 2); + ASSERT_EQ(matrix.cols(), 2); + EXPECT_FLOAT_EQ(matrix(0, 0), 1.0F); + EXPECT_FLOAT_EQ(matrix(1, 1), 4.0F); + std::remove(path.c_str()); +} + +TEST(VectorIoTest, RejectsInconsistentAndTruncatedRecordsWithoutReplacingOutput) { + const std::string path = ::testing::TempDir() + "rabitq_invalid.fvecs"; + RowMajorMatrix matrix(1, 1); + matrix(0, 0) = 7.0F; + + { + std::ofstream output(path, std::ios::binary); + const uint32_t cols = 2; + const uint32_t wrong_cols = 3; + const float values[] = {1.0F, 2.0F, 3.0F, 4.0F}; + write_value(output, cols); + output.write(reinterpret_cast(values), 2 * sizeof(float)); + write_value(output, wrong_cols); + output.write(reinterpret_cast(values + 2), 2 * sizeof(float)); + } + EXPECT_THROW(load_vecs(path.c_str(), matrix), std::runtime_error); + ASSERT_EQ(matrix.rows(), 1); + EXPECT_FLOAT_EQ(matrix(0, 0), 7.0F); + + { + std::ofstream output(path, std::ios::binary | std::ios::trunc); + const uint32_t cols = 2; + const float value = 1.0F; + write_value(output, cols); + write_value(output, value); + } + EXPECT_THROW(load_vecs(path.c_str(), matrix), std::runtime_error); + ASSERT_EQ(matrix.rows(), 1); + EXPECT_FLOAT_EQ(matrix(0, 0), 7.0F); + std::remove(path.c_str()); +} + +TEST(BinaryMatrixIoTest, ValidatesPayloadBeforeReplacingOutput) { + const std::string path = ::testing::TempDir() + "rabitq_matrix.bin"; + RowMajorMatrix matrix(1, 1); + matrix(0, 0) = 7.0F; + + { + std::ofstream output(path, std::ios::binary); + const uint32_t rows = 2; + const uint32_t cols = 2; + const float values[] = {1.0F, 2.0F, 3.0F}; + write_value(output, rows); + write_value(output, cols); + output.write(reinterpret_cast(values), sizeof(values)); + } + EXPECT_THROW(load_bin(path.c_str(), matrix), std::runtime_error); + ASSERT_EQ(matrix.rows(), 1); + EXPECT_FLOAT_EQ(matrix(0, 0), 7.0F); + + { + std::ofstream output(path, std::ios::binary | std::ios::trunc); + const uint32_t rows = 2; + const uint32_t cols = 2; + const float values[] = {1.0F, 2.0F, 3.0F, 4.0F}; + write_value(output, rows); + write_value(output, cols); + output.write(reinterpret_cast(values), sizeof(values)); + } + load_bin(path.c_str(), matrix); + ASSERT_EQ(matrix.rows(), 2); + ASSERT_EQ(matrix.cols(), 2); + EXPECT_FLOAT_EQ(matrix(1, 1), 4.0F); + std::remove(path.c_str()); +} + +} // namespace +} // namespace rabitqlib diff --git a/tests/unit/rabitqlib/utils/mask_ip_test.cpp b/tests/unit/rabitqlib/utils/mask_ip_test.cpp index 0ca5869..3df3ab9 100644 --- a/tests/unit/rabitqlib/utils/mask_ip_test.cpp +++ b/tests/unit/rabitqlib/utils/mask_ip_test.cpp @@ -16,10 +16,12 @@ TEST(MaskIpX0Q, BackendsMatchScalarAcrossBlocksAndAlignments) { if (!cpu::has_avx2()) { GTEST_SKIP() << "Binary dot product tests require AVX2/FMA"; } - using Function = float (*)(const float*, const uint64_t*, size_t); - std::vector functions{simd::mask_ip_x0_q_avx2, mask_ip_x0_q}; + using Function = float (*)(const float*, const uint8_t*, size_t); + std::vector functions{ + static_cast(simd::mask_ip_x0_q_avx2), + static_cast(mask_ip_x0_q)}; if (cpu::has_avx512_core()) { - functions.push_back(simd::mask_ip_x0_q_avx512); + functions.push_back(static_cast(simd::mask_ip_x0_q_avx512)); } for (size_t dim : {0, 64, 128, 192, 256, 448, 576, 1024, 4096}) { @@ -50,8 +52,7 @@ TEST(MaskIpX0Q, BackendsMatchScalarAcrossBlocksAndAlignments) { if (dim != 0) { std::memcpy(storage.data() + offset, words.data(), dim / 8); } - const auto* codes = - reinterpret_cast(storage.data() + offset); + const auto* codes = storage.data() + offset; for (auto function : functions) { const float result = function(query.data() + 1, codes, dim); EXPECT_NEAR(result, expected, 2e-6 * std::max(1.0, sum_abs)); diff --git a/tests/unit/rabitqlib/utils/memory_test.cpp b/tests/unit/rabitqlib/utils/memory_test.cpp new file mode 100644 index 0000000..4785c29 --- /dev/null +++ b/tests/unit/rabitqlib/utils/memory_test.cpp @@ -0,0 +1,38 @@ +#include "rabitqlib/utils/memory.hpp" + +#include + +#include +#include +#include +#include + +#include "rabitqlib/utils/array.hpp" + +namespace rabitqlib::memory { +namespace { + +TEST(AlignedAllocatorTest, RejectsSizeThatOverflowsAlignmentRounding) { + AlignedAllocator allocator; + EXPECT_THROW( + (void)allocator.allocate(std::numeric_limits::max()), + std::bad_array_new_length + ); +} + +TEST(AlignedAllocationTest, RejectsSizeThatOverflowsAlignmentRounding) { + EXPECT_THROW( + (align_allocate<64, char>(std::numeric_limits::max())), + std::bad_array_new_length + ); +} + +TEST(ArrayTest, RejectsDimensionProductOverflow) { + EXPECT_THROW( + (Array(std::vector{std::numeric_limits::max(), 2})), + std::bad_array_new_length + ); +} + +} // namespace +} // namespace rabitqlib::memory diff --git a/tests/unit/rabitqlib/utils/rotator_test.cpp b/tests/unit/rabitqlib/utils/rotator_test.cpp index ec2d88a..9662377 100644 --- a/tests/unit/rabitqlib/utils/rotator_test.cpp +++ b/tests/unit/rabitqlib/utils/rotator_test.cpp @@ -3,13 +3,14 @@ #include #include +#include +#include #include -#include #include +#include #include #include "test_data.hpp" -#include "test_helpers.hpp" using namespace rabitqlib; using namespace rabitq_test; @@ -41,6 +42,26 @@ TEST_F(RotatorTest, DefaultRotatorType) { EXPECT_GE(padded_dim, dim); } +TEST_F(RotatorTest, RejectsPaddedDimensionSmallerThanInput) { + EXPECT_THROW( + (choose_rotator(dim, RotatorType::FhtKacRotator, dim / 2)), + std::invalid_argument + ); + EXPECT_THROW( + (choose_rotator(dim, RotatorType::MatrixRotator, dim / 2)), + std::invalid_argument + ); +} + +TEST_F(RotatorTest, RejectsZeroDimension) { + EXPECT_THROW( + (choose_rotator(0, RotatorType::FhtKacRotator)), std::invalid_argument + ); + EXPECT_THROW( + (choose_rotator(0, RotatorType::MatrixRotator)), std::invalid_argument + ); +} + uint8_t bitreverse8(uint8_t x) { x = (((x & 0x55) << 1) | ((x & 0xAA) >> 1)); x = (((x & 0x33) << 2) | ((x & 0xCC) >> 2)); diff --git a/tests/unit/rabitqlib/utils/space_test.cpp b/tests/unit/rabitqlib/utils/space_test.cpp index 9044ed8..21b0ca8 100644 --- a/tests/unit/rabitqlib/utils/space_test.cpp +++ b/tests/unit/rabitqlib/utils/space_test.cpp @@ -4,12 +4,15 @@ #include #include +#include #include #include #include "rabitqlib/simd/pack_excode_dispatch.hpp" #include "rabitqlib/simd/space_dispatch.hpp" +#include "rabitqlib/simd/warmup_dispatch.hpp" #include "rabitqlib/utils/cpu_features.hpp" +#include "rabitqlib/utils/warmup_space.hpp" using namespace rabitqlib; @@ -21,9 +24,9 @@ TEST(PackBinary, SupportsUnalignedOutput) { } alignas(uint64_t) std::array storage{}; - auto* output = reinterpret_cast(storage.data() + 1); + auto* output = storage.data() + 1; ASSERT_NE(reinterpret_cast(output) % alignof(uint64_t), 0U); - pack_binary(binary_code.data(), output, dim); + pack_binary_to_bytes(binary_code.data(), output, dim); const std::array packed{ load_unaligned_u64(storage.data() + 1), @@ -48,9 +51,9 @@ TEST(MaskIpX0Q, SupportsUnalignedCodes) { } alignas(uint64_t) std::array storage{}; - auto* codes = reinterpret_cast(storage.data() + 1); + auto* codes = storage.data() + 1; ASSERT_NE(reinterpret_cast(codes) % alignof(uint64_t), 0U); - pack_binary(binary_code.data(), codes, dim); + pack_binary_to_bytes(binary_code.data(), codes, dim); if (cpu::has_avx2()) { EXPECT_FLOAT_EQ(simd::mask_ip_x0_q_avx2(query.data(), codes, dim), expected); @@ -61,6 +64,51 @@ TEST(MaskIpX0Q, SupportsUnalignedCodes) { EXPECT_FLOAT_EQ(mask_ip_x0_q(query.data(), codes, dim), expected); } +TEST(WarmupIpX0Q, SupportsUnalignedCodes) { + constexpr float delta = 0.5F; + constexpr float vl = -0.25F; + for (size_t dim : {64UL, 128UL, 192UL, 448UL, 512UL, 576UL}) { + SCOPED_TRACE(dim); + std::vector data_bits(dim); + std::vector query_bits(dim); + size_t data_popcount = 0; + size_t intersection_popcount = 0; + for (size_t i = 0; i < dim; ++i) { + data_bits[i] = (i % 3) == 0; + query_bits[i] = (i % 5) == 0; + data_popcount += data_bits[i] != 0; + intersection_popcount += data_bits[i] != 0 && query_bits[i] != 0; + } + + std::vector storage((dim / 8) + 1); + auto* codes = storage.data() + 1; + ASSERT_NE(reinterpret_cast(codes) % alignof(uint64_t), 0U); + pack_binary_to_bytes(data_bits.data(), codes, dim); + std::vector query(dim / 64); + pack_binary(query_bits.data(), query.data(), dim); + const float expected = delta * static_cast(intersection_popcount) + + vl * static_cast(data_popcount); + + if (cpu::has_avx2()) { + EXPECT_FLOAT_EQ( + simd::warmup_ip_x0_q_512_avx2(codes, query.data(), delta, vl, dim, 1), + expected + ); + } + if (cpu::has_avx512_popcnt()) { + EXPECT_FLOAT_EQ( + simd::warmup_ip_x0_q_512_avx512(codes, query.data(), delta, vl, dim, 1), + expected + ); + } + if (cpu::has_avx2()) { + EXPECT_FLOAT_EQ( + warmup_ip_x0_q_512(codes, query.data(), delta, vl, dim, 1), expected + ); + } + } +} + TEST(Select_IP_Func, returns_stable_function_pointer) { auto ip_func = select_excode_ipfunc(0); ASSERT_NE(ip_func, nullptr);