From 84c75965b16fe33e2c5e9a23fc748cd015d174fc Mon Sep 17 00:00:00 2001 From: gouyt13clear Date: Fri, 4 Sep 2026 21:17:28 +0800 Subject: [PATCH] feat: add quantized SymphonyQG vectors --- README.md | 10 +- docs/docs/index/qg.md | 21 +- include/rabitqlib/index/symqg/qg.hpp | 321 ++++++++++++++++-- include/rabitqlib/index/symqg/qg_builder.hpp | 25 +- include/rabitqlib/quantization/rabitq.hpp | 19 ++ python_bindings/symqg_bindings.cpp | 33 +- sample/cpp/symqg_indexing.cpp | 15 +- sample/python/symqg_indexing.py | 19 +- tests/python/test_symqg.py | 25 ++ tests/unit/rabitqlib/index/qg_test.cpp | 68 ++++ .../rabitqlib/quantization/rabitq_test.cpp | 32 ++ 11 files changed, 535 insertions(+), 53 deletions(-) diff --git a/README.md b/README.md index 363db89..2ed560e 100644 --- a/README.md +++ b/README.md @@ -26,6 +26,14 @@ +## News + +- **September 2026 — Quantized SymphonyQG:** SymphonyQG now supports optional + 4-bit and 8-bit RaBitQ vector storage. Select QG-quant with + `quantization_bits=4` or `quantization_bits=8`; vanilla raw-vector QG remains + the default. See the [SymphonyQG documentation](docs/docs/index/qg.md) for + details. + ## Install ```bash @@ -242,7 +250,7 @@ algorithm guidance is available in the [documentation](docs/docs/index.md). | **Quantizer** | Integrating RaBitQ into an existing system | Low-level 1-bit or multi-bit encoding and distance estimation. | | **IVF** | Memory-efficient partitioned search | Stores quantized codes without retaining the raw dataset. | | **HNSW** | Graph search with compact vectors | Adds graph links and searches directly from quantized codes. | -| **SymphonyQG** | Query speed when more memory is available | Retains raw vectors and stores per-neighborhood quantization data. | +| **SymphonyQG** | Fast graph search with a configurable memory/accuracy tradeoff | Uses raw vectors by default, or optional packed 4-bit/8-bit RaBitQ vectors, alongside per-neighborhood quantization data. | IVF and SymphonyQG use [FastScan](https://arxiv.org/abs/1704.07355) for batched estimates, while HNSW uses single-code AVX2 or AVX-512 kernels. diff --git a/docs/docs/index/qg.md b/docs/docs/index/qg.md index 17ff2ef..8ff9f2d 100644 --- a/docs/docs/index/qg.md +++ b/docs/docs/index/qg.md @@ -4,11 +4,23 @@ is a graph-based index originating from the [NGT library](https://github.com/yahoojapan/NGT). This implementation comes from the [SymphonyQG](https://dl.acm.org/doi/abs/10.1145/3709730) project. For -each vertex it stores the raw vector, a fixed-size neighbor list, and batched +each vertex vanilla QG stores the raw vector, a fixed-size neighbor list, and batched one-bit RaBitQ data for those neighbors. This layout uses more memory than the raw vectors alone, but lets graph traversal estimate a group of neighbor distances with FastScan while computing exact distances for visited vertices. +QG-quant replaces each raw vector with an independent packed 4- or 8-bit RaBitQ +code. All codes use the dataset's global centroid and are not combined with the +one-bit neighbor codes. Pass `quantization_bits` as `4` or `8` to select +QG-quant, or leave it at `0` for vanilla QG. + +During graph construction, candidate discovery uses a floating-point source +vector against the stored qg-quant candidate codes. Those estimated distances +are retained for candidate ordering and source-to-candidate pruning terms, while +candidate-to-candidate pruning comparisons and graph refinement use the available +raw build vectors. The raw vectors are not retained in the completed qg-quant +index. + Memory and performance depend on the dimension, degree, build window, and search window. See `sample/cpp/symqg_indexing.cpp` and `sample/cpp/symqg_querying.cpp` for complete programs. @@ -25,7 +37,8 @@ QuantizedGraph::QuantizedGraph( size_t dim, size_t max_deg, MetricType metric_type = METRIC_L2, - RotatorType rotator_type = RotatorType::FhtKacRotator + RotatorType rotator_type = RotatorType::FhtKacRotator, + size_t quantization_bits = 0 ); QGBuilder::QGBuilder( @@ -38,6 +51,7 @@ QGBuilder::QGBuilder( - **num**: Number of vertices (vectors) in the dataset. - **dim**: Dimension of the dataset. - **max_deg**: Degree bound of QG, must be a multiple of 32. +- **quantization_bits**: `0` for vanilla QG, or `4`/`8` for QG-quant. - **index**: Previously initialized QG. - **ef_build**: Search window size during indexing. - **data**: Pointer to the dataset, size of num * dim. @@ -72,6 +86,9 @@ Each indexed element is stored in the following layout. [Edges] ``` +For QG-quant, the first block becomes `[Packed 4/8-bit RaBitQ code + factors]`. +The index also stores one rotated global centroid shared by all rows. + `Batch data for QG` contains one-bit codes and estimator factors for the element's neighbors, organized in FastScan batches of 32. Consequently, `max_deg` must be a multiple of 32. diff --git a/include/rabitqlib/index/symqg/qg.hpp b/include/rabitqlib/index/symqg/qg.hpp index a3c26ce..31b50dd 100644 --- a/include/rabitqlib/index/symqg/qg.hpp +++ b/include/rabitqlib/index/symqg/qg.hpp @@ -10,8 +10,10 @@ #include #include #include -#include +#include +#include #include +#include #include #include "rabitqlib/defines.hpp" @@ -31,6 +33,32 @@ namespace rabitqlib::symqg { +template +class QuantizedQuery { + private: + const T* rotated_query_; + T k1xsumq_; + T g_add_; + + public: + QuantizedQuery( + const T* rotated_query, const T* centroid, size_t padded_dim, MetricType metric_type + ) + : rotated_query_(rotated_query) { + k1xsumq_ = + std::accumulate(rotated_query, rotated_query + padded_dim, static_cast(0)) / + -2; + + g_add_ = metric_type == METRIC_IP + ? -dot_product(rotated_query, centroid, padded_dim) + : euclidean_sqr(rotated_query, centroid, padded_dim); + } + + [[nodiscard]] const T* rotated_query() const { return rotated_query_; } + [[nodiscard]] T k1xsumq() const { return k1xsumq_; } + [[nodiscard]] T g_add() const { return g_add_; } +}; + template class QuantizedGraph { friend class QGBuilder; @@ -44,6 +72,10 @@ class QuantizedGraph { 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 + const T* build_data_ = nullptr; // non-owning, used only while QGBuilder runs + std::vector centroid_; // rotated global centroid for qg-quant + ex_ipfunc quantized_ip_func_ = nullptr; Array< char, @@ -52,13 +84,13 @@ class QuantizedGraph { char, 1 << 22, true>> - data_; // vectors + graph + quantization codes + factors + data_; // vectors/codes + graph quantization data + edges std::unique_ptr> rotator_; // data rotator std::unique_ptr visited_list_pool_ = nullptr; - // Position of different data in each row (RawData + QuantizationCodes + Factors + - // neighborIDs) Since we guarantee the degree for each vertex equals degree_bound - // (multiple of 32), we do not need to store the degree for each vertex + // Position of row data (raw vector or packed qg-quant vector), neighbor + // quantization data, and neighbor IDs. Since every degree equals degree_bound_ + // (a multiple of 32), the degree does not need to be stored per vertex. size_t batch_data_offset_ = 0; // offset of qg batch data size_t neighbor_offset_ = 0; // offset of neighbors size_t row_offset_ = 0; // length of entire row @@ -70,6 +102,8 @@ class QuantizedGraph { void copy_vectors(const T*); + void set_quantization_centroid(const T* centroid); + [[nodiscard]] T* get_vector(PID data_id) { return reinterpret_cast(&data_.at(row_offset_ * data_id)); } @@ -78,6 +112,26 @@ class QuantizedGraph { return reinterpret_cast(&data_.at(row_offset_ * data_id)); } + [[nodiscard]] const T* get_build_vector(PID data_id) const { + return build_data_ + (dim_ * data_id); + } + + [[nodiscard]] char* get_quantized_vector(PID data_id) { + return &data_.at(row_offset_ * data_id); + } + + [[nodiscard]] const char* get_quantized_vector(PID data_id) const { + return &data_.at(row_offset_ * data_id); + } + + void prepare_query(const T*, std::vector&, std::optional>&) const; + + T point_distance(const T*, const QuantizedQuery*, PID) const; + + T quantized_distance(const QuantizedQuery&, PID) const; + + 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_); } @@ -103,7 +157,8 @@ class QuantizedGraph { void update_qg(PID, const std::vector>&); - void update_results(buffer::SearchBuffer&, VisitedSet&, const T*); + void + update_results(buffer::SearchBuffer&, VisitedSet&, const T*, const QuantizedQuery*); void scan_neighbors( const BatchQuery&, PID, T*, buffer::SearchBuffer&, VisitedSet&, size_t @@ -115,7 +170,8 @@ class QuantizedGraph { size_t dim, size_t max_deg, MetricType metric_type = METRIC_L2, - RotatorType rotator_type = RotatorType::FhtKacRotator + RotatorType rotator_type = RotatorType::FhtKacRotator, + size_t quantization_bits = 0 ); explicit QuantizedGraph() = default; @@ -132,6 +188,10 @@ class QuantizedGraph { [[nodiscard]] auto metric_type() const { return this->metric_type_; } + [[nodiscard]] auto quantization_bits() const { return this->quantization_bits_; } + + [[nodiscard]] bool is_quantized() const { return quantization_bits_ != 0; } + void set_ep(PID entry) { this->entry_point_ = entry; }; void save(const char*) const; @@ -151,7 +211,12 @@ class QuantizedGraph { template inline QuantizedGraph::QuantizedGraph( - size_t num, size_t dim, size_t max_deg, MetricType metric_type, RotatorType rotator_type + size_t num, + size_t dim, + size_t max_deg, + MetricType metric_type, + RotatorType rotator_type, + size_t quantization_bits ) : num_points_(num) , degree_bound_(max_deg) @@ -159,7 +224,8 @@ inline QuantizedGraph::QuantizedGraph( , padded_dim_(dim) , raw_dist_func_((metric_type == METRIC_IP) ? dot_product_dis : euclidean_sqr) , metric_type_(metric_type) - , rotator_type_(rotator_type) { + , rotator_type_(rotator_type) + , quantization_bits_(quantization_bits) { validate_configuration(); initialize(); } @@ -181,10 +247,63 @@ inline void QuantizedGraph::validate_configuration() const { "QuantizedGraph point count exceeds the search-buffer ID limit" ); } + if (quantization_bits_ != 0 && quantization_bits_ != 4 && quantization_bits_ != 8) { + throw std::invalid_argument( + "QuantizedGraph quantization bits must be 0 (vanilla), 4, or 8" + ); + } + if (quantization_bits_ != 0 && !std::is_same_v) { + throw std::invalid_argument("QuantizedGraph qg-quant currently requires float data" + ); + } } template inline void QuantizedGraph::copy_vectors(const T* data) { + build_data_ = data; + if (quantization_bits_ != 0) { + if constexpr (!std::is_same_v) { + throw std::logic_error("qg-quant currently requires float data"); + } else { + if (centroid_.size() != padded_dim_) { + throw std::logic_error( + "qg-quant centroid must be set before copying vectors" + ); + } +#pragma omp parallel + { + std::vector rotated_data(padded_dim_); + std::vector quantized_data(padded_dim_); +#pragma omp for schedule(dynamic) + for (size_t i = 0; i < num_points_; ++i) { + rotator_->rotate(data + (dim_ * i), rotated_data.data()); + ExDataMap output( + get_quantized_vector(i), padded_dim_, quantization_bits_ + ); + T unused_f_error = 0; + quant::quantize_full_single( + rotated_data.data(), + centroid_.data(), + padded_dim_, + quantization_bits_, + quantized_data.data(), + output.f_add_ex(), + output.f_rescale_ex(), + unused_f_error, + metric_type_ + ); + quant::rabitq_impl::ex_bits::packing_rabitqplus_code( + quantized_data.data(), + output.ex_code(), + padded_dim_, + quantization_bits_ + ); + } + } + std::cout << "\tVectors quantized to " << quantization_bits_ << " bits\n"; + return; + } + } #pragma omp parallel for schedule(dynamic) for (size_t i = 0; i < num_points_; ++i) { const T* src = data + (dim_ * i); @@ -194,12 +313,30 @@ inline void QuantizedGraph::copy_vectors(const T* data) { std::cout << "\tVectors Copied\n"; } +template +inline void QuantizedGraph::set_quantization_centroid(const T* centroid) { + if (quantization_bits_ == 0) { + return; + } + centroid_.resize(padded_dim_); + rotator_->rotate(centroid, centroid_.data()); +} + template inline void QuantizedGraph::save(const char* filename) const { std::cout << "Saving quantized graph to " << filename << '\n'; std::ofstream output(filename, std::ios::binary); assert(output.is_open()); + constexpr uint64_t kFormatMagic = 0x5147524142495451ULL; // "QGRABITQ" + constexpr uint32_t kFormatVersion = 1; + if (quantization_bits_ != 0) { + output.write(reinterpret_cast(&kFormatMagic), sizeof(kFormatMagic)); + output.write( + reinterpret_cast(&kFormatVersion), sizeof(kFormatVersion) + ); + } + /* Basic variants */ output.write(reinterpret_cast(&num_points_), sizeof(size_t)); output.write(reinterpret_cast(°ree_bound_), sizeof(size_t)); @@ -208,6 +345,14 @@ inline void QuantizedGraph::save(const char* filename) const { output.write(reinterpret_cast(&entry_point_), sizeof(PID)); output.write(reinterpret_cast(&rotator_type_), sizeof(RotatorType)); output.write(reinterpret_cast(&metric_type_), sizeof(MetricType)); + if (quantization_bits_ != 0) { + output.write( + reinterpret_cast(&quantization_bits_), sizeof(quantization_bits_) + ); + output.write( + reinterpret_cast(centroid_.data()), padded_dim_ * sizeof(T) + ); + } /* Data */ data_.save(output); @@ -232,6 +377,22 @@ inline void QuantizedGraph::load(const char* filename) { std::ifstream input(filename, std::ios::binary); assert(input.is_open()); + constexpr uint64_t kFormatMagic = 0x5147524142495451ULL; // "QGRABITQ" + constexpr uint32_t kFormatVersion = 1; + uint64_t magic = 0; + input.read(reinterpret_cast(&magic), sizeof(magic)); + if (magic == kFormatMagic) { + uint32_t version = 0; + input.read(reinterpret_cast(&version), sizeof(version)); + if (version != kFormatVersion) { + throw std::runtime_error("Unsupported QuantizedGraph file version"); + } + } else { + // Files produced before qg-quant have no header and always contain raw vectors. + input.clear(); + input.seekg(0); + } + /* Basic variants */ input.read(reinterpret_cast(&num_points_), sizeof(size_t)); input.read(reinterpret_cast(°ree_bound_), sizeof(size_t)); @@ -240,12 +401,24 @@ inline void QuantizedGraph::load(const char* filename) { input.read(reinterpret_cast(&entry_point_), sizeof(PID)); input.read(reinterpret_cast(&rotator_type_), sizeof(RotatorType)); input.read(reinterpret_cast(&metric_type_), sizeof(MetricType)); + if (magic == kFormatMagic) { + input.read( + reinterpret_cast(&quantization_bits_), sizeof(quantization_bits_) + ); + } else { + quantization_bits_ = 0; + } raw_dist_func_ = (metric_type_ == METRIC_IP) ? dot_product_dis : euclidean_sqr; validate_configuration(); initialize(); + if (quantization_bits_ != 0) { + centroid_.resize(padded_dim_); + input.read(reinterpret_cast(centroid_.data()), padded_dim_ * sizeof(T)); + } + /* Data */ data_.load(input); @@ -273,9 +446,8 @@ inline void QuantizedGraph::search( T* __restrict__ dists ) { std::vector rotated_query(padded_dim_); - rotator_->rotate(query, rotated_query.data()); - - // init query + std::optional> quantized_query; + prepare_query(query, rotated_query, quantized_query); BatchQuery q_obj(rotated_query.data(), padded_dim_); buffer::SearchBuffer search_pool(ef_); @@ -294,7 +466,9 @@ inline void QuantizedGraph::search( } vis->set(cur_node); - q_obj.set_g_add(raw_dist_func_(query, get_vector(cur_node), dim_)); + q_obj.set_g_add( + point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_node) + ); scan_neighbors( q_obj, cur_node, est_dist.data(), search_pool, *vis, this->degree_bound_ @@ -302,13 +476,38 @@ inline void QuantizedGraph::search( res_pool.insert(cur_node, q_obj.g_add()); } - update_results(res_pool, *vis, query); + update_results(res_pool, *vis, query, quantized_query ? &*quantized_query : nullptr); visited_list_pool_->release_vis_list(vis); res_pool.copy_results(results, dists); } -// scan a data row (including data vec and quantization codes for its neighbors) -// Store estimated neighbor distances; the caller computes the current vertex exactly. +template +inline void QuantizedGraph::prepare_query( + const T* query, + std::vector& rotated_query, + std::optional>& quantized_query +) const { + rotator_->rotate(query, rotated_query.data()); + + if (quantization_bits_ != 0) { + quantized_query.emplace( + rotated_query.data(), centroid_.data(), padded_dim_, metric_type_ + ); + } +} + +template +inline T QuantizedGraph::point_distance( + const T* raw_query, const QuantizedQuery* quantized_query, PID data_id +) const { + if (quantized_query != nullptr) { + return quantized_distance(*quantized_query, data_id); + } + return raw_dist_func_(raw_query, get_vector(data_id), dim_); +} + +// Scan a data row and store estimated neighbor distances. The caller scores the current +// vertex from either its raw vector (vanilla QG) or its 4/8-bit code (qg-quant). template void QuantizedGraph::scan_neighbors( const BatchQuery& q_obj, @@ -341,7 +540,10 @@ void QuantizedGraph::scan_neighbors( template inline void QuantizedGraph::update_results( - buffer::SearchBuffer& result_pool, VisitedSet& vis, const T* query + buffer::SearchBuffer& result_pool, + VisitedSet& vis, + const T* query, + const QuantizedQuery* quantized_query ) { if (result_pool.is_full()) { return; @@ -355,7 +557,7 @@ inline void QuantizedGraph::update_results( if (!vis.get(cur_neighbor)) { vis.set(cur_neighbor); result_pool.insert( - cur_neighbor, raw_dist_func_(query, get_vector(cur_neighbor), dim_) + cur_neighbor, point_distance(query, quantized_query, cur_neighbor) ); } } @@ -377,7 +579,9 @@ inline void QuantizedGraph::initialize() { assert(padded_dim_ % 64 == 0); assert(padded_dim_ >= dim_); - this->batch_data_offset_ = dim_ * sizeof(T); // pos of packed code (aligned) + 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)); @@ -388,6 +592,61 @@ inline void QuantizedGraph::initialize() { ); visited_list_pool_ = std::make_unique(1, num_points_); + if (quantization_bits_ != 0) { + quantized_ip_func_ = select_excode_ipfunc(quantization_bits_); + } +} + +template +inline T QuantizedGraph::quantized_distance(const QuantizedQuery& query, PID data_id) + const { + if constexpr (!std::is_same_v) { + throw std::logic_error("qg-quant currently requires float data"); + } else { + ConstExDataMap data( + get_quantized_vector(data_id), padded_dim_, quantization_bits_ + ); + return quant::full_est_dist( + data.ex_code(), + query.rotated_query(), + quantized_ip_func_, + padded_dim_, + quantization_bits_, + data.f_add_ex(), + data.f_rescale_ex(), + query.g_add(), + query.k1xsumq() + ); + } +} + +template +inline void QuantizedGraph::reconstruct_quantized_vector(PID data_id, T* reconstructed) + const { + ConstExDataMap data(get_quantized_vector(data_id), padded_dim_, quantization_bits_); + std::vector quantized_data(padded_dim_); + if (quantization_bits_ == 8) { + std::copy(data.ex_code(), data.ex_code() + padded_dim_, quantized_data.begin()); + } else { + for (size_t i = 0; i < padded_dim_; i += 16) { + uint64_t packed = 0; + std::memcpy(&packed, data.ex_code() + (i / 2), sizeof(packed)); + for (size_t j = 0; j < 8; ++j) { + const uint8_t pair = static_cast(packed >> (j * 8)); + quantized_data[i + j] = pair & 0x0f; + quantized_data[i + 8 + j] = pair >> 4; + } + } + } + quant::reconstruct_full_vec( + quantized_data.data(), + centroid_.data(), + padded_dim_, + quantization_bits_, + data.f_rescale_ex(), + reconstructed, + metric_type_ + ); } // find candidate neighbors for cur_id, exclude the vertex itself @@ -399,19 +658,15 @@ inline void QuantizedGraph::find_candidates( VisitedSet& vis, const std::vector& degrees ) const { - const T* query = get_vector(cur_id); + const T* query = get_build_vector(cur_id); std::vector rotated_query(padded_dim_); - rotator_->rotate(query, rotated_query.data()); - - // init query + std::optional> quantized_query; + prepare_query(query, rotated_query, quantized_query); BatchQuery q_obj(rotated_query.data(), padded_dim_); // insert entry point to initialize search buffer buffer::SearchBuffer tmp_pool(search_ef); - tmp_pool.insert(this->entry_point_, 1e10); - memory::mem_prefetch_l1( - reinterpret_cast(get_vector(this->entry_point_)), 10 - ); + tmp_pool.insert(this->entry_point_, std::numeric_limits::max()); /* Current version of fast scan compute 32 distances */ std::vector est_dist(degree_bound_); // estimated distances @@ -422,7 +677,9 @@ inline void QuantizedGraph::find_candidates( } vis.set(cur_candi); auto cur_degree = degrees[cur_candi]; - q_obj.set_g_add(raw_dist_func_(query, get_vector(cur_candi), dim_)); + q_obj.set_g_add( + point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_candi) + ); scan_neighbors(q_obj, cur_candi, est_dist.data(), tmp_pool, vis, cur_degree); if (cur_candi != cur_id) { results.emplace_back(cur_candi, q_obj.g_add()); @@ -450,10 +707,14 @@ inline void QuantizedGraph::update_qg( std::vector rotated_data(cur_degree * padded_dim_); std::vector rotated_centroid(padded_dim_); for (size_t i = 0; i < cur_degree; ++i) { - const T* neighbor_vec = get_vector(new_neighbors[i].id); + const T* neighbor_vec = get_build_vector(new_neighbors[i].id); this->rotator_->rotate(neighbor_vec, &rotated_data[i * padded_dim_]); } - this->rotator_->rotate(get_vector(cur_id), rotated_centroid.data()); + if (quantization_bits_ == 0) { + this->rotator_->rotate(get_build_vector(cur_id), rotated_centroid.data()); + } else { + reconstruct_quantized_vector(cur_id, rotated_centroid.data()); + } // quantize batches for current vertex auto* batch_data = get_batch_data(cur_id); diff --git a/include/rabitqlib/index/symqg/qg_builder.hpp b/include/rabitqlib/index/symqg/qg_builder.hpp index 8b75bfa..c83e07f 100644 --- a/include/rabitqlib/index/symqg/qg_builder.hpp +++ b/include/rabitqlib/index/symqg/qg_builder.hpp @@ -75,18 +75,21 @@ class QGBuilder { std::vector centroid = compute_centroid(data, num_nodes_, dim_, num_threads_); + qg_.set_quantization_centroid(centroid.data()); + qg_.copy_vectors(data); + PID entry_point = exact_nn( data, centroid.data(), num_nodes_, dim_, num_threads_, qg_.raw_dist_func_ ); std::cout << "Setting entry_point to " << entry_point << '\n' << std::flush; - qg_.set_ep(entry_point); - qg_.copy_vectors(data); random_init(); } + ~QGBuilder() { qg_.build_data_ = nullptr; } + void build(size_t num_iter = 3) { if (num_iter < 2) { std::cerr << "The number of iterations for building QG must be at least 2\n"; @@ -139,7 +142,7 @@ inline void QGBuilder::add_pruned_edges( while (new_result.size() < degree_bound_ && start < pruned_list.size()) { const auto& cur = pruned_list[start]; bool occlude = false; - const float* cur_data = qg_.get_vector(cur.id); + const float* cur_data = qg_.get_build_vector(cur.id); float dik_sqr = cur.distance; if (nei_set.find(cur.id) != nei_set.end()) { @@ -151,7 +154,8 @@ inline void QGBuilder::add_pruned_edges( if (dij_sqr > dik_sqr) { break; } - float djk_sqr = qg_.raw_dist_func_(qg_.get_vector(nei.id), cur_data, dim_); + float djk_sqr = + qg_.raw_dist_func_(qg_.get_build_vector(nei.id), cur_data, dim_); float cosine = (dik_sqr + dij_sqr - djk_sqr) / (2 * std::sqrt(dij_sqr * dik_sqr)); if (cosine > threshold) { @@ -200,7 +204,7 @@ inline void QGBuilder::heuristic_prune( } pruned_results.emplace_back(pool[start]); // add current candidate to result - const float* data_j = qg_.get_vector(candidate_id); + const float* data_j = qg_.get_build_vector(candidate_id); // i : current vertex // j : neighbor added in this iter @@ -210,7 +214,7 @@ inline void QGBuilder::heuristic_prune( continue; } float dik = pool[k].distance; - auto djk = qg_.raw_dist_func_(data_j, qg_.get_vector(pool[k].id), dim_); + auto djk = qg_.raw_dist_func_(data_j, qg_.get_build_vector(pool[k].id), dim_); if (djk < dik) { if (refine && pruned_neighbors_[cur_id].size() < kMaxPrunedSize) { @@ -319,11 +323,12 @@ inline void QGBuilder::random_init() { } } - const float* cur_data = qg_.get_vector(i); + const float* cur_data = qg_.get_build_vector(i); new_neighbors_[i].reserve(degree_bound_); for (PID cur_neigh : neighbor_set) { new_neighbors_[i].emplace_back( - cur_neigh, qg_.raw_dist_func_(cur_data, qg_.get_vector(cur_neigh), dim_) + cur_neigh, + qg_.raw_dist_func_(cur_data, qg_.get_build_vector(cur_neigh), dim_) ); } @@ -385,7 +390,9 @@ inline void QGBuilder::graph_refine() { if (rand_id != static_cast(i) && ids.find(rand_id) == ids.end()) { new_result.emplace_back( rand_id, - qg_.raw_dist_func_(qg_.get_vector(rand_id), qg_.get_vector(i), dim_) + qg_.raw_dist_func_( + qg_.get_build_vector(rand_id), qg_.get_build_vector(i), dim_ + ) ); ids.emplace(rand_id); } diff --git a/include/rabitqlib/quantization/rabitq.hpp b/include/rabitqlib/quantization/rabitq.hpp index 229b184..8ef4e23 100644 --- a/include/rabitqlib/quantization/rabitq.hpp +++ b/include/rabitqlib/quantization/rabitq.hpp @@ -393,6 +393,25 @@ inline void reconstruct_vec( (ConstRowMajorArrayMap(quantized_vec, 1, dim).template cast() * delta) + vl; } +template +inline void reconstruct_full_vec( + const TP* quantized_vec, + const T* centroid, + size_t dim, + size_t bits, + T f_rescale, + T* results, + MetricType metric_type = METRIC_L2 +) { + const T code_center = -static_cast((static_cast(1) << bits) - 1) / 2; + const T scale = metric_type == METRIC_L2 ? -f_rescale / 2 : -f_rescale; + RowMajorArrayMap result_arr(results, 1, dim); + result_arr = + ConstRowMajorArrayMap(centroid, 1, dim) + + scale * (ConstRowMajorArrayMap(quantized_vec, 1, dim).template cast() + + code_center); +} + template inline TF full_est_dist( const TI* quantized_vec, diff --git a/python_bindings/symqg_bindings.cpp b/python_bindings/symqg_bindings.cpp index 768bf2a..dcd3398 100644 --- a/python_bindings/symqg_bindings.cpp +++ b/python_bindings/symqg_bindings.cpp @@ -16,11 +16,22 @@ namespace rabitqlib::python_bindings { class SymqgIndex { public: - SymqgIndex(size_t dim, size_t max_degree, const std::string& metric = "l2") - : dim_(dim), max_degree_(max_degree), metric_(metric_from_string(metric)) { + SymqgIndex( + size_t dim, + size_t max_degree, + const std::string& metric = "l2", + size_t quantization_bits = 0 + ) + : dim_(dim) + , max_degree_(max_degree) + , metric_(metric_from_string(metric)) + , quantization_bits_(quantization_bits) { if (max_degree == 0 || max_degree % rabitqlib::fastscan::kBatchSize != 0) { throw std::invalid_argument("max_degree must be a positive multiple of 32"); } + if (quantization_bits != 0 && quantization_bits != 4 && quantization_bits != 8) { + throw std::invalid_argument("quantization_bits must be 0, 4, or 8"); + } } void build(py::handle data, size_t ef_construction, size_t num_threads = 1) { @@ -31,7 +42,12 @@ class SymqgIndex { num_points_ = static_cast(data_array.shape(0)); index_ = std::make_unique>( - num_points_, dim_, max_degree_, metric_, rabitqlib::RotatorType::FhtKacRotator + num_points_, + dim_, + max_degree_, + metric_, + rabitqlib::RotatorType::FhtKacRotator, + quantization_bits_ ); rabitqlib::symqg::QGBuilder builder( @@ -100,6 +116,7 @@ class SymqgIndex { wrapper.dim_ = wrapper.index_->dimension(); wrapper.max_degree_ = wrapper.index_->degree_bound(); wrapper.metric_ = wrapper.index_->metric_type(); + wrapper.quantization_bits_ = wrapper.index_->quantization_bits(); wrapper.built_ = true; return wrapper; } @@ -109,6 +126,7 @@ class SymqgIndex { [[nodiscard]] size_t num_points() const { return num_points_; } [[nodiscard]] bool is_built() const { return built_; } [[nodiscard]] std::string metric() const { return metric_to_string(metric_); } + [[nodiscard]] size_t quantization_bits() const { return quantization_bits_; } private: SymqgIndex() = default; @@ -117,6 +135,7 @@ class SymqgIndex { size_t max_degree_ = 0; size_t num_points_ = 0; rabitqlib::MetricType metric_ = rabitqlib::METRIC_L2; + size_t quantization_bits_ = 0; bool built_ = false; std::unique_ptr> index_; }; @@ -129,10 +148,11 @@ void register_symqg(py::module_& m) { py::class_(m, "SymqgIndex") .def( - py::init(), + py::init(), py::arg("dim"), py::arg("max_degree"), - py::arg("metric") = "l2" + py::arg("metric") = "l2", + py::arg("quantization_bits") = 0 ) .def( "build", @@ -155,5 +175,6 @@ void register_symqg(py::module_& m) { .def_property_readonly("max_degree", &SymqgIndex::max_degree) .def_property_readonly("num_points", &SymqgIndex::num_points) .def_property_readonly("is_built", &SymqgIndex::is_built) - .def_property_readonly("metric", &SymqgIndex::metric); + .def_property_readonly("metric", &SymqgIndex::metric) + .def_property_readonly("quantization_bits", &SymqgIndex::quantization_bits); } diff --git a/sample/cpp/symqg_indexing.cpp b/sample/cpp/symqg_indexing.cpp index c15b5c4..f511e69 100644 --- a/sample/cpp/symqg_indexing.cpp +++ b/sample/cpp/symqg_indexing.cpp @@ -18,7 +18,8 @@ int main(int argc, char** argv) { << "arg2: degree bound for symqg, must be a multiple of 32\n" << "arg3: ef for indexing \n" << "arg4: path for saving index\n" - << "arg5: metric type (\"l2\" or \"ip\"), l2 by default\n"; + << "arg5: metric type (\"l2\" or \"ip\"), l2 by default\n" + << "arg6: vector quantization bits (0, 4, or 8), 0 by default\n"; exit(1); } @@ -39,6 +40,7 @@ int main(int argc, char** argv) { } else if (metric_type == rabitqlib::METRIC_L2) { std::cout << "Metric Type: L2\n"; } + size_t quantization_bits = argc > 6 ? static_cast(atoi(argv[6])) : 0; data_type data; @@ -46,7 +48,14 @@ int main(int argc, char** argv) { rabitqlib::StopW stopw; - index_type qg(data.rows(), data.cols(), degree, metric_type); + index_type qg( + data.rows(), + data.cols(), + degree, + metric_type, + rabitqlib::RotatorType::FhtKacRotator, + quantization_bits + ); rabitqlib::symqg::QGBuilder builder(qg, ef, data.data()); @@ -60,4 +69,4 @@ int main(int argc, char** argv) { qg.save(index_file); return 0; -} \ No newline at end of file +} diff --git a/sample/python/symqg_indexing.py b/sample/python/symqg_indexing.py index 7bdf692..4c2db6b 100644 --- a/sample/python/symqg_indexing.py +++ b/sample/python/symqg_indexing.py @@ -11,6 +11,7 @@ MAX_DEGREE = 32 # degree bound for SymphonyQG EF_CONSTRUCTION = 200 # ef for indexing METRIC = "l2" # "l2" or "ip" +QUANTIZATION_BITS = 0 # 0 for vanilla QG, or 4/8 for QG-quant NUM_THREADS = 16 # number of threads for build # ────────────────────────────────────────────── @@ -24,10 +25,16 @@ def main(args=None) -> None: n, dim = data.shape print( f"\nBuilding SymphonyQG index: n={n}, dim={dim}, MaxDegree={args.max_degree}, " - f"ef={args.ef_construction}, metric={args.metric}" + f"ef={args.ef_construction}, metric={args.metric}, " + f"quantization_bits={args.quantization_bits}" ) - idx = SymqgIndex(dim=dim, max_degree=args.max_degree, metric=args.metric) + idx = SymqgIndex( + dim=dim, + max_degree=args.max_degree, + metric=args.metric, + quantization_bits=args.quantization_bits, + ) t0 = time() idx.build(data, ef_construction=args.ef_construction, num_threads=args.num_threads) @@ -65,6 +72,14 @@ def main(args=None) -> None: choices=["l2", "ip"], help="Distance metric (l2 or ip)", ) + parser.add_argument( + "--quantization-bits", + dest="quantization_bits", + type=int, + choices=[0, 4, 8], + default=QUANTIZATION_BITS, + help="Vector quantization bits: 0 for vanilla QG, or 4/8 for QG-quant", + ) parser.add_argument( "--num-threads", dest="num_threads", diff --git a/tests/python/test_symqg.py b/tests/python/test_symqg.py index ff938ae..566dca7 100644 --- a/tests/python/test_symqg.py +++ b/tests/python/test_symqg.py @@ -33,6 +33,7 @@ def test_properties(built_symqg): assert built_symqg.max_degree == _MAX_DEGREE assert built_symqg.num_points == N_VECTORS assert built_symqg.metric == "l2" + assert built_symqg.quantization_bits == 0 # ── search output shape and dtype ───────────────────────────────────────────── @@ -113,6 +114,30 @@ def test_invalid_degree_raises(): SymqgIndex(DIM, max_degree=16) +def test_invalid_quantization_bits_raises(): + with pytest.raises(Exception): + SymqgIndex(DIM, max_degree=_MAX_DEGREE, quantization_bits=6) + + +@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) + idx.build(base_data, ef_construction=_EF_BUILD) + assert idx.quantization_bits == bits + + ids, dists = idx.search(query_data, k=5, ef=_EF) + assert np.all(ids < N_VECTORS) + assert np.all(np.isfinite(dists)) + + path = str(tmp_path / f"symqg_quant{bits}.index") + idx.save(path) + loaded = SymqgIndex.load(path) + assert loaded.quantization_bits == bits + loaded_ids, loaded_dists = loaded.search(query_data, k=5, ef=_EF) + np.testing.assert_array_equal(loaded_ids, ids) + np.testing.assert_allclose(loaded_dists, dists, rtol=1e-5) + + # ── save / load roundtrip ───────────────────────────────────────────────────── diff --git a/tests/unit/rabitqlib/index/qg_test.cpp b/tests/unit/rabitqlib/index/qg_test.cpp index 8fb5d00..9f46343 100644 --- a/tests/unit/rabitqlib/index/qg_test.cpp +++ b/tests/unit/rabitqlib/index/qg_test.cpp @@ -2,8 +2,11 @@ #include #include +#include #include +#include #include +#include #include #include "rabitqlib/index/symqg/qg_builder.hpp" @@ -25,6 +28,22 @@ TEST(QuantizedGraphConfigurationTest, RejectsDegreeThatCannotExcludeSelf) { ); } +TEST(QuantizedGraphConfigurationTest, AcceptsOnlySupportedVectorQuantizationBits) { + EXPECT_NO_THROW( + (QuantizedGraph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator, 0)) + ); + EXPECT_NO_THROW( + (QuantizedGraph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator, 4)) + ); + EXPECT_NO_THROW( + (QuantizedGraph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator, 8)) + ); + EXPECT_THROW( + (QuantizedGraph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator, 6)), + std::invalid_argument + ); +} + TEST(QuantizedGraphLifecycleTest, DestroysConcreteRotatorThroughBasePointer) { QuantizedGraph graph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator); EXPECT_EQ(graph.num_vertices(), 33U); @@ -90,5 +109,54 @@ TEST(QGEstimatorTest, AccumulatesAcrossUint16SafeChunks) { } } +TEST(QGQuantTest, SearchesAndRoundTripsFourAndEightBitIndexes) { + constexpr size_t kNumPoints = 33; + constexpr size_t kDim = 64; + constexpr size_t kDegree = 32; + std::vector data(kNumPoints * kDim); + for (size_t i = 0; i < data.size(); ++i) { + data[i] = std::sin(static_cast(i) * 0.13F) + + std::cos(static_cast(i) * 0.07F); + } + + for (size_t bits : {4U, 8U}) { + SCOPED_TRACE(bits); + QuantizedGraph graph( + kNumPoints, kDim, kDegree, METRIC_L2, RotatorType::MatrixRotator, bits + ); + { + QGBuilder builder(graph, kDegree, data.data(), 1); + builder.build(2); + } + graph.set_ef(kNumPoints); + + std::array ids{}; + std::array distances{}; + graph.search(data.data(), ids.size(), ids.data(), distances.data()); + for (size_t i = 0; i < ids.size(); ++i) { + EXPECT_LT(ids[i], kNumPoints); + EXPECT_TRUE(std::isfinite(distances[i])); + } + + const std::string path = + ::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); + EXPECT_TRUE(loaded.is_quantized()); + EXPECT_EQ(loaded.quantization_bits(), bits); + + std::array loaded_ids{}; + std::array loaded_distances{}; + loaded.search( + data.data(), loaded_ids.size(), loaded_ids.data(), loaded_distances.data() + ); + EXPECT_EQ(loaded_ids, ids); + EXPECT_EQ(loaded_distances, distances); + std::remove(path.c_str()); + } +} + } // namespace } // namespace rabitqlib::symqg diff --git a/tests/unit/rabitqlib/quantization/rabitq_test.cpp b/tests/unit/rabitqlib/quantization/rabitq_test.cpp index f8103a4..c6b0b4e 100644 --- a/tests/unit/rabitqlib/quantization/rabitq_test.cpp +++ b/tests/unit/rabitqlib/quantization/rabitq_test.cpp @@ -93,6 +93,38 @@ TEST(RabitqOneBitTest, FullQuantizationInitializesCodesAndFactors) { EXPECT_TRUE(std::isfinite(f_error)); } +TEST(RabitqFullQuantizationTest, ReconstructsFromCodeAndEstimatorFactors) { + constexpr size_t kDim = 64; + constexpr size_t kBits = 4; + std::array data{}; + std::array centroid{}; + std::array code{}; + std::array reconstructed{}; + for (size_t i = 0; i < kDim; ++i) { + data[i] = static_cast(i % 11) - 5.0F; + centroid[i] = static_cast(i % 3) * 0.25F; + } + + float f_add = 0; + float f_rescale = 0; + float f_error = 0; + quantize_full_single( + data.data(), centroid.data(), kDim, kBits, code.data(), f_add, f_rescale, f_error + ); + reconstruct_full_vec( + code.data(), centroid.data(), kDim, kBits, f_rescale, reconstructed.data() + ); + + const float code_center = -static_cast((1 << kBits) - 1) / 2; + const float scale = -f_rescale / 2; + for (size_t i = 0; i < kDim; ++i) { + EXPECT_FLOAT_EQ( + reconstructed[i], + centroid[i] + scale * (static_cast(code[i]) + code_center) + ); + } +} + // An exactly-zero residual coordinate must be encoded consistently by both halves of // the split code. one_bit_code() uses (residual > 0) and so calls zero negative; if // ex_bits_code() calls it positive, the assembled code lands 2^ex_bits away from where