From 5b27d789fdbff12479c2165c0303eff6341242a6 Mon Sep 17 00:00:00 2001 From: gouyt13clear Date: Wed, 16 Sep 2026 11:44:52 +0800 Subject: [PATCH] refactor: unify runtime SIMD kernel dispatch - Align dispatch across distance estimation, rotation, and matrix kernels - Isolate Eigen implementations for each ISA backend - Preserve portable wheel performance and search accuracy - Clean up includes and improve static-analysis coverage - Add regression tests and document dispatch conventions --- AGENTS.md | 5 +- CMakeLists.txt | 12 +- CONTRIBUTING.md | 48 +- docs/docs/index/qg.md | 12 +- include/rabitqlib/index/estimator.hpp | 77 +- include/rabitqlib/index/hnsw/hnsw.hpp | 30 +- include/rabitqlib/index/ivf/initializer.hpp | 23 +- include/rabitqlib/index/ivf/ivf.hpp | 11 +- .../rabitqlib/index/symqg/detail/pipnn.hpp | 39 +- include/rabitqlib/index/symqg/qg.hpp | 889 ++---------------- include/rabitqlib/index/symqg/qg_builder.hpp | 20 +- include/rabitqlib/quantization/rabitq.hpp | 1 + include/rabitqlib/simd/dispatch.hpp | 6 +- include/rabitqlib/simd/estimator_dispatch.hpp | 103 ++ include/rabitqlib/simd/hnsw_dispatch.hpp | 29 + include/rabitqlib/simd/matrix_dispatch.hpp | 123 +++ .../rabitqlib/simd/pack_excode_dispatch.hpp | 29 +- .../rabitqlib/simd/quantization_dispatch.hpp | 3 + include/rabitqlib/simd/rotator_dispatch.hpp | 40 +- include/rabitqlib/utils/rotator.hpp | 99 +- python_bindings/symqg_bindings.cpp | 6 +- sample/cpp/symqg_indexing.cpp | 2 +- sample/cpp/symqg_querying.cpp | 2 +- scripts/check-includes.sh | 5 +- scripts/check-tidy.sh | 2 +- src/index/ivf_search.cpp | 17 + src/index/qg.cpp | 796 ++++++++++++++++ src/simd/dispatch.cpp | 481 ++++++---- src/simd/estimator_avx2.cpp | 61 ++ src/simd/estimator_avx512.cpp | 30 + src/simd/estimator_generic.cpp | 56 ++ src/simd/estimator_kernels.hpp | 172 ++++ src/simd/fastscan_generic.cpp | 18 + src/simd/matrix_avx2.cpp | 45 + src/simd/matrix_avx512.cpp | 45 + src/simd/matrix_generic.cpp | 45 + src/simd/matrix_kernels.hpp | 85 ++ src/simd/quantization_generic.cpp | 24 + src/simd/rotator_avx2.cpp | 15 + src/simd/rotator_avx512.cpp | 15 + src/simd/rotator_kernels.hpp | 104 ++ .../{space_float.cpp => space_generic.cpp} | 0 tests/python/test_ivf.py | 23 + .../unit/rabitqlib/fastscan/fastscan_test.cpp | 96 ++ .../unit/rabitqlib/index/initializer_test.cpp | 43 + tests/unit/rabitqlib/index/ivf_test.cpp | 14 + tests/unit/rabitqlib/index/qg_test.cpp | 134 +-- .../rabitqlib/utils/matrix_dispatch_test.cpp | 109 +++ tests/unit/rabitqlib/utils/rotator_test.cpp | 90 ++ 49 files changed, 2804 insertions(+), 1330 deletions(-) create mode 100644 include/rabitqlib/simd/estimator_dispatch.hpp create mode 100644 include/rabitqlib/simd/hnsw_dispatch.hpp create mode 100644 include/rabitqlib/simd/matrix_dispatch.hpp create mode 100644 src/index/ivf_search.cpp create mode 100644 src/index/qg.cpp create mode 100644 src/simd/estimator_avx2.cpp create mode 100644 src/simd/estimator_avx512.cpp create mode 100644 src/simd/estimator_generic.cpp create mode 100644 src/simd/estimator_kernels.hpp create mode 100644 src/simd/fastscan_generic.cpp create mode 100644 src/simd/matrix_avx2.cpp create mode 100644 src/simd/matrix_avx512.cpp create mode 100644 src/simd/matrix_generic.cpp create mode 100644 src/simd/matrix_kernels.hpp create mode 100644 src/simd/quantization_generic.cpp create mode 100644 src/simd/rotator_kernels.hpp rename src/simd/{space_float.cpp => space_generic.cpp} (100%) create mode 100644 tests/unit/rabitqlib/utils/matrix_dispatch_test.cpp diff --git a/AGENTS.md b/AGENTS.md index ecd626a..f27322d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -18,7 +18,7 @@ factors estimate L2 distance or inner product. | `include/rabitqlib/index/{ivf,hnsw,symqg}/` | Index construction, persistence, and search | | `include/rabitqlib/index/{query,estimator}.hpp` | Query state and distance estimation | | `include/rabitqlib/simd/`, `src/simd/` | Kernel declarations, implementations, and dispatch | -| `src/index/` | Compiled HNSW search kernels | +| `src/index/` | Compiled SymphonyQG implementation, HNSW search kernels, and IVF candidate insertion | | `src/utils/cpu_features.cpp` | Generic x86 feature detection | | `include/rabitqlib/utils/` | Rotation, allocation, buffers, I/O, and helpers | | `python_bindings/` | pybind11 extension and index wrappers | @@ -110,6 +110,9 @@ Recommended: - Public/generic code calls centralized dispatch entry points. Keep ISA-specific translation units and their flags in `CMakeLists.txt` synchronized with feature predicates in `src/simd/dispatch.cpp` and detection in `src/utils/cpu_features.cpp`, including HNSW source groups. +- Use the shared resolver in `src/simd/dispatch.cpp` for cached selection, including HNSW. + Keep calculations in backend source files; see the dispatch coverage table in + [CONTRIBUTING.md](CONTRIBUTING.md#dispatch-conventions-and-coverage). - Dispatch resolves function pointers during static initialization. Detection must use safe generic code; never execute a high-ISA kernel to find out whether the CPU supports it. - Semantic kernel changes must cover every implementation and a backend-independent reference diff --git a/CMakeLists.txt b/CMakeLists.txt index 08f9e04..383a77d 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -45,9 +45,15 @@ function(rabitq_enable_warnings target) endfunction() set(RABITQ_COMMON_SOURCES + src/index/ivf_search.cpp + src/index/qg.cpp src/utils/cpu_features.cpp src/simd/dispatch.cpp - src/simd/space_float.cpp + src/simd/space_generic.cpp + src/simd/fastscan_generic.cpp + src/simd/quantization_generic.cpp + src/simd/estimator_generic.cpp + src/simd/matrix_generic.cpp ) set(RABITQ_AVX2_SOURCES @@ -56,6 +62,8 @@ set(RABITQ_AVX2_SOURCES src/simd/space_excode_avx2.cpp src/simd/space_avx2.cpp src/simd/fastscan_avx2.cpp + src/simd/estimator_avx2.cpp + src/simd/matrix_avx2.cpp src/simd/warmup_avx2.cpp src/simd/rotator_avx2.cpp ) @@ -66,6 +74,8 @@ set(RABITQ_AVX512_SOURCES src/simd/space_excode_avx512.cpp src/simd/space_avx512.cpp src/simd/fastscan_avx512.cpp + src/simd/estimator_avx512.cpp + src/simd/matrix_avx512.cpp src/simd/rotator_avx512.cpp ) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 8010697..21913df 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -231,7 +231,8 @@ cmake -S . -B build-includes -G Ninja \ The script checks library sources using their compilation database and checks headers as main files, since this clang-tidy check does not report findings in -included headers. Private headers are checked with AVX2 and AVX-512 flags. +included headers. Private headers are checked with AVX2 and AVX-512 flags. New +untracked files are included, and tracked files deleted from the working tree are skipped. Vendored files are excluded. The script also ignores suggestions to include Eigen and hnswlib implementation headers behind their existing public headers; these vendor snapshots lack the export annotations needed by include-cleaner. @@ -399,6 +400,51 @@ requests labeled `duplicate` or `invalid` are omitted from release notes. Never execute a high-ISA implementation merely to test whether that ISA is supported; detection must happen in generic code first. +#### Dispatch conventions and coverage + +- Keep backend-neutral declarations in `include/rabitqlib/simd/*_dispatch.hpp` and + implementations in `src/simd/*_{generic,avx2,avx512}.cpp`. HNSW keeps its existing + `src/index/` implementations and compatibility namespaces. +- Select a cached function pointer through `resolve_kernel` in `src/simd/dispatch.cpp`. + The order is AVX-512, AVX2/FMA, then the existing generic implementation or a descriptive + unsupported-operation exception. Public wrappers do not repeat feature checks. +- Preserve stricter predicates: population-count kernels need AVX512_VPOPCNTDQ, and HNSW's + AVX-512 core variant also needs AVX2 for its warmup implementation. Do not infer support + from a backend name or from `__AVX*__` macros in a public header. +- Put calculations and scratch-storage helpers outside the dispatcher. Shared implementation + headers use internal linkage so independently compiled backends retain their own bodies. + Include the FFHT implementation inside the private kernel namespace for the same reason. +- Pass ordinary pointers, sizes, and library-owned query state across ISA boundaries. Never + pass Eigen matrix/packet objects between backends. The private matrix implementation + header includes Eigen under a namespace selected by each backend translation unit. + Otherwise, Eigen emits identically named out-of-line template helpers, which the linker + can merge across incompatible ISA builds. + Keep this isolation when adding matrix kernels; do not modify the vendor snapshot. +- Preserve existing public names and namespaces, including legacy functions with `_avx` in + their name that now dispatch at runtime. The explicit `_avx2` and `_avx512` entry points + are for selected kernels and capability-guarded backend tests. + +The current first-party kernel audit is summarized below. Dispatching an index operation +covers its arithmetic kernels, not every scalar loop in construction and search. + +| Area | Dispatch coverage / deliberate boundary | +| --- | --- | +| Raw float distances, norms, packed-code products | Runtime AVX2/AVX-512 selection; generic raw-float fallback | +| Quantizer rescale search | SIMD bounded search; the certified scalar event sweep remains the fallback | +| Integer scalar quantization, extra-code packing, transpose, sign masks | Existing runtime selection; byte layouts and rounding rules are unchanged | +| Standard and high-accuracy FastScan | Runtime selection; generic LUT construction stays separate from selection | +| IVF and float SymphonyQG batch correction | Complete estimator runs in the selected backend; non-float template paths remain generic | +| FHT/Kac rotation | Complete rotation and scaling run in selected ISA translation units; imported FFHT AVX butterflies are unchanged | +| Float matrix rotation and PiPNN construction | Matrix products, row norms, and lower-triangle pairwise distances use isolated matrix backends | +| HNSW search and IVF centroid routing | Cached HNSW search selection; centroid routing uses the common raw-distance dispatcher | +| Quantization orchestration, reconstruction, non-float utilities | Template/control code remains generic; no blanket native tuning or reduction-order rewrite | +| Graph scheduling, candidate queues, I/O, allocation, random initialization | Generic control code; IVF one-bit candidate insertion stays in a small compiled function to avoid inlining-induced register spills; thread scheduling and seeds remain caller-owned | +| Example KMeans training | Uses external FAISS, whose build and dispatch are independent of this package | + +Portable wheels continue to disable `RABITQ_ENABLE_NATIVE_OPTIMIZATION`. Adding an optimized +backend does not add support for generic-CPU quantized search or AArch64; those require +complete implementations and separate compatibility validation. + ### Change quantization or packing Check all of these together: diff --git a/docs/docs/index/qg.md b/docs/docs/index/qg.md index 2d1e15b..d6833c2 100644 --- a/docs/docs/index/qg.md +++ b/docs/docs/index/qg.md @@ -39,17 +39,19 @@ window; `ef` controls the query search window. Python defaults to one thread. ### C++ -The C++ API uses `rabitqlib::symqg::QuantizedGraph` and `QGBuilder`: +The C++ API uses the float-only `rabitqlib::symqg::QuantizedGraph` and `QGBuilder`. +`QuantizedGraph` is a non-template class; use `QuantizedGraph` instead of +`QuantizedGraph`. Its implementation is compiled in `src/index/qg.cpp`: ```cpp -QuantizedGraph( +QuantizedGraph( size_t num, size_t dim, size_t max_deg, MetricType metric_type = METRIC_L2, RotatorType rotator_type = RotatorType::FhtKacRotator, size_t quantization_bits = 0 ); QGBuilder( - QuantizedGraph& index, uint32_t ef_build, const float* data, + QuantizedGraph& index, uint32_t ef_build, const float* data, size_t num_threads = std::numeric_limits::max(), QGInitialization init = QGInitialization::PiPNN ); @@ -64,7 +66,7 @@ The builder handles initialization internally. using namespace rabitqlib::symqg; // data contains rows * cols floats. -QuantizedGraph qg(rows, cols, 32); +QuantizedGraph qg(rows, cols, 32); { QGBuilder builder(qg, 200, data.data(), 32, QGInitialization::PiPNN); builder.build(); @@ -97,7 +99,7 @@ C++ search accepts one vector in the original input dimension and writes `k` IDs and distances: ```cpp -QuantizedGraph qg; +QuantizedGraph qg; qg.load("qg_example.index"); qg.set_ef(100); diff --git a/include/rabitqlib/index/estimator.hpp b/include/rabitqlib/index/estimator.hpp index b6b00ab..a607890 100644 --- a/include/rabitqlib/index/estimator.hpp +++ b/include/rabitqlib/index/estimator.hpp @@ -3,12 +3,13 @@ #include #include #include +#include #include "rabitqlib/defines.hpp" #include "rabitqlib/fastscan/fastscan.hpp" -#include "rabitqlib/fastscan/highacc_fastscan.hpp" #include "rabitqlib/index/query.hpp" #include "rabitqlib/quantization/data_layout.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" #include "rabitqlib/utils/space.hpp" #include "rabitqlib/utils/warmup_space.hpp" @@ -33,74 +34,9 @@ inline void split_batch_estdist( float* ip_x0_qr, bool use_hacc ) { - constexpr size_t kSafeChunkDim = 1024; - ConstBatchDataMap cur_batch(batch_data, padded_dim); - std::array accu_values{}; - RowMajorArrayMap accu_arr(accu_values.data(), 1, fastscan::kBatchSize); - const auto* codes_ptr = cur_batch.bin_code(); - const auto* lut_ptr = q_obj.lut(); - if (use_hacc) { - std::array accu_res; - size_t remaining_dim = padded_dim; - - while (remaining_dim > kSafeChunkDim) { - fastscan::accumulate_hacc(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); - codes_ptr += kSafeChunkDim << 2; - lut_ptr += kSafeChunkDim << 3; - for (size_t i = 0; i < fastscan::kBatchSize; ++i) { - accu_arr.data()[i] += accu_res[i]; - } - remaining_dim -= kSafeChunkDim; - } - - fastscan::accumulate_hacc(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); - for (size_t i = 0; i < fastscan::kBatchSize; ++i) { - accu_arr.data()[i] += accu_res[i]; - } - } else { - std::array accu_res; - size_t remaining_dim = padded_dim; - - while (remaining_dim > kSafeChunkDim) { - fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); - codes_ptr += kSafeChunkDim << 2; - lut_ptr += kSafeChunkDim << 2; - for (size_t i = 0; i < fastscan::kBatchSize; ++i) { - accu_arr.data()[i] += accu_res[i]; - } - remaining_dim -= kSafeChunkDim; - } - - fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); - for (size_t i = 0; i < fastscan::kBatchSize; ++i) { - accu_arr.data()[i] += accu_res[i]; - } - } - - 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( - f_rescale_values.data(), 1, fastscan::kBatchSize - ); - ConstRowMajorArrayMap f_error_arr( - f_error_values.data(), 1, fastscan::kBatchSize + simd::split_batch_estdist( + batch_data, q_obj, padded_dim, est_distance, low_distance, ip_x0_qr, use_hacc ); - - RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); - RowMajorArrayMap ip_x0_qr_arr(ip_x0_qr, 1, fastscan::kBatchSize); - RowMajorArrayMap low_dist_arr(low_distance, 1, fastscan::kBatchSize); - - ip_x0_qr_arr = q_obj.delta() * (accu_arr.template cast()) + q_obj.sum_vl_lut(); - - est_dist_arr = - f_add_arr + q_obj.g_add() + f_rescale_arr * (ip_x0_qr_arr + q_obj.k1xsumq()); - - low_dist_arr = est_dist_arr - f_error_arr * q_obj.g_error(); } /** @@ -178,6 +114,11 @@ template inline void qg_batch_estdist( const char* batch_data, const BatchQuery& q_obj, size_t padded_dim, T* est_distance ) { + if constexpr (std::is_same_v && std::is_same_v) { + simd::qg_batch_estdist(batch_data, q_obj, padded_dim, est_distance); + return; + } + // Each 4-dimensional codebook can contribute at most 255, so 1024 dimensions // produce at most 255 * (1024 / 4) = 65280 in the uint16_t FastScan result. constexpr size_t kSafeChunkDim = 1024; diff --git a/include/rabitqlib/index/hnsw/hnsw.hpp b/include/rabitqlib/index/hnsw/hnsw.hpp index e85f6e9..f2bfc11 100644 --- a/include/rabitqlib/index/hnsw/hnsw.hpp +++ b/include/rabitqlib/index/hnsw/hnsw.hpp @@ -27,8 +27,8 @@ #include "rabitqlib/index/query.hpp" #include "rabitqlib/quantization/data_layout.hpp" #include "rabitqlib/quantization/rabitq.hpp" +#include "rabitqlib/simd/hnsw_dispatch.hpp" #include "rabitqlib/utils/buffer.hpp" -#include "rabitqlib/utils/cpu_features.hpp" #include "rabitqlib/utils/memory.hpp" #include "rabitqlib/utils/rotator.hpp" #include "rabitqlib/utils/space.hpp" @@ -44,22 +44,6 @@ using maxheap = std::priority_queue; template using minheap = std::priority_queue, std::greater>; -class HierarchicalNSW; - -namespace detail { - -maxheap> search_knn_avx2(HierarchicalNSW&, const float*, size_t); - -maxheap> search_knn_avx512_core( - HierarchicalNSW&, const float*, size_t -); - -maxheap> search_knn_avx512_popcnt( - HierarchicalNSW&, const float*, size_t -); - -} // namespace detail - class HierarchicalNSW { public: explicit HierarchicalNSW(){}; @@ -1108,17 +1092,7 @@ inline std::vector>> HierarchicalNSW::search( inline maxheap> HierarchicalNSW::search_knn( const float* rotated_query, size_t TOPK ) { - if (rabitqlib::cpu::has_avx512_popcnt()) { - return detail::search_knn_avx512_popcnt(*this, rotated_query, TOPK); - } - if (rabitqlib::cpu::has_avx512_core() && rabitqlib::cpu::has_avx2()) { - return detail::search_knn_avx512_core(*this, rotated_query, TOPK); - } - if (rabitqlib::cpu::has_avx2()) { - return detail::search_knn_avx2(*this, rotated_query, TOPK); - } - - throw std::runtime_error("HNSW search requires AVX2/FMA or AVX512 support"); + return detail::search_knn(*this, rotated_query, TOPK); } template diff --git a/include/rabitqlib/index/ivf/initializer.hpp b/include/rabitqlib/index/ivf/initializer.hpp index d5458b3..57c34ff 100644 --- a/include/rabitqlib/index/ivf/initializer.hpp +++ b/include/rabitqlib/index/ivf/initializer.hpp @@ -172,12 +172,33 @@ class FlatInitializer : public Initializer { } }; +// Keep centroid routing on the same runtime-selected distance kernels as flat IVF. +class CentroidL2Space : public hnswlib::SpaceInterface { + private: + size_t dim_; + + static float distance(const void* a, const void* b, const void* dim) { + return euclidean_sqr( + static_cast(a), + static_cast(b), + *static_cast(dim) + ); + } + + public: + explicit CentroidL2Space(size_t dim) : dim_(dim) {} + + size_t get_data_size() override { return dim_ * sizeof(float); } + hnswlib::DISTFUNC get_dist_func() override { return distance; } + void* get_dist_func_param() override { return &dim_; } +}; + class HNSWInitializer : public Initializer { private: int M_ = 16; int ef_construction_ = 400; hnswlib::HierarchicalNSW* alg_hnsw_ = nullptr; - hnswlib::L2Space space_; + CentroidL2Space space_; public: explicit HNSWInitializer(size_t d, size_t k) : Initializer(d, k), space_(d) { diff --git a/include/rabitqlib/index/ivf/ivf.hpp b/include/rabitqlib/index/ivf/ivf.hpp index 80e814a..a563140 100644 --- a/include/rabitqlib/index/ivf/ivf.hpp +++ b/include/rabitqlib/index/ivf/ivf.hpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include #include @@ -29,6 +30,10 @@ #include "rabitqlib/utils/tools.hpp" namespace rabitqlib::ivf { +namespace detail { +void insert_candidates(buffer::SearchBuffer&, const PID*, const float*, size_t); +} // namespace detail + class IVF { private: using ByteStorage = @@ -717,11 +722,7 @@ inline void IVF::scan_one_batch( // Without reranking data, return the one-bit estimates directly. if (ex_bits_ == 0 && !raw_reranking_) { - for (size_t i = 0; i < num_points; ++i) { - PID id = ids[i]; - float ex_dist = est_distance[i]; - knns.insert(id, ex_dist); - } + detail::insert_candidates(knns, ids, est_distance.data(), num_points); return; } diff --git a/include/rabitqlib/index/symqg/detail/pipnn.hpp b/include/rabitqlib/index/symqg/detail/pipnn.hpp index eb93fe2..a1257f7 100644 --- a/include/rabitqlib/index/symqg/detail/pipnn.hpp +++ b/include/rabitqlib/index/symqg/detail/pipnn.hpp @@ -17,6 +17,7 @@ #include #include "rabitqlib/defines.hpp" +#include "rabitqlib/simd/matrix_dispatch.hpp" #include "rabitqlib/utils/buffer.hpp" #include "rabitqlib/utils/space.hpp" #include "rabitqlib/utils/tools.hpp" @@ -88,20 +89,9 @@ inline void gather(const float* data, size_t dim, const Bucket& ids, Workspace& } inline void pairwise(Workspace& work, size_t size, size_t dim, MetricType metric) { - const auto count = static_cast(size); - RowMajorMatrixMap points(work.points, count, static_cast(dim)); - RowMajorMatrixMap distances(work.distances, count, count); - VectorMap norms(work.norms, count); - distances.setZero(); - distances.template selfadjointView().rankUpdate(points); - norms = points.rowwise().squaredNorm(); - for (Eigen::Index i = 0; i < count; ++i) { - for (Eigen::Index j = 0; j <= i; ++j) { - const float dot = distances(i, j); - distances(i, j) = - metric == METRIC_L2 ? std::max(0.0F, norms[i] + norms[j] - 2 * dot) : -dot; - } - } + simd::pairwise_distances_lower( + work.points, work.distances, work.norms, size, dim, metric == METRIC_IP + ); } inline std::vector partition( @@ -130,7 +120,8 @@ inline std::vector partition( for (size_t i = 0; i < leader_count; ++i) { std::copy_n(data + leaders[i] * dim, dim, leader_data.data() + i * dim); } - const Vector leader_norms = leader_data.rowwise().squaredNorm(); + Vector leader_norms(leader_count); + simd::row_norms(leader_data.data(), leader_norms.data(), leader_count, dim); std::vector assignments(ids.size() * fanout); const auto assign_tile = [&](size_t begin, Workspace& work) { const size_t count = std::min(kTileSize, ids.size() - begin); @@ -146,8 +137,10 @@ inline std::vector partition( for (size_t i = 0; i < count; ++i) { std::copy_n(data + ids[begin + i] * dim, dim, points.data() + i * dim); } - distances.noalias() = points * leader_data.transpose(); - norms = points.rowwise().squaredNorm(); + simd::matrix_product_transposed( + points.data(), leader_data.data(), distances.data(), count, dim, leader_count + ); + simd::row_norms(points.data(), norms.data(), count, dim); work.candidates.resize(leader_count); for (size_t i = 0; i < count; ++i) { for (size_t j = 0; j < leader_count; ++j) { @@ -315,14 +308,14 @@ inline InitialGraph build_initial_graph( #pragma omp parallel for num_threads(threads) schedule(static) for (size_t begin = 0; begin < count; begin += kTileSize) { const size_t rows = std::min(kTileSize, count - begin); - ConstRowMajorMatrixMap points( + simd::matrix_product( data + begin * dim, - static_cast(rows), - static_cast(dim) + projections.data(), + sketches.data() + begin * kHashBits, + rows, + dim, + kHashBits ); - sketches - .middleRows(static_cast(begin), static_cast(rows)) - .noalias() = points * projections; } const size_t capacity = degree * 5 / 2; std::vector table(count * capacity); diff --git a/include/rabitqlib/index/symqg/qg.hpp b/include/rabitqlib/index/symqg/qg.hpp index 2dcbdbd..f7799ad 100644 --- a/include/rabitqlib/index/symqg/qg.hpp +++ b/include/rabitqlib/index/symqg/qg.hpp @@ -1,89 +1,65 @@ #pragma once -#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/query.hpp" #include "rabitqlib/quantization/data_layout.hpp" -#include "rabitqlib/quantization/pack_excode.hpp" -#include "rabitqlib/quantization/rabitq.hpp" #include "rabitqlib/utils/buffer.hpp" -#include "rabitqlib/utils/io.hpp" #include "rabitqlib/utils/memory.hpp" #include "rabitqlib/utils/rotator.hpp" #include "rabitqlib/utils/space.hpp" -#include "rabitqlib/utils/tools.hpp" #include "rabitqlib/utils/visited_set.hpp" namespace rabitqlib::symqg { -template class QuantizedQuery { private: - const T* rotated_query_; - T k1xsumq_; - T g_add_; + const float* rotated_query_; + float k1xsumq_; + float 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_; } + const float* rotated_query, + const float* centroid, + size_t padded_dim, + MetricType metric_type + ); + [[nodiscard]] const float* rotated_query() const; + [[nodiscard]] float k1xsumq() const; + [[nodiscard]] float g_add() const; }; -template class QuantizedGraph { friend class QGBuilder; 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) = nullptr; // raw-vector distance - 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 + // Raw-vector distance. + float (*raw_dist_func_)(const float*, const float*, size_t) = nullptr; + 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 + std::vector centroid_; // rotated global centroid for qg-quant ex_ipfunc quantized_ip_func_ = nullptr; - using RowStorage = std::vector>; + 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. + // Typed storage establishes float lifetimes for raw vectors; packed portions + // are accessed only as bytes, never as float values. RowStorage data_; - std::unique_ptr> rotator_; // data rotator + std::unique_ptr> rotator_; // data rotator // Position of row data (raw vector or packed qg-quant vector), neighbor // quantization data, and neighbor IDs. Since every degree equals degree_bound_ @@ -94,23 +70,11 @@ class QuantizedGraph { 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_add(size_t lhs, size_t 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 checked_multiply(size_t lhs, size_t rhs); - [[nodiscard]] static size_t padded_dimension(size_t dim) { - return (checked_add(dim, 63) / 64) * 64; - } + [[nodiscard]] static size_t padded_dimension(size_t dim); void validate_configuration() const; @@ -118,75 +82,59 @@ class QuantizedGraph { void initialize(); - void copy_vectors(const T*, size_t); - - void set_quantization_centroid(const T* centroid); + void copy_vectors(const float*, size_t); - [[nodiscard]] char* get_row_data(PID data_id) { - return reinterpret_cast(get_vector(data_id)); - } + void set_quantization_centroid(const float* centroid); - [[nodiscard]] const char* get_row_data(PID data_id) const { - return reinterpret_cast(get_vector(data_id)); - } + [[nodiscard]] char* get_row_data(PID data_id); - [[nodiscard]] T* get_vector(PID data_id) { - return data_.data() + ((row_offset_ / sizeof(T)) * data_id); - } + [[nodiscard]] const char* get_row_data(PID data_id) const; - [[nodiscard]] const T* get_vector(PID data_id) const { - return data_.data() + ((row_offset_ / sizeof(T)) * data_id); - } + [[nodiscard]] float* get_vector(PID data_id); - [[nodiscard]] char* get_quantized_vector(PID data_id) { return get_row_data(data_id); } + [[nodiscard]] const float* get_vector(PID data_id) const; - [[nodiscard]] const char* get_quantized_vector(PID data_id) const { - return get_row_data(data_id); - } + [[nodiscard]] char* get_quantized_vector(PID data_id); - void prepare_query(const T*, std::vector&, std::optional>&) const; + [[nodiscard]] const char* get_quantized_vector(PID data_id) const; - const T* prepare_build_query(PID, std::vector&, std::optional>&) + void prepare_query(const float*, std::vector&, std::optional&) const; - T point_distance(const T*, const QuantizedQuery*, PID) const; + const float* + prepare_build_query(PID, std::vector&, std::optional&) const; + + float point_distance(const float*, const QuantizedQuery*, PID) const; - T quantized_distance(const QuantizedQuery&, PID) const; + float quantized_distance(const QuantizedQuery&, PID) const; - void reconstruct_quantized_vector(PID, T*) const; + void reconstruct_quantized_vector(PID, float*) const; - [[nodiscard]] char* get_batch_data(PID data_id) { - return get_row_data(data_id) + batch_data_offset_; - } + [[nodiscard]] char* get_batch_data(PID data_id); - [[nodiscard]] const char* get_batch_data(PID data_id) const { - return get_row_data(data_id) + batch_data_offset_; - } + [[nodiscard]] const char* get_batch_data(PID data_id) const; - [[nodiscard]] rabitqlib::detail::PackedArrayView get_neighbors(PID data_id) { - return rabitqlib::detail::PackedArrayView( - get_row_data(data_id) + neighbor_offset_ - ); - } + [[nodiscard]] rabitqlib::detail::PackedArrayView get_neighbors(PID data_id); [[nodiscard]] rabitqlib::detail::ConstPackedArrayView get_neighbors(PID data_id - ) const { - return rabitqlib::detail::ConstPackedArrayView( - get_row_data(data_id) + neighbor_offset_ - ); - } + ) const; void - find_candidates(PID, size_t, std::vector>&, VisitedSet&, const std::vector&) + find_candidates(PID, size_t, std::vector>&, VisitedSet&, const std::vector&) const; - void update_qg(PID, const std::vector>&); + void update_qg(PID, const std::vector>&); void - update_results(buffer::SearchBuffer&, VisitedSet&, const T*, const QuantizedQuery*); + update_results(buffer::SearchBuffer&, VisitedSet&, const float*, const QuantizedQuery*); void scan_neighbors( - const BatchQuery&, PID, T*, buffer::SearchBuffer&, VisitedSet&, size_t + const BatchQuery&, + PID, + float*, + buffer::SearchBuffer&, + VisitedSet&, + size_t ) const; public: @@ -199,35 +147,30 @@ class QuantizedGraph { size_t quantization_bits = 0 ); - explicit QuantizedGraph() = default; + explicit QuantizedGraph(); - ~QuantizedGraph() = default; + ~QuantizedGraph(); QuantizedGraph(const QuantizedGraph&) = delete; QuantizedGraph& operator=(const QuantizedGraph&) = delete; - QuantizedGraph(QuantizedGraph&&) noexcept = default; - QuantizedGraph& operator=(QuantizedGraph&&) noexcept = default; + QuantizedGraph(QuantizedGraph&&) noexcept; + QuantizedGraph& operator=(QuantizedGraph&&) noexcept; - [[nodiscard]] auto num_vertices() const { return this->num_points_; } + [[nodiscard]] size_t num_vertices() const; - [[nodiscard]] auto dimension() const { return this->dim_; } + [[nodiscard]] size_t dimension() const; - [[nodiscard]] auto degree_bound() const { return this->degree_bound_; } + [[nodiscard]] size_t degree_bound() const; - [[nodiscard]] auto entry_point() const { return this->entry_point_; } + [[nodiscard]] PID entry_point() const; - [[nodiscard]] auto metric_type() const { return this->metric_type_; } + [[nodiscard]] MetricType metric_type() const; - [[nodiscard]] auto quantization_bits() const { return this->quantization_bits_; } + [[nodiscard]] size_t quantization_bits() const; - [[nodiscard]] bool is_quantized() const { return quantization_bits_ != 0; } + [[nodiscard]] bool is_quantized() const; - void set_ep(PID entry) { - if (entry >= num_points_) { - throw std::invalid_argument("QuantizedGraph entry point is out of range"); - } - entry_point_ = entry; - } + void set_ep(PID entry); void save(const char*) const; @@ -237,707 +180,11 @@ class QuantizedGraph { /* search and copy results to KNN */ void search( - const T* __restrict__ query, + const float* __restrict__ query, uint32_t knn, uint32_t* __restrict__ results, - T* __restrict__ dists + float* __restrict__ dists ); }; -template -inline QuantizedGraph::QuantizedGraph( - 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) - , dim_(dim) - , padded_dim_(dim) - , raw_dist_func_((metric_type == METRIC_IP) ? dot_product_dis : euclidean_sqr) - , metric_type_(metric_type) - , rotator_type_(rotator_type) - , quantization_bits_(quantization_bits) { - validate_configuration(); - initialize(); -} - -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" - ); - } - if (degree_bound_ >= num_points_) { - throw std::invalid_argument( - "QuantizedGraph degree bound must be smaller than the number of points" - ); - } - if (num_points_ > buffer::kSearchBufferMaxPointCount) { - throw std::invalid_argument( - "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" - ); - } - 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, 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"); - } else { - if (centroid_.size() != padded_dim_) { - throw std::logic_error( - "qg-quant centroid must be set before copying vectors" - ); - } -#pragma omp parallel num_threads(thread_count) - { - 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 f_add; - T f_rescale; - T unused_f_error = 0; - quant::quantize_full_single( - rotated_data.data(), - centroid_.data(), - padded_dim_, - quantization_bits_, - quantized_data.data(), - 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(), - padded_dim_, - quantization_bits_ - ); - } - } - return; - } - } -#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); - std::copy(src, src + dim_, dst); - } -} - -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 { - 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; - 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)); - output.write(reinterpret_cast(&dim_), sizeof(size_t)); - output.write(reinterpret_cast(&padded_dim_), sizeof(size_t)); - 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 */ - 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"); - } - - std::ifstream input(filename, std::ios::binary); - if (!input.is_open()) { - 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; - read_value(magic, "format marker"); - if (magic == kFormatMagic) { - uint32_t version = 0; - read_value(version, "format 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); - } - - 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) { - read_value(loaded.quantization_bits_, "quantization bits"); - } else { - loaded.quantization_bits_ = 0; - } - - 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(); - - 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); - - 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"); - } - - loaded.initialize(); - - 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; -} - -template -inline void QuantizedGraph::search( - const T* __restrict__ query, - uint32_t k, - 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); - BatchQuery q_obj(rotated_query.data(), padded_dim_, metric_type_); - - buffer::SearchBuffer search_pool(ef_); - // init search buffer - search_pool.insert(this->entry_point_, std::numeric_limits::max()); - - buffer::SearchBuffer res_pool(k); // result buffer - thread_local VisitedSet visited; - thread_local size_t visited_size = 0; - if (visited_size != num_points_) { - visited.initialize(num_points_, num_points_ / 10); - visited_size = num_points_; - } - visited.clear(); - auto* vis = &visited; - - std::vector est_dist(degree_bound_); // estimated distances - - while (search_pool.has_next()) { - PID cur_node = search_pool.pop(); - if (vis->get(cur_node)) { - continue; - } - vis->set(cur_node); - - const T vertex_distance = - point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_node); - q_obj.set_g_add(vertex_distance); - - scan_neighbors( - q_obj, cur_node, est_dist.data(), search_pool, *vis, this->degree_bound_ - ); - res_pool.insert(cur_node, vertex_distance); - } - - 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); -} - -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, - PID data_id, - T* est_dist, - buffer::SearchBuffer& search_pool, - VisitedSet& vis, - size_t cur_degree -) const { - const auto* batch_data = get_batch_data(data_id); - for (size_t i = 0; i < cur_degree; i += fastscan::kBatchSize) { - qg_batch_estdist(batch_data, q_obj, padded_dim_, est_dist + i); - batch_data += QGBatchDataMap::data_bytes(padded_dim_); - } - - 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; - for (size_t lane = 0; lane < fastscan::kBatchSize; ++lane) { - candidate_mask |= static_cast(!(est_dist[begin + lane] > threshold)) - << lane; - } - - // Construction can leave a partial batch with stale IDs in its unused lanes. - const size_t remaining = cur_degree - begin; - if (remaining < fastscan::kBatchSize) { - candidate_mask &= (uint32_t{1} << remaining) - 1; - } - - while (candidate_mask != 0) { - const auto lane = static_cast(__builtin_ctz(candidate_mask)); - candidate_mask &= candidate_mask - 1; - const size_t i = begin + lane; - 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(get_row_data(search_pool.next_id()), 10); - } - } -} - -template -inline void QuantizedGraph::update_results( - buffer::SearchBuffer& result_pool, - VisitedSet& vis, - const T* query, - const QuantizedQuery* quantized_query -) { - if (result_pool.is_full()) { - return; - } - - 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 = neighbors[i]; - if (!vis.get(cur_neighbor)) { - vis.set(cur_neighbor); - result_pool.insert( - cur_neighbor, point_distance(query, quantized_query, cur_neighbor) - ); - } - } - if (result_pool.is_full()) { - break; - } - } -} - -// initialize const offsets & data array -template -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) - ); - 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_)); - - assert(padded_dim_ % 64 == 0); - assert(padded_dim_ >= dim_); - - initialize_layout(); - - 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_); - } -} - -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_ - ); -} - -// Construction sources come from owned raw rows or from the existing RaBitQ codes. -// Reconstructed sources are already rotated; never pass them through prepare_query. -template -inline const T* QuantizedGraph::prepare_build_query( - PID id, std::vector& rotated, std::optional>& prepared -) const { - if (is_quantized()) { - rotated.resize(padded_dim_); - reconstruct_quantized_vector(id, rotated.data()); - prepared.emplace(rotated.data(), centroid_.data(), padded_dim_, metric_type_); - return rotated.data(); - } - prepared.reset(); - return get_vector(id); -} - -// find candidate neighbors for cur_id, exclude the vertex itself -template -inline void QuantizedGraph::find_candidates( - PID cur_id, - size_t search_ef, - std::vector>& results, - VisitedSet& vis, - const std::vector& degrees -) const { - std::vector rotated_query(padded_dim_); - std::optional> quantized_query; - const T* query = prepare_build_query(cur_id, rotated_query, quantized_query); - if (!is_quantized()) { - rotator_->rotate(query, rotated_query.data()); - } - BatchQuery q_obj(rotated_query.data(), padded_dim_, metric_type_); - - // insert entry point to initialize search buffer - buffer::SearchBuffer tmp_pool(search_ef); - 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 - while (tmp_pool.has_next()) { - auto cur_candi = tmp_pool.pop(); - if (vis.get(cur_candi)) { - continue; - } - vis.set(cur_candi); - auto cur_degree = degrees[cur_candi]; - const T vertex_distance = - point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_candi); - q_obj.set_g_add(vertex_distance); - scan_neighbors(q_obj, cur_candi, est_dist.data(), tmp_pool, vis, cur_degree); - if (cur_candi != cur_id) { - results.emplace_back(cur_candi, vertex_distance); - } - } -} - -// based on new neighbor lists to update quantization code and factors -template -inline void QuantizedGraph::update_qg( - PID cur_id, const std::vector>& new_neighbors -) { - size_t cur_degree = new_neighbors.size(); - - if (cur_degree == 0) { - return; - } - // copy neighbors - auto neighbors = get_neighbors(cur_id); - for (size_t i = 0; i < cur_degree; ++i) { - neighbors[i] = new_neighbors[i].id; - } - - // rotated data - std::vector rotated_data(cur_degree * padded_dim_); - std::vector rotated_centroid(padded_dim_); - for (size_t i = 0; i < cur_degree; ++i) { - if (quantization_bits_ == 0) { - const T* neighbor_vec = get_vector(new_neighbors[i].id); - this->rotator_->rotate(neighbor_vec, &rotated_data[i * padded_dim_]); - } else { - reconstruct_quantized_vector( - new_neighbors[i].id, &rotated_data[i * padded_dim_] - ); - } - } - if (quantization_bits_ == 0) { - this->rotator_->rotate(get_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); - 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( - rotated_data.data() + (i * padded_dim_), - rotated_centroid.data(), - 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) - ); - } - - batch_data += QGBatchDataMap::data_bytes(padded_dim_); - } -} } // namespace rabitqlib::symqg diff --git a/include/rabitqlib/index/symqg/qg_builder.hpp b/include/rabitqlib/index/symqg/qg_builder.hpp index 80f63cc..a9abdab 100644 --- a/include/rabitqlib/index/symqg/qg_builder.hpp +++ b/include/rabitqlib/index/symqg/qg_builder.hpp @@ -32,7 +32,7 @@ class QGBuilder { friend struct QGConstructionTestAccess; private: - QuantizedGraph& qg_; + QuantizedGraph& qg_; size_t ef_build_; // size of search pool for indexing size_t num_threads_; // number of threads used for indexing size_t num_nodes_; // num of data points @@ -59,7 +59,7 @@ class QGBuilder { void initialize_storage(const float* data); - QGBuilder(QuantizedGraph& index, uint32_t ef_build, size_t num_threads) + QGBuilder(QuantizedGraph& index, uint32_t ef_build, size_t num_threads) : qg_{index} , ef_build_{ef_build} , num_threads_{std::max(1, std::min(num_threads, total_threads()))} @@ -69,7 +69,7 @@ class QGBuilder { public: explicit QGBuilder( - QuantizedGraph& index, + QuantizedGraph& index, uint32_t ef_build, const float* data, size_t num_threads = std::numeric_limits::max(), @@ -220,7 +220,7 @@ inline void QGBuilder::initialize_storage(const float* data) { PID entry_point = 0; if (qg_.is_quantized()) { - QuantizedQuery query( + QuantizedQuery query( qg_.centroid_.data(), qg_.centroid_.data(), qg_.padded_dim_, qg_.metric_type_ ); float best = std::numeric_limits::max(); @@ -257,7 +257,7 @@ inline void QGBuilder::add_pruned_edges( } std::vector reconstructed; - std::optional> prepared; + std::optional prepared; while (new_result.size() < degree_bound_ && start < pruned_list.size()) { const auto& cur = pruned_list[start]; bool occlude = false; @@ -314,7 +314,7 @@ inline void QGBuilder::heuristic_prune( size_t start = 0; // start position std::vector reconstructed; - std::optional> prepared; + std::optional prepared; while (pruned_results.size() < degree_bound_ && start < poolsize) { auto candidate_id = pool[start].id; @@ -371,7 +371,7 @@ inline void QGBuilder::search_new_neighbors(bool refine) { // their scores on demand, after the caller can release the input/CSR. if (new_neighbors_[cur_id].empty() && degrees_[cur_id] != 0) { std::vector reconstructed; - std::optional> prepared; + std::optional prepared; const float* source = qg_.prepare_build_query(cur_id, reconstructed, prepared); const auto ids = qg_.get_neighbors(cur_id); for (size_t j = 0; j < degrees_[cur_id]; ++j) { @@ -444,7 +444,7 @@ inline void QGBuilder::add_reverse_edges(bool refine) { if (qg_.is_quantized() && !tmp_pool.empty()) { // RaBitQ estimates are directional: score destination -> source afresh. std::vector reconstructed; - std::optional> prepared; + std::optional prepared; qg_.prepare_build_query(data_id, reconstructed, prepared); for (auto& candidate : tmp_pool) { candidate.distance = qg_.quantized_distance(*prepared, candidate.id); @@ -474,7 +474,7 @@ inline void QGBuilder::random_init() { } std::vector reconstructed; - std::optional> prepared; + std::optional prepared; const float* cur_data = qg_.prepare_build_query(i, reconstructed, prepared); new_neighbors_[i].reserve(degree_bound_); for (PID cur_neigh : neighbor_set) { @@ -536,7 +536,7 @@ inline void QGBuilder::graph_refine() { ids.emplace(neighbor.id); } std::vector reconstructed; - std::optional> prepared; + std::optional prepared; const float* source = qg_.prepare_build_query(i, reconstructed, prepared); while (new_result.size() < degree_bound_) { PID rand_id = rand_integer(0, static_cast(num_nodes_) - 1); diff --git a/include/rabitqlib/quantization/rabitq.hpp b/include/rabitqlib/quantization/rabitq.hpp index 5b4176f..0c91a81 100644 --- a/include/rabitqlib/quantization/rabitq.hpp +++ b/include/rabitqlib/quantization/rabitq.hpp @@ -1,6 +1,7 @@ #pragma once #include +#include #include #include "rabitqlib/defines.hpp" diff --git a/include/rabitqlib/simd/dispatch.hpp b/include/rabitqlib/simd/dispatch.hpp index 7b6a5d8..0885760 100644 --- a/include/rabitqlib/simd/dispatch.hpp +++ b/include/rabitqlib/simd/dispatch.hpp @@ -1,12 +1,12 @@ #pragma once #include - -#include "rabitqlib/utils/space.hpp" +#include +#include namespace rabitqlib::simd { -using ExcodeIpTable = std::array; +using ExcodeIpTable = std::array; ExcodeIpTable resolve_excode_ip_table(); diff --git a/include/rabitqlib/simd/estimator_dispatch.hpp b/include/rabitqlib/simd/estimator_dispatch.hpp new file mode 100644 index 0000000..065b1ae --- /dev/null +++ b/include/rabitqlib/simd/estimator_dispatch.hpp @@ -0,0 +1,103 @@ +#pragma once + +#include +#include + +namespace rabitqlib { +template +class SplitBatchQuery; +template +class BatchQuery; + +namespace simd { +void split_batch_estdist( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +); + +void split_batch_estdist_generic( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +); + +void split_batch_estdist_avx2( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +); + +void split_batch_estdist_avx512( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +); + +void qg_batch_estdist( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +); +void qg_batch_estdist_generic( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +); +void qg_batch_estdist_avx2( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +); +void qg_batch_estdist_avx512( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +); + +uint32_t qg_candidate_mask_generic(const float* distances, float threshold); +uint32_t qg_candidate_mask_avx2(const float* distances, float threshold); + +uint32_t qg_batch_estdist_mask( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +); +uint32_t qg_batch_estdist_mask_generic( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +); +uint32_t qg_batch_estdist_mask_avx2( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +); +} // namespace simd +} // namespace rabitqlib diff --git a/include/rabitqlib/simd/hnsw_dispatch.hpp b/include/rabitqlib/simd/hnsw_dispatch.hpp new file mode 100644 index 0000000..231daff --- /dev/null +++ b/include/rabitqlib/simd/hnsw_dispatch.hpp @@ -0,0 +1,29 @@ +#pragma once + +#include +#include +#include + +#include "rabitqlib/defines.hpp" + +namespace rabitqlib::hnsw { +class HierarchicalNSW; +namespace detail { +std::priority_queue> search_knn( + HierarchicalNSW&, const float*, size_t +); + +std::priority_queue> search_knn_avx2( + HierarchicalNSW&, const float*, size_t +); + +std::priority_queue> search_knn_avx512_core( + HierarchicalNSW&, const float*, size_t +); + +std::priority_queue> search_knn_avx512_popcnt( + HierarchicalNSW&, const float*, size_t +); + +} // namespace detail +} // namespace rabitqlib::hnsw diff --git a/include/rabitqlib/simd/matrix_dispatch.hpp b/include/rabitqlib/simd/matrix_dispatch.hpp new file mode 100644 index 0000000..aadbd8c --- /dev/null +++ b/include/rabitqlib/simd/matrix_dispatch.hpp @@ -0,0 +1,123 @@ +#pragma once + +#include + +namespace rabitqlib::simd { +// Dense row-major float32 kernels; inputs and outputs must not overlap. +// These preserve the caller's Eigen operations and surrounding thread scheduling. +void matrix_product( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_generic( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_avx2( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_avx512( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_transposed( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_transposed_generic( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_transposed_avx2( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void matrix_product_transposed_avx512( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +); + +void row_norms(const float* data, float* result, size_t rows, size_t dim); + +void row_norms_generic(const float* data, float* result, size_t rows, size_t dim); + +void row_norms_avx2(const float* data, float* result, size_t rows, size_t dim); + +void row_norms_avx512(const float* data, float* result, size_t rows, size_t dim); + +void pairwise_distances_lower( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +); + +void pairwise_distances_lower_generic( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +); + +void pairwise_distances_lower_avx2( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +); + +void pairwise_distances_lower_avx512( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +); +} // namespace rabitqlib::simd diff --git a/include/rabitqlib/simd/pack_excode_dispatch.hpp b/include/rabitqlib/simd/pack_excode_dispatch.hpp index 91002d4..f61ed9e 100644 --- a/include/rabitqlib/simd/pack_excode_dispatch.hpp +++ b/include/rabitqlib/simd/pack_excode_dispatch.hpp @@ -5,25 +5,40 @@ namespace rabitqlib::simd { +void packing_2bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + +void packing_3bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + +void packing_4bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + +void packing_5bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + +void packing_6bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + +void packing_7bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_2bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_3bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_4bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_5bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_6bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_7bit_excode_avx2(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); void packing_2bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_3bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_4bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_5bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); + void packing_6bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_7bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_2bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_3bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_4bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_5bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_6bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); -void packing_7bit_excode(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); +void packing_7bit_excode_avx512(const uint8_t* o_raw, uint8_t* o_compact, size_t dim); } // namespace rabitqlib::simd diff --git a/include/rabitqlib/simd/quantization_dispatch.hpp b/include/rabitqlib/simd/quantization_dispatch.hpp index bdf3b21..d17cfe2 100644 --- a/include/rabitqlib/simd/quantization_dispatch.hpp +++ b/include/rabitqlib/simd/quantization_dispatch.hpp @@ -10,6 +10,9 @@ namespace rabitqlib::simd { double best_rescale_factor( const float* magnitudes, size_t dim, int max_code, double start, double end ); +double best_rescale_factor_generic( + const float* magnitudes, size_t dim, int max_code, double start, double end +); double best_rescale_factor_avx2( const float* magnitudes, size_t dim, int max_code, double start, double end ); diff --git a/include/rabitqlib/simd/rotator_dispatch.hpp b/include/rabitqlib/simd/rotator_dispatch.hpp index eec30c9..8003627 100644 --- a/include/rabitqlib/simd/rotator_dispatch.hpp +++ b/include/rabitqlib/simd/rotator_dispatch.hpp @@ -5,12 +5,46 @@ namespace rabitqlib::simd { +void flip_sign(const uint8_t* flip, float* data, size_t dim); + +void kacs_walk(float* data, size_t len); + +void fht_rotate( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +); + void flip_sign_avx2(const uint8_t* flip, float* data, size_t dim); -void flip_sign_avx512(const uint8_t* flip, float* data, size_t dim); + void kacs_walk_avx2(float* data, size_t len); + +void fht_rotate_avx2( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +); + +void flip_sign_avx512(const uint8_t* flip, float* data, size_t dim); + void kacs_walk_avx512(float* data, size_t len); -void flip_sign(const uint8_t* flip, float* data, size_t dim); -void kacs_walk(float* data, size_t len); +void fht_rotate_avx512( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +); } // namespace rabitqlib::simd diff --git a/include/rabitqlib/utils/rotator.hpp b/include/rabitqlib/utils/rotator.hpp index 4c401a2..18bbea9 100644 --- a/include/rabitqlib/utils/rotator.hpp +++ b/include/rabitqlib/utils/rotator.hpp @@ -6,7 +6,6 @@ #include #include #include -#include #include #include #include @@ -14,8 +13,8 @@ #include #include "rabitqlib/defines.hpp" +#include "rabitqlib/simd/matrix_dispatch.hpp" #include "rabitqlib/simd/rotator_dispatch.hpp" -#include "rabitqlib/utils/fht_avx.hpp" #include "rabitqlib/utils/space.hpp" #include "rabitqlib/utils/tools.hpp" @@ -113,9 +112,18 @@ class MatrixRotator : public Rotator { } void rotate(const T* vec, T* rotated_vec) const override { - ConstRowMajorMatrixMap v(vec, 1, this->dim_); - RowMajorMatrixMap rv(rotated_vec, 1, this->padded_dim_); - rv = v * this->rand_mat_; + if constexpr (std::is_same_v) { + // Match Eigen's alias-safe assignment: callers may rotate in place. + std::vector result(this->padded_dim_); + simd::matrix_product( + vec, rand_mat_.data(), result.data(), 1, this->dim_, this->padded_dim_ + ); + std::copy(result.begin(), result.end(), rotated_vec); + } else { + ConstRowMajorMatrixMap v(vec, 1, this->dim_); + RowMajorMatrixMap rv(rotated_vec, 1, this->padded_dim_); + rv = v * this->rand_mat_; + } } }; @@ -126,7 +134,6 @@ static inline void flip_sign(const uint8_t* flip, float* data, size_t dim) { class FhtKacRotator : public Rotator { private: std::vector flip_; - std::function fht_float_ = helper_float_6; size_t trunc_dim_ = 0; float fac_ = 0; @@ -152,28 +159,8 @@ class FhtKacRotator : public Rotator { trunc_dim_ = 1 << bottom_log_dim; fac_ = 1.0F / std::sqrt(static_cast(trunc_dim_)); - switch (bottom_log_dim) { - case 6: - this->fht_float_ = helper_float_6; - break; - case 7: - this->fht_float_ = helper_float_7; - break; - case 8: - this->fht_float_ = helper_float_8; - break; - case 9: - this->fht_float_ = helper_float_9; - break; - case 10: - this->fht_float_ = helper_float_10; - break; - case 11: - this->fht_float_ = helper_float_11; - break; - default: - // TODO(lib): should we do more? - throw std::invalid_argument("Unsupported dimension for FhtKacRotator"); + if (bottom_log_dim < 6 || bottom_log_dim > 11) { + throw std::invalid_argument("Unsupported dimension for FhtKacRotator"); } } FhtKacRotator() = default; @@ -207,7 +194,6 @@ class FhtKacRotator : public Rotator { this->dim_ = other.dim_; this->padded_dim_ = other.padded_dim_; this->flip_ = other.flip_; - this->fht_float_ = other.fht_float_; this->trunc_dim_ = other.trunc_dim_; this->fac_ = other.fac_; return *this; @@ -216,58 +202,9 @@ class FhtKacRotator : public Rotator { static void kacs_walk(float* data, size_t len) { simd::kacs_walk(data, len); } void rotate(const float* data, float* rotated_vec) const override { - std::memcpy(rotated_vec, data, sizeof(float) * dim_); - std::fill(rotated_vec + dim_, rotated_vec + padded_dim_, 0); - - if (trunc_dim_ == padded_dim_) { - flip_sign(flip_.data(), rotated_vec, padded_dim_); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - - flip_sign(flip_.data() + (padded_dim_ / kByteLen), rotated_vec, padded_dim_); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - - flip_sign( - flip_.data() + (2 * padded_dim_ / kByteLen), rotated_vec, padded_dim_ - ); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - - flip_sign( - flip_.data() + (3 * padded_dim_ / kByteLen), rotated_vec, padded_dim_ - ); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - - return; - } - - size_t start = padded_dim_ - trunc_dim_; - - flip_sign(flip_.data(), rotated_vec, padded_dim_); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - kacs_walk(rotated_vec, padded_dim_); - - flip_sign(flip_.data() + (padded_dim_ / kByteLen), rotated_vec, padded_dim_); - fht_float_(rotated_vec + start); - vec_rescale(rotated_vec + start, trunc_dim_, fac_); - kacs_walk(rotated_vec, padded_dim_); - - flip_sign(flip_.data() + (2 * padded_dim_ / kByteLen), rotated_vec, padded_dim_); - fht_float_(rotated_vec); - vec_rescale(rotated_vec, trunc_dim_, fac_); - kacs_walk(rotated_vec, padded_dim_); - - flip_sign(flip_.data() + (3 * padded_dim_ / kByteLen), rotated_vec, padded_dim_); - fht_float_(rotated_vec + start); - vec_rescale(rotated_vec + start, trunc_dim_, fac_); - kacs_walk(rotated_vec, padded_dim_); - - // This can be removed if we don't care about the absolute value of - // similarities. - vec_rescale(rotated_vec, padded_dim_, 0.25F); + simd::fht_rotate( + data, rotated_vec, dim_, padded_dim_, trunc_dim_, fac_, flip_.data() + ); } }; } // namespace rotator_impl diff --git a/python_bindings/symqg_bindings.cpp b/python_bindings/symqg_bindings.cpp index 9a54bf0..60dfe48 100644 --- a/python_bindings/symqg_bindings.cpp +++ b/python_bindings/symqg_bindings.cpp @@ -66,7 +66,7 @@ class SymqgIndex { } num_points_ = static_cast(data_array.shape(0)); - index_ = std::make_unique>( + index_ = std::make_unique( num_points_, dim_, max_degree_, @@ -167,7 +167,7 @@ class SymqgIndex { static SymqgIndex load(const std::string& path) { SymqgIndex wrapper; - wrapper.index_ = std::make_unique>(); + wrapper.index_ = std::make_unique(); wrapper.index_->load(path.c_str()); wrapper.num_points_ = wrapper.index_->num_vertices(); wrapper.dim_ = wrapper.index_->dimension(); @@ -194,7 +194,7 @@ class SymqgIndex { rabitqlib::MetricType metric_ = rabitqlib::METRIC_L2; size_t quantization_bits_ = 0; bool built_ = false; - std::unique_ptr> index_; + std::unique_ptr index_; }; } // namespace rabitqlib::python_bindings diff --git a/sample/cpp/symqg_indexing.cpp b/sample/cpp/symqg_indexing.cpp index 6853bab..2a8348f 100644 --- a/sample/cpp/symqg_indexing.cpp +++ b/sample/cpp/symqg_indexing.cpp @@ -12,7 +12,7 @@ #include "rabitqlib/utils/stopw.hpp" using PID = rabitqlib::PID; -using index_type = rabitqlib::symqg::QuantizedGraph; +using index_type = rabitqlib::symqg::QuantizedGraph; using data_type = rabitqlib::RowMajorArray; using gt_type = rabitqlib::RowMajorArray; diff --git a/sample/cpp/symqg_querying.cpp b/sample/cpp/symqg_querying.cpp index 4453bba..8ea343b 100644 --- a/sample/cpp/symqg_querying.cpp +++ b/sample/cpp/symqg_querying.cpp @@ -7,7 +7,7 @@ #include "rabitqlib/utils/stopw.hpp" using PID = rabitqlib::PID; -using index_type = rabitqlib::symqg::QuantizedGraph; +using index_type = rabitqlib::symqg::QuantizedGraph; using data_type = rabitqlib::RowMajorArray; using gt_type = rabitqlib::RowMajorArray; diff --git a/scripts/check-includes.sh b/scripts/check-includes.sh index 1075f9a..e130c7d 100755 --- a/scripts/check-includes.sh +++ b/scripts/check-includes.sh @@ -21,6 +21,8 @@ fi check_file() { local file="$1" + # Tracked deletions remain in git ls-files until staged. + [[ -f "$file" ]] || return 0 local status=0 local report report="$(mktemp)" @@ -59,7 +61,8 @@ export -f check_file # The child shell expands its positional argument. # shellcheck disable=SC2016 -git ls-files -z -- 'src/*.cpp' 'src/*.hpp' 'include/rabitqlib/*.hpp' \ +git ls-files --cached --others --exclude-standard -z -- 'src/*.cpp' 'src/*.hpp' 'include/rabitqlib/*.hpp' \ ':(exclude)include/rabitqlib/third/**' \ ':(exclude)include/rabitqlib/utils/fht_avx.hpp' \ + | sort -zu \ | xargs -0 -r -n 1 -P "${INCLUDE_JOBS:-2}" bash -c 'check_file "$1"' _ diff --git a/scripts/check-tidy.sh b/scripts/check-tidy.sh index cc7125d..220e511 100755 --- a/scripts/check-tidy.sh +++ b/scripts/check-tidy.sh @@ -62,7 +62,7 @@ for include_dir in "${system_include_dirs[@]}"; do done mapfile -t first_party_headers < <( - git -C "$repo_root" ls-files -- '*.h' '*.hpp' \ + git -C "$repo_root" ls-files --cached --others --exclude-standard -- '*.h' '*.hpp' \ ':(exclude)include/rabitqlib/third/**' \ ':(exclude)include/rabitqlib/utils/fht_avx.hpp' ) diff --git a/src/index/ivf_search.cpp b/src/index/ivf_search.cpp new file mode 100644 index 0000000..96174f6 --- /dev/null +++ b/src/index/ivf_search.cpp @@ -0,0 +1,17 @@ +#include + +#include "rabitqlib/defines.hpp" +#include "rabitqlib/index/ivf/ivf.hpp" +#include "rabitqlib/utils/buffer.hpp" + +namespace rabitqlib::ivf::detail { +// Keep the scalar top-k loop out of the large inlined IVF search body. This avoids +// register spills in wheel builds without changing candidate order or tie handling. +void insert_candidates( + buffer::SearchBuffer& knns, const PID* ids, const float* distances, size_t count +) { + for (size_t i = 0; i < count; ++i) { + knns.insert(ids[i], distances[i]); + } +} +} // namespace rabitqlib::ivf::detail diff --git a/src/index/qg.cpp b/src/index/qg.cpp new file mode 100644 index 0000000..05d4bde --- /dev/null +++ b/src/index/qg.cpp @@ -0,0 +1,796 @@ +#include "rabitqlib/index/symqg/qg.hpp" + +#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/query.hpp" +#include "rabitqlib/quantization/data_layout.hpp" +#include "rabitqlib/quantization/pack_excode.hpp" +#include "rabitqlib/quantization/rabitq.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" +#include "rabitqlib/utils/buffer.hpp" +#include "rabitqlib/utils/io.hpp" +#include "rabitqlib/utils/memory.hpp" +#include "rabitqlib/utils/rotator.hpp" +#include "rabitqlib/utils/space.hpp" +#include "rabitqlib/utils/visited_set.hpp" + +namespace rabitqlib::symqg { + +QuantizedQuery::QuantizedQuery( + const float* rotated_query, + const float* centroid, + size_t padded_dim, + MetricType metric_type +) + : rotated_query_(rotated_query) { + k1xsumq_ = std::accumulate(rotated_query, rotated_query + padded_dim, 0.0F) / -2; + g_add_ = metric_type == METRIC_IP ? -dot_product(rotated_query, centroid, padded_dim) + : euclidean_sqr(rotated_query, centroid, padded_dim); +} +const float* QuantizedQuery::rotated_query() const { return rotated_query_; } +float QuantizedQuery::k1xsumq() const { return k1xsumq_; } +float QuantizedQuery::g_add() const { return g_add_; } + +size_t QuantizedGraph::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; +} + +size_t QuantizedGraph::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; +} + +size_t QuantizedGraph::padded_dimension(size_t dim) { + return (checked_add(dim, 63) / 64) * 64; +} + +char* QuantizedGraph::get_row_data(PID data_id) { + return reinterpret_cast(get_vector(data_id)); +} + +const char* QuantizedGraph::get_row_data(PID data_id) const { + return reinterpret_cast(get_vector(data_id)); +} + +float* QuantizedGraph::get_vector(PID data_id) { + return data_.data() + ((row_offset_ / sizeof(float)) * data_id); +} + +const float* QuantizedGraph::get_vector(PID data_id) const { + return data_.data() + ((row_offset_ / sizeof(float)) * data_id); +} + +char* QuantizedGraph::get_quantized_vector(PID data_id) { return get_row_data(data_id); } + +const char* QuantizedGraph::get_quantized_vector(PID data_id) const { + return get_row_data(data_id); +} + +char* QuantizedGraph::get_batch_data(PID data_id) { + return get_row_data(data_id) + batch_data_offset_; +} + +const char* QuantizedGraph::get_batch_data(PID data_id) const { + return get_row_data(data_id) + batch_data_offset_; +} + +rabitqlib::detail::PackedArrayView QuantizedGraph::get_neighbors(PID data_id) { + return rabitqlib::detail::PackedArrayView( + get_row_data(data_id) + neighbor_offset_ + ); +} + +rabitqlib::detail::ConstPackedArrayView QuantizedGraph::get_neighbors(PID data_id +) const { + return rabitqlib::detail::ConstPackedArrayView( + get_row_data(data_id) + neighbor_offset_ + ); +} + +size_t QuantizedGraph::num_vertices() const { return this->num_points_; } + +size_t QuantizedGraph::dimension() const { return this->dim_; } + +size_t QuantizedGraph::degree_bound() const { return this->degree_bound_; } + +PID QuantizedGraph::entry_point() const { return this->entry_point_; } + +MetricType QuantizedGraph::metric_type() const { return this->metric_type_; } + +size_t QuantizedGraph::quantization_bits() const { return this->quantization_bits_; } + +bool QuantizedGraph::is_quantized() const { return quantization_bits_ != 0; } + +void QuantizedGraph::set_ep(PID entry) { + if (entry >= num_points_) { + throw std::invalid_argument("QuantizedGraph entry point is out of range"); + } + entry_point_ = entry; +} + +QuantizedGraph::QuantizedGraph() = default; + +QuantizedGraph::~QuantizedGraph() = default; + +QuantizedGraph::QuantizedGraph(QuantizedGraph&&) noexcept = default; + +QuantizedGraph& QuantizedGraph::operator=(QuantizedGraph&&) noexcept = default; + +QuantizedGraph::QuantizedGraph( + 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) + , dim_(dim) + , padded_dim_(dim) + , raw_dist_func_((metric_type == METRIC_IP) ? dot_product_dis : euclidean_sqr) + , metric_type_(metric_type) + , rotator_type_(rotator_type) + , quantization_bits_(quantization_bits) { + validate_configuration(); + initialize(); +} + +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" + ); + } + if (degree_bound_ >= num_points_) { + throw std::invalid_argument( + "QuantizedGraph degree bound must be smaller than the number of points" + ); + } + if (num_points_ > buffer::kSearchBufferMaxPointCount) { + throw std::invalid_argument( + "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" + ); + } +} + +void QuantizedGraph::copy_vectors(const float* data, size_t num_threads) { + const int thread_count = static_cast(num_threads); + if (quantization_bits_ != 0) { + if (centroid_.size() != padded_dim_) { + throw std::logic_error("qg-quant centroid must be set before copying vectors"); + } +#pragma omp parallel num_threads(thread_count) + { + 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_ + ); + float f_add; + float f_rescale; + float unused_f_error = 0; + quant::quantize_full_single( + rotated_data.data(), + centroid_.data(), + padded_dim_, + quantization_bits_, + quantized_data.data(), + 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(), padded_dim_, quantization_bits_ + ); + } + } + return; + } +#pragma omp parallel for schedule(dynamic) num_threads(thread_count) + for (size_t i = 0; i < num_points_; ++i) { + const float* src = data + (dim_ * i); + float* dst = get_vector(i); + std::copy(src, src + dim_, dst); + } +} + +void QuantizedGraph::set_quantization_centroid(const float* centroid) { + if (quantization_bits_ == 0) { + return; + } + centroid_.resize(padded_dim_); + rotator_->rotate(centroid, centroid_.data()); +} + +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; + 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)); + output.write(reinterpret_cast(&dim_), sizeof(size_t)); + output.write(reinterpret_cast(&padded_dim_), sizeof(size_t)); + 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()), + static_cast(padded_dim_ * sizeof(float)) + ); + } + + /* Data */ + output.write( + get_row_data(0), + static_cast(checked_multiply(num_points_, row_offset_)) + ); + + /* Rotator */ + this->rotator_->save(output); + + output.flush(); + output.close(); +} + +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"); + } + + std::ifstream input(filename, std::ios::binary); + if (!input.is_open()) { + 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; + read_value(magic, "format marker"); + if (magic == kFormatMagic) { + uint32_t version = 0; + read_value(version, "format 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); + } + + 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) { + read_value(loaded.quantization_bits_, "quantization bits"); + } else { + loaded.quantization_bits_ = 0; + } + + 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(); + + const size_t centroid_bytes = loaded.quantization_bits_ == 0 + ? 0 + : checked_multiply(loaded.padded_dim_, sizeof(float)); + 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(float), 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); + + 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"); + } + + loaded.initialize(); + + 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); +} + +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; +} + +void QuantizedGraph::search( + const float* __restrict__ query, + uint32_t k, + uint32_t* __restrict__ results, + float* __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); + BatchQuery q_obj(rotated_query.data(), padded_dim_, metric_type_); + + buffer::SearchBuffer search_pool(ef_); + // init search buffer + search_pool.insert(this->entry_point_, std::numeric_limits::max()); + + buffer::SearchBuffer res_pool(k); // result buffer + thread_local VisitedSet visited; + thread_local size_t visited_size = 0; + if (visited_size != num_points_) { + visited.initialize(num_points_, num_points_ / 10); + visited_size = num_points_; + } + visited.clear(); + auto* vis = &visited; + + std::vector est_dist(degree_bound_); // estimated distances + + while (search_pool.has_next()) { + PID cur_node = search_pool.pop(); + if (vis->get(cur_node)) { + continue; + } + vis->set(cur_node); + + const float vertex_distance = + point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_node); + q_obj.set_g_add(vertex_distance); + + scan_neighbors( + q_obj, cur_node, est_dist.data(), search_pool, *vis, this->degree_bound_ + ); + res_pool.insert(cur_node, vertex_distance); + } + + 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); +} + +void QuantizedGraph::prepare_query( + const float* 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_ + ); + } +} + +float QuantizedGraph::point_distance( + const float* 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). +void QuantizedGraph::scan_neighbors( + const BatchQuery& q_obj, + PID data_id, + float* est_dist, + buffer::SearchBuffer& search_pool, + VisitedSet& vis, + size_t cur_degree +) const { + const auto* batch_data = get_batch_data(data_id); + const auto neighbors = get_neighbors(data_id); + for (size_t begin = 0; begin < cur_degree; begin += fastscan::kBatchSize) { + const float threshold = search_pool.top_dist(); + uint32_t candidate_mask = simd::qg_batch_estdist_mask( + batch_data, q_obj, padded_dim_, est_dist + begin, threshold + ); + batch_data += QGBatchDataMap::data_bytes(padded_dim_); + + // Construction can leave a partial batch with stale IDs in its unused lanes. + const size_t remaining = cur_degree - begin; + if (remaining < fastscan::kBatchSize) { + candidate_mask &= (uint32_t{1} << remaining) - 1; + } + + while (candidate_mask != 0) { + const auto lane = static_cast(__builtin_ctz(candidate_mask)); + candidate_mask &= candidate_mask - 1; + const size_t i = begin + lane; + PID cur_neighbor = neighbors[i]; + float 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(get_row_data(search_pool.next_id()), 10); + } + } +} + +void QuantizedGraph::update_results( + buffer::SearchBuffer& result_pool, + VisitedSet& vis, + const float* query, + const QuantizedQuery* quantized_query +) { + if (result_pool.is_full()) { + return; + } + + 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 = neighbors[i]; + if (!vis.get(cur_neighbor)) { + vis.set(cur_neighbor); + result_pool.insert( + cur_neighbor, point_distance(query, quantized_query, cur_neighbor) + ); + } + } + if (result_pool.is_full()) { + break; + } + } +} + +// initialize const offsets & data array +void QuantizedGraph::initialize_layout() { + if (quantization_bits_ == 0) { + batch_data_offset_ = checked_multiply(dim_, sizeof(float)); + } else { + const size_t code_bits = checked_multiply(padded_dim_, quantization_bits_); + batch_data_offset_ = checked_add(code_bits / 8, checked_multiply(sizeof(float), 2)); + } + + const size_t binary_batch_bytes = + checked_multiply(padded_dim_, fastscan::kBatchSize) / 8; + const size_t factor_bytes = + checked_multiply(checked_multiply(sizeof(float), 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) + ); + row_offset_ = + checked_add(neighbor_offset_, checked_multiply(degree_bound_, sizeof(PID))); +} + +void QuantizedGraph::initialize() { + padded_dim_ = padded_dimension(dim_); + rotator_.reset(choose_rotator(dim_, rotator_type_, padded_dim_)); + + assert(padded_dim_ % 64 == 0); + assert(padded_dim_ >= dim_); + + initialize_layout(); + + const size_t data_bytes = checked_multiply(num_points_, row_offset_); + assert(row_offset_ % sizeof(float) == 0); + data_ = RowStorage(data_bytes / sizeof(float)); + + if (quantization_bits_ != 0) { + quantized_ip_func_ = select_excode_ipfunc(quantization_bits_); + } +} + +float QuantizedGraph::quantized_distance(const QuantizedQuery& query, PID data_id) const { + 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() + ); +} + +void QuantizedGraph::reconstruct_quantized_vector(PID data_id, float* 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_ + ); +} + +// Construction sources come from owned raw rows or from the existing RaBitQ codes. +// Reconstructed sources are already rotated; never pass them through prepare_query. +const float* QuantizedGraph::prepare_build_query( + PID id, std::vector& rotated, std::optional& prepared +) const { + if (is_quantized()) { + rotated.resize(padded_dim_); + reconstruct_quantized_vector(id, rotated.data()); + prepared.emplace(rotated.data(), centroid_.data(), padded_dim_, metric_type_); + return rotated.data(); + } + prepared.reset(); + return get_vector(id); +} + +// find candidate neighbors for cur_id, exclude the vertex itself +void QuantizedGraph::find_candidates( + PID cur_id, + size_t search_ef, + std::vector>& results, + VisitedSet& vis, + const std::vector& degrees +) const { + std::vector rotated_query(padded_dim_); + std::optional quantized_query; + const float* query = prepare_build_query(cur_id, rotated_query, quantized_query); + if (!is_quantized()) { + rotator_->rotate(query, rotated_query.data()); + } + BatchQuery q_obj(rotated_query.data(), padded_dim_, metric_type_); + + // insert entry point to initialize search buffer + buffer::SearchBuffer tmp_pool(search_ef); + 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 + while (tmp_pool.has_next()) { + auto cur_candi = tmp_pool.pop(); + if (vis.get(cur_candi)) { + continue; + } + vis.set(cur_candi); + auto cur_degree = degrees[cur_candi]; + const float vertex_distance = + point_distance(query, quantized_query ? &*quantized_query : nullptr, cur_candi); + q_obj.set_g_add(vertex_distance); + scan_neighbors(q_obj, cur_candi, est_dist.data(), tmp_pool, vis, cur_degree); + if (cur_candi != cur_id) { + results.emplace_back(cur_candi, vertex_distance); + } + } +} + +// based on new neighbor lists to update quantization code and factors +void QuantizedGraph::update_qg( + PID cur_id, const std::vector>& new_neighbors +) { + size_t cur_degree = new_neighbors.size(); + + if (cur_degree == 0) { + return; + } + // copy neighbors + auto neighbors = get_neighbors(cur_id); + for (size_t i = 0; i < cur_degree; ++i) { + neighbors[i] = new_neighbors[i].id; + } + + // rotated data + std::vector rotated_data(cur_degree * padded_dim_); + std::vector rotated_centroid(padded_dim_); + for (size_t i = 0; i < cur_degree; ++i) { + if (quantization_bits_ == 0) { + const float* neighbor_vec = get_vector(new_neighbors[i].id); + this->rotator_->rotate(neighbor_vec, &rotated_data[i * padded_dim_]); + } else { + reconstruct_quantized_vector( + new_neighbors[i].id, &rotated_data[i * padded_dim_] + ); + } + } + if (quantization_bits_ == 0) { + this->rotator_->rotate(get_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); + 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( + rotated_data.data() + (i * padded_dim_), + rotated_centroid.data(), + batch_size, + padded_dim_, + batch_data, + metric_type_ + ); + if (batch_size < fastscan::kBatchSize) { + QGBatchDataMap batch(batch_data, padded_dim_); + std::fill( + batch.f_add() + static_cast(batch_size), + batch.f_add() + fastscan::kBatchSize, + static_cast(0) + ); + std::fill( + batch.f_rescale() + static_cast(batch_size), + batch.f_rescale() + fastscan::kBatchSize, + static_cast(0) + ); + } + + batch_data += QGBatchDataMap::data_bytes(padded_dim_); + } +} +} // namespace rabitqlib::symqg diff --git a/src/simd/dispatch.cpp b/src/simd/dispatch.cpp index 10f9e14..82f2789 100644 --- a/src/simd/dispatch.cpp +++ b/src/simd/dispatch.cpp @@ -2,13 +2,18 @@ #include #include +#include #include #include +#include #include "rabitqlib/defines.hpp" #include "rabitqlib/fastscan/fastscan.hpp" #include "rabitqlib/fastscan/highacc_fastscan.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" #include "rabitqlib/simd/fastscan_dispatch.hpp" +#include "rabitqlib/simd/hnsw_dispatch.hpp" +#include "rabitqlib/simd/matrix_dispatch.hpp" #include "rabitqlib/simd/pack_excode_dispatch.hpp" #include "rabitqlib/simd/quantization_dispatch.hpp" #include "rabitqlib/simd/rotator_dispatch.hpp" @@ -17,22 +22,144 @@ #include "rabitqlib/utils/cpu_features.hpp" #include "rabitqlib/utils/space.hpp" #include "rabitqlib/utils/warmup_space.hpp" -#include "rescale_search.hpp" namespace rabitqlib::simd { +namespace { +// Resolve once during static initialization; wrappers never repeat CPU checks. +// The override is used only by kernels needing a stricter AVX-512 subset. +template +Function resolve_kernel( + Function avx512, + Function avx2, + Function fallback, + bool avx512_supported = cpu::has_avx512_core() +) { + if (avx512_supported) { + return avx512; + } + if (cpu::has_avx2()) { + return avx2; + } + return fallback; +} -const auto kEuclideanSqrFn = cpu::has_avx512_core() ? euclidean_sqr_avx512 - : cpu::has_avx2() ? euclidean_sqr_avx2 - : euclidean_sqr_generic; -const auto kDotProductFn = cpu::has_avx512_core() ? dot_product_avx512 - : cpu::has_avx2() ? dot_product_avx2 - : dot_product_generic; -const auto kDotProductDisFn = cpu::has_avx512_core() ? dot_product_dis_avx512 - : cpu::has_avx2() ? dot_product_dis_avx2 - : dot_product_dis_generic; -const auto kL2normSqrFn = cpu::has_avx512_core() ? l2norm_sqr_avx512 - : cpu::has_avx2() ? l2norm_sqr_avx2 - : l2norm_sqr_generic; +template +Function resolve_kernel(Function avx2, Function fallback) { + if (cpu::has_avx2()) { + return avx2; + } + return fallback; +} + +} // namespace + +const auto kMatrixProductFn = + resolve_kernel(matrix_product_avx512, matrix_product_avx2, matrix_product_generic); + +void matrix_product( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + kMatrixProductFn(left, right, result, rows, inner, cols); +} + +const auto kMatrixProductTransposedFn = resolve_kernel( + matrix_product_transposed_avx512, + matrix_product_transposed_avx2, + matrix_product_transposed_generic +); + +void matrix_product_transposed( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + kMatrixProductTransposedFn(left, right, result, rows, inner, cols); +} + +const auto kRowNormsFn = + resolve_kernel(row_norms_avx512, row_norms_avx2, row_norms_generic); + +void row_norms(const float* data, float* result, size_t rows, size_t dim) { + kRowNormsFn(data, result, rows, dim); +} + +const auto kPairwiseDistancesLowerFn = resolve_kernel( + pairwise_distances_lower_avx512, + pairwise_distances_lower_avx2, + pairwise_distances_lower_generic +); + +void pairwise_distances_lower( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +) { + kPairwiseDistancesLowerFn(data, result, norms_data, size, dim, inner_product); +} + +const auto kQgBatchEstdistFn = resolve_kernel( + qg_batch_estdist_avx512, qg_batch_estdist_avx2, qg_batch_estdist_generic +); + +void qg_batch_estdist( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +) { + kQgBatchEstdistFn(batch_data, q_obj, padded_dim, est_distance); +} + +const auto kQgBatchEstdistMaskFn = + resolve_kernel(qg_batch_estdist_mask_avx2, qg_batch_estdist_mask_generic); + +uint32_t qg_batch_estdist_mask( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +) { + return kQgBatchEstdistMaskFn(batch_data, q_obj, padded_dim, est_distance, threshold); +} + +const auto kSplitBatchEstdistFn = resolve_kernel( + split_batch_estdist_avx512, split_batch_estdist_avx2, split_batch_estdist_generic +); + +void split_batch_estdist( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +) { + kSplitBatchEstdistFn( + batch_data, q_obj, padded_dim, est_distance, low_distance, ip_x0_qr, use_hacc + ); +} + +const auto kEuclideanSqrFn = + resolve_kernel(euclidean_sqr_avx512, euclidean_sqr_avx2, euclidean_sqr_generic); +const auto kDotProductFn = + resolve_kernel(dot_product_avx512, dot_product_avx2, dot_product_generic); +const auto kDotProductDisFn = + resolve_kernel(dot_product_dis_avx512, dot_product_dis_avx2, dot_product_dis_generic); +const auto kL2normSqrFn = + resolve_kernel(l2norm_sqr_avx512, l2norm_sqr_avx2, l2norm_sqr_generic); float euclidean_sqr(const float* a, const float* b, size_t dim) { return kEuclideanSqrFn(a, b, dim); @@ -48,22 +175,6 @@ float dot_product_dis(const float* a, const float* b, size_t dim) { float l2norm_sqr(const float* a, size_t dim) { return kL2normSqrFn(a, dim); } -namespace detail { - -RescaleScratch& get_thread_local_rescale_scratch(size_t dim) { - // Search evaluation is synchronous and does not re-enter the quantizer. - // Retain the largest buffers on each worker; smaller vectors overwrite only - // their active prefix without shrinking or zero-initializing the storage. - thread_local RescaleScratch scratch; - if (scratch.magnitudes.size() < dim) - scratch.magnitudes.resize(dim); - if (scratch.reciprocals.size() < dim) - scratch.reciprocals.resize(dim); - return scratch; -} - -} // namespace detail - [[noreturn]] static void missing_feature(const char* feature_name) { throw std::runtime_error( std::string(feature_name) + " requires AVX2/FMA or AVX512 support" @@ -79,20 +190,9 @@ static float ip_fxu0( return 0.0F; } -static double request_scalar_rescale_search(const float*, size_t, int, double, double) { - return -1; -} - -using BestRescaleFactorFn = double (*)(const float*, size_t, int, double, double); -const BestRescaleFactorFn kBestRescaleFactorFn = [] { - if (cpu::has_avx512_core()) { - return best_rescale_factor_avx512; - } else if (cpu::has_avx2()) { - return best_rescale_factor_avx2; - } else { - return request_scalar_rescale_search; - } -}(); +const auto kBestRescaleFactorFn = resolve_kernel( + best_rescale_factor_avx512, best_rescale_factor_avx2, best_rescale_factor_generic +); double best_rescale_factor( const float* magnitudes, size_t dim, int max_code, double start, double end @@ -100,6 +200,26 @@ double best_rescale_factor( return kBestRescaleFactorFn(magnitudes, dim, max_code, start, end); } +static void +missing_fht_rotate(const float*, float*, size_t, size_t, size_t, float, const uint8_t*) { + missing_feature("sign flip"); +} + +const auto kFhtRotateFn = + resolve_kernel(fht_rotate_avx512, fht_rotate_avx2, missing_fht_rotate); + +void fht_rotate( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +) { + kFhtRotateFn(data, rotated_vec, dim, padded_dim, trunc_dim, fac, flip); +} + static float missing_excode_ip(const float*, const uint8_t*, size_t) { missing_feature("excode ip functions"); } @@ -155,8 +275,8 @@ static float missing_warmup_ip_x0_q_512( } ExcodeIpTable resolve_excode_ip_table() { - if (cpu::has_avx512_core()) { - return { + return resolve_kernel( + ExcodeIpTable{ ip_fxu0, excode_ipimpl::ip16_fxu1_avx512, excode_ipimpl::ip64_fxu2_avx512, @@ -166,9 +286,8 @@ ExcodeIpTable resolve_excode_ip_table() { excode_ipimpl::ip64_fxu6_avx512, excode_ipimpl::ip64_fxu7_avx512, excode_ipimpl::ip16_fxu8_avx512, - }; - } else if (cpu::has_avx2()) { - return { + }, + ExcodeIpTable{ ip_fxu0, excode_ipimpl::ip16_fxu1_avx2, excode_ipimpl::ip64_fxu2_avx2, @@ -178,9 +297,8 @@ ExcodeIpTable resolve_excode_ip_table() { excode_ipimpl::ip64_fxu6_avx2, excode_ipimpl::ip64_fxu7_avx2, excode_ipimpl::ip16_fxu8_avx2, - }; - } else { - return { + }, + ExcodeIpTable{ ip_fxu0, missing_excode_ip, missing_excode_ip, @@ -190,78 +308,44 @@ ExcodeIpTable resolve_excode_ip_table() { missing_excode_ip, missing_excode_ip, missing_excode_ip, - }; - } -} - -using FlipSignFn = void (*)(const uint8_t*, float*, size_t); -const FlipSignFn kFlipSignFn = [] { - if (cpu::has_avx512_core()) { - return flip_sign_avx512; - } else if (cpu::has_avx2()) { - return flip_sign_avx2; - } else { - return missing_flip_sign; - } -}(); - -using KacsWalkFn = void (*)(float*, size_t); -const KacsWalkFn kKacsWalkFn = [] { - if (cpu::has_avx512_core()) { - return kacs_walk_avx512; - } else if (cpu::has_avx2()) { - return kacs_walk_avx2; - } else { - return missing_kacs_walk; - } -}(); - -using ScalarQuantizeUint8Fn = void (*)(uint8_t*, const float*, size_t, float, float); -const ScalarQuantizeUint8Fn kScalarQuantizeUint8Fn = [] { - if (cpu::has_avx512_core()) { - return scalar_quantize_uint8_avx512; - } else if (cpu::has_avx2()) { - return scalar_quantize_uint8_avx2; - } else { - return missing_scalar_quantize_uint8; - } -}(); - -using ScalarQuantizeUint16Fn = void (*)(uint16_t*, const float*, size_t, float, float); -const ScalarQuantizeUint16Fn kScalarQuantizeUint16Fn = [] { - if (cpu::has_avx512_core()) { - return scalar_quantize_uint16_avx512; - } else if (cpu::has_avx2()) { - return scalar_quantize_uint16_avx2; - } else { - return missing_scalar_quantize_uint16; - } -}(); - -using PackExcodeFn = void (*)(const uint8_t*, uint8_t*, size_t); - -static PackExcodeFn resolve_pack_excode_fn(PackExcodeFn avx512_fn, PackExcodeFn avx2_fn) { - if (cpu::has_avx512_core()) { - return avx512_fn; - } else if (cpu::has_avx2()) { - return avx2_fn; - } else { - return missing_pack_excode; - } + } + ); } -const PackExcodeFn kPacking2BitExcodeFn = - resolve_pack_excode_fn(packing_2bit_excode_avx512, packing_2bit_excode_avx2); -const PackExcodeFn kPacking3BitExcodeFn = - resolve_pack_excode_fn(packing_3bit_excode_avx512, packing_3bit_excode_avx2); -const PackExcodeFn kPacking4BitExcodeFn = - resolve_pack_excode_fn(packing_4bit_excode_avx512, packing_4bit_excode_avx2); -const PackExcodeFn kPacking5BitExcodeFn = - resolve_pack_excode_fn(packing_5bit_excode_avx512, packing_5bit_excode_avx2); -const PackExcodeFn kPacking6BitExcodeFn = - resolve_pack_excode_fn(packing_6bit_excode_avx512, packing_6bit_excode_avx2); -const PackExcodeFn kPacking7BitExcodeFn = - resolve_pack_excode_fn(packing_7bit_excode_avx512, packing_7bit_excode_avx2); +const auto kFlipSignFn = + resolve_kernel(flip_sign_avx512, flip_sign_avx2, missing_flip_sign); + +const auto kKacsWalkFn = + resolve_kernel(kacs_walk_avx512, kacs_walk_avx2, missing_kacs_walk); + +const auto kScalarQuantizeUint8Fn = resolve_kernel( + scalar_quantize_uint8_avx512, scalar_quantize_uint8_avx2, missing_scalar_quantize_uint8 +); + +const auto kScalarQuantizeUint16Fn = resolve_kernel( + scalar_quantize_uint16_avx512, + scalar_quantize_uint16_avx2, + missing_scalar_quantize_uint16 +); + +const auto kPacking2BitExcodeFn = resolve_kernel( + packing_2bit_excode_avx512, packing_2bit_excode_avx2, missing_pack_excode +); +const auto kPacking3BitExcodeFn = resolve_kernel( + packing_3bit_excode_avx512, packing_3bit_excode_avx2, missing_pack_excode +); +const auto kPacking4BitExcodeFn = resolve_kernel( + packing_4bit_excode_avx512, packing_4bit_excode_avx2, missing_pack_excode +); +const auto kPacking5BitExcodeFn = resolve_kernel( + packing_5bit_excode_avx512, packing_5bit_excode_avx2, missing_pack_excode +); +const auto kPacking6BitExcodeFn = resolve_kernel( + packing_6bit_excode_avx512, packing_6bit_excode_avx2, missing_pack_excode +); +const auto kPacking7BitExcodeFn = resolve_kernel( + packing_7bit_excode_avx512, packing_7bit_excode_avx2, missing_pack_excode +); void flip_sign(const uint8_t* flip, float* data, size_t dim) { kFlipSignFn(flip, data, dim); @@ -319,38 +403,24 @@ const ex_ipfunc kIp64Fxu5AvxFn = kExcodeIpTable[5]; const ex_ipfunc kIp64Fxu6AvxFn = kExcodeIpTable[6]; const ex_ipfunc kIp64Fxu7AvxFn = kExcodeIpTable[7]; -using NewTransposeBinFn = void (*)(const uint16_t*, uint64_t*, size_t, size_t); -const NewTransposeBinFn kNewTransposeBinFn = [] { - if (cpu::has_avx512_core()) { - return simd::new_transpose_bin_avx512; - } else if (cpu::has_avx2()) { - return simd::new_transpose_bin_avx2; - } else { - return simd::missing_new_transpose_bin; - } -}(); - -using NewTransposeBin512Fn = void (*)(const uint8_t*, uint64_t*, size_t, size_t); -const NewTransposeBin512Fn kNewTransposeBin512Fn = [] { - if (cpu::has_avx512_core()) { - return simd::new_transpose_bin_512_avx512; - } else if (cpu::has_avx2()) { - return simd::new_transpose_bin_512_avx2; - } else { - return simd::missing_new_transpose_bin_512; - } -}(); +const auto kNewTransposeBinFn = rabitqlib::simd::resolve_kernel( + simd::new_transpose_bin_avx512, + simd::new_transpose_bin_avx2, + simd::missing_new_transpose_bin +); + +const auto kNewTransposeBin512Fn = rabitqlib::simd::resolve_kernel( + simd::new_transpose_bin_512_avx512, + simd::new_transpose_bin_512_avx2, + simd::missing_new_transpose_bin_512 +); using MaskIpX0QFn = float (*)(const float*, const uint8_t*, size_t); -const MaskIpX0QFn kMaskIpX0QFn = [] { - if (cpu::has_avx512_core()) { - return static_cast(simd::mask_ip_x0_q_avx512); - } else if (cpu::has_avx2()) { - return static_cast(simd::mask_ip_x0_q_avx2); - } else { - return simd::missing_mask_ip_x0_q; - } -}(); +const MaskIpX0QFn kMaskIpX0QFn = rabitqlib::simd::resolve_kernel( + static_cast(simd::mask_ip_x0_q_avx512), + static_cast(simd::mask_ip_x0_q_avx2), + simd::missing_mask_ip_x0_q +); ex_ipfunc select_excode_ipfunc(size_t ex_bits) { if (ex_bits <= 8) { @@ -424,59 +494,32 @@ float mask_ip_x0_q(const float* query, const uint64_t* data, size_t padded_dim) namespace rabitqlib::fastscan { -void simd::pack_lut_generic(size_t dim, const float* query, float* lut) { - for (size_t group = 0; group < dim / 4; ++group) { - lut[0] = 0; - for (size_t j = 1; j < 16; ++j) { - lut[j] = lut[j - LOWBIT(j)] + query[kPos[j]]; - } - query += 4; - lut += 16; - } -} - -using PackLutFn = void (*)(size_t, const float*, float*); -const PackLutFn kPackLutFn = cpu::has_avx512_core() ? simd::pack_lut_avx512 - : cpu::has_avx2() ? simd::pack_lut_avx2 - : simd::pack_lut_generic; +const auto kPackLutFn = rabitqlib::simd::resolve_kernel( + simd::pack_lut_avx512, simd::pack_lut_avx2, simd::pack_lut_generic +); template <> void pack_lut(size_t dim, const float* __restrict__ query, float* __restrict__ lut) { kPackLutFn(dim, query, lut); } -using AccumulateFn = void (*)(const uint8_t*, const uint8_t*, uint16_t*, size_t); -const AccumulateFn kAccumulateFn = [] { - if (cpu::has_avx512_core()) { - return simd::accumulate_avx512; - } else if (cpu::has_avx2()) { - return simd::accumulate_avx2; - } else { - return rabitqlib::simd::missing_fastscan_accumulate; - } -}(); - -using TransferLutHaccFn = void (*)(const uint16_t*, size_t, uint8_t*); -const TransferLutHaccFn kTransferLutHaccFn = [] { - if (cpu::has_avx512_core()) { - return simd::transfer_lut_hacc_avx512; - } else if (cpu::has_avx2()) { - return simd::transfer_lut_hacc_avx2; - } else { - return rabitqlib::simd::missing_fastscan_transfer_lut_hacc; - } -}(); - -using AccumulateHaccFn = void (*)(const uint8_t*, const uint8_t*, int32_t*, size_t); -const AccumulateHaccFn kAccumulateHaccFn = [] { - if (cpu::has_avx512_core()) { - return simd::accumulate_hacc_avx512; - } else if (cpu::has_avx2()) { - return simd::accumulate_hacc_avx2; - } else { - return rabitqlib::simd::missing_fastscan_accumulate_hacc; - } -}(); +const auto kAccumulateFn = rabitqlib::simd::resolve_kernel( + simd::accumulate_avx512, + simd::accumulate_avx2, + rabitqlib::simd::missing_fastscan_accumulate +); + +const auto kTransferLutHaccFn = rabitqlib::simd::resolve_kernel( + simd::transfer_lut_hacc_avx512, + simd::transfer_lut_hacc_avx2, + rabitqlib::simd::missing_fastscan_transfer_lut_hacc +); + +const auto kAccumulateHaccFn = rabitqlib::simd::resolve_kernel( + simd::accumulate_hacc_avx512, + simd::accumulate_hacc_avx2, + rabitqlib::simd::missing_fastscan_accumulate_hacc +); void accumulate( const uint8_t* __restrict__ codes, @@ -519,15 +562,13 @@ namespace rabitqlib { using WarmupIpX0Q512Fn = float (*)(const uint8_t*, const uint64_t*, float, float, size_t, size_t); -const WarmupIpX0Q512Fn kWarmupIpX0Q512Fn = [] { - if (rabitqlib::cpu::has_avx512_popcnt()) { - return static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx512); - } else if (rabitqlib::cpu::has_avx2()) { - return static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx2); - } else { - return rabitqlib::simd::missing_warmup_ip_x0_q_512; - } -}(); +const WarmupIpX0Q512Fn kWarmupIpX0Q512Fn = + rabitqlib::simd::resolve_kernel( + static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx512), + static_cast(rabitqlib::simd::warmup_ip_x0_q_512_avx2), + rabitqlib::simd::missing_warmup_ip_x0_q_512, + cpu::has_avx512_popcnt() + ); float warmup_ip_x0_q_512( const uint8_t* data, @@ -554,3 +595,29 @@ float warmup_ip_x0_q_512( } } // namespace rabitqlib + +namespace rabitqlib::hnsw::detail { +namespace { +using SearchKnnFn = + std::priority_queue> (*)(HierarchicalNSW&, const float*, size_t); +std::priority_queue> missing_search_knn( + HierarchicalNSW&, const float*, size_t +) { + throw std::runtime_error("HNSW search requires AVX2/FMA or AVX512 support"); +} +// The core variant uses AVX2 warmup; the popcount variant has its own stricter tier. +const SearchKnnFn kSearchKnnFn = cpu::has_avx512_popcnt() + ? search_knn_avx512_popcnt + : rabitqlib::simd::resolve_kernel( + search_knn_avx512_core, + search_knn_avx2, + missing_search_knn, + cpu::has_avx512_core() && cpu::has_avx2() + ); +} // namespace +std::priority_queue> search_knn( + HierarchicalNSW& index, const float* query, size_t topk +) { + return kSearchKnnFn(index, query, topk); +} +} // namespace rabitqlib::hnsw::detail diff --git a/src/simd/estimator_avx2.cpp b/src/simd/estimator_avx2.cpp new file mode 100644 index 0000000..8dafabf --- /dev/null +++ b/src/simd/estimator_avx2.cpp @@ -0,0 +1,61 @@ +#include + +#include +#include + +#include "estimator_kernels.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" + +namespace rabitqlib::simd { +namespace { +uint32_t candidate_mask_impl(const float* distances, float threshold) { + const __m256 limit = _mm256_set1_ps(threshold); + uint32_t mask = 0; + for (size_t lane = 0; lane < 32; lane += 8) { + const __m256 values = _mm256_loadu_ps(distances + lane); + const __m256 candidates = _mm256_cmp_ps(values, limit, _CMP_NGT_UQ); + mask |= static_cast(_mm256_movemask_ps(candidates)) << lane; + } + return mask; +} +} // namespace + +void split_batch_estdist_avx2( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +) { + split_batch_estdist_impl( + batch_data, q_obj, padded_dim, est_distance, low_distance, ip_x0_qr, use_hacc + ); +} + +void qg_batch_estdist_avx2( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +) { + qg_batch_estdist_impl(batch_data, q_obj, padded_dim, est_distance); +} + +uint32_t qg_candidate_mask_avx2(const float* distances, float threshold) { + return candidate_mask_impl(distances, threshold); +} + +uint32_t qg_batch_estdist_mask_avx2( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +) { + qg_batch_estdist_impl(batch_data, q_obj, padded_dim, est_distance); + return candidate_mask_impl(est_distance, threshold); +} + +} // namespace rabitqlib::simd diff --git a/src/simd/estimator_avx512.cpp b/src/simd/estimator_avx512.cpp new file mode 100644 index 0000000..15abe7c --- /dev/null +++ b/src/simd/estimator_avx512.cpp @@ -0,0 +1,30 @@ +#include + +#include "estimator_kernels.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" + +namespace rabitqlib::simd { +void split_batch_estdist_avx512( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +) { + split_batch_estdist_impl( + batch_data, q_obj, padded_dim, est_distance, low_distance, ip_x0_qr, use_hacc + ); +} + +void qg_batch_estdist_avx512( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +) { + qg_batch_estdist_impl(batch_data, q_obj, padded_dim, est_distance); +} + +} // namespace rabitqlib::simd diff --git a/src/simd/estimator_generic.cpp b/src/simd/estimator_generic.cpp new file mode 100644 index 0000000..9050e9f --- /dev/null +++ b/src/simd/estimator_generic.cpp @@ -0,0 +1,56 @@ +#include +#include + +#include "estimator_kernels.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" + +namespace rabitqlib::simd { +namespace { +uint32_t candidate_mask_impl(const float* distances, float threshold) { + uint32_t mask = 0; + for (size_t lane = 0; lane < 32; ++lane) { + mask |= static_cast(!(distances[lane] > threshold)) << lane; + } + return mask; +} +} // namespace + +void split_batch_estdist_generic( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +) { + split_batch_estdist_impl( + batch_data, q_obj, padded_dim, est_distance, low_distance, ip_x0_qr, use_hacc + ); +} + +void qg_batch_estdist_generic( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +) { + qg_batch_estdist_impl(batch_data, q_obj, padded_dim, est_distance); +} + +uint32_t qg_candidate_mask_generic(const float* distances, float threshold) { + return candidate_mask_impl(distances, threshold); +} + +uint32_t qg_batch_estdist_mask_generic( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float threshold +) { + qg_batch_estdist_impl(batch_data, q_obj, padded_dim, est_distance); + return candidate_mask_impl(est_distance, threshold); +} + +} // namespace rabitqlib::simd diff --git a/src/simd/estimator_kernels.hpp b/src/simd/estimator_kernels.hpp new file mode 100644 index 0000000..ba2073b --- /dev/null +++ b/src/simd/estimator_kernels.hpp @@ -0,0 +1,172 @@ +#pragma once + +#include +#include +#include + +#include "rabitqlib/defines.hpp" +#include "rabitqlib/fastscan/fastscan.hpp" +#include "rabitqlib/fastscan/highacc_fastscan.hpp" +#include "rabitqlib/index/query.hpp" +#include "rabitqlib/quantization/data_layout.hpp" + +namespace rabitqlib::simd { +namespace { +// Compile the complete batch estimator at each ISA so accumulation conversion and +// factor correction use the same vector width as the scanning kernels. +inline void split_batch_estdist_impl( + const char* batch_data, + const SplitBatchQuery& q_obj, + size_t padded_dim, + float* est_distance, + float* low_distance, + float* ip_x0_qr, + bool use_hacc +) { + constexpr size_t kSafeChunkDim = 1024; + ConstBatchDataMap cur_batch(batch_data, padded_dim); + std::array accu_values{}; + RowMajorArrayMap accu_arr(accu_values.data(), 1, fastscan::kBatchSize); + const auto* codes_ptr = cur_batch.bin_code(); + const auto* lut_ptr = q_obj.lut(); + if (use_hacc) { + std::array accu_res; + size_t remaining_dim = padded_dim; + + while (remaining_dim > kSafeChunkDim) { + fastscan::accumulate_hacc(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); + codes_ptr += kSafeChunkDim << 2; + lut_ptr += kSafeChunkDim << 3; + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_arr.data()[i] += accu_res[i]; + } + remaining_dim -= kSafeChunkDim; + } + + fastscan::accumulate_hacc(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_arr.data()[i] += accu_res[i]; + } + } else { + std::array accu_res; + size_t remaining_dim = padded_dim; + + while (remaining_dim > kSafeChunkDim) { + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); + codes_ptr += kSafeChunkDim << 2; + lut_ptr += kSafeChunkDim << 2; + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_arr.data()[i] += accu_res[i]; + } + remaining_dim -= kSafeChunkDim; + } + + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_arr.data()[i] += accu_res[i]; + } + } + + 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( + f_rescale_values.data(), 1, fastscan::kBatchSize + ); + ConstRowMajorArrayMap f_error_arr( + f_error_values.data(), 1, fastscan::kBatchSize + ); + + RowMajorArrayMap est_dist_arr(est_distance, 1, fastscan::kBatchSize); + RowMajorArrayMap ip_x0_qr_arr(ip_x0_qr, 1, fastscan::kBatchSize); + RowMajorArrayMap low_dist_arr(low_distance, 1, fastscan::kBatchSize); + + ip_x0_qr_arr = q_obj.delta() * (accu_arr.template cast()) + q_obj.sum_vl_lut(); + + est_dist_arr = + f_add_arr + q_obj.g_add() + f_rescale_arr * (ip_x0_qr_arr + q_obj.k1xsumq()); + + low_dist_arr = est_dist_arr - f_error_arr * q_obj.g_error(); +} + +inline void qg_batch_estdist_impl( + const char* batch_data, + const BatchQuery& q_obj, + size_t padded_dim, + float* est_distance +) { + using T = float; + using TA = uint16_t; + + // Each 4-dimensional codebook can contribute at most 255, so 1024 dimensions + // produce at most 255 * (1024 / 4) = 65280 in the uint16_t FastScan result. + constexpr size_t kSafeChunkDim = 1024; + ConstQGBatchDataMap cur_batch(batch_data, padded_dim); + + if (padded_dim <= kSafeChunkDim) { + std::array accu_res{}; + fastscan::accumulate( + cur_batch.bin_code(), q_obj.lut(), accu_res.data(), padded_dim + ); + + ConstRowMajorArrayMap ip_arr(accu_res.data(), 1, fastscan::kBatchSize); + 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); + + est_dist_arr = f_add_arr + q_obj.g_add() + + (f_rescale_arr * (q_obj.delta() * (ip_arr.template cast()) + + q_obj.sum_vl_lut() + q_obj.k1xsumq())); + return; + } + + std::array accu_values{}; + std::array accu_res{}; + const auto* codes_ptr = cur_batch.bin_code(); + const auto* lut_ptr = q_obj.lut(); + size_t remaining_dim = padded_dim; + + while (remaining_dim > kSafeChunkDim) { + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), kSafeChunkDim); + codes_ptr += kSafeChunkDim << 2; + lut_ptr += kSafeChunkDim << 2; + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_values[i] += accu_res[i]; + } + remaining_dim -= kSafeChunkDim; + } + + fastscan::accumulate(codes_ptr, lut_ptr, accu_res.data(), remaining_dim); + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + accu_values[i] += accu_res[i]; + } + + ConstRowMajorArrayMap ip_arr(accu_values.data(), 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); + + est_dist_arr = f_add_arr + q_obj.g_add() + + (f_rescale_arr * (q_obj.delta() * (ip_arr.template cast()) + + q_obj.sum_vl_lut() + q_obj.k1xsumq())); +} + +} // namespace +} // namespace rabitqlib::simd diff --git a/src/simd/fastscan_generic.cpp b/src/simd/fastscan_generic.cpp new file mode 100644 index 0000000..52c6644 --- /dev/null +++ b/src/simd/fastscan_generic.cpp @@ -0,0 +1,18 @@ +#include + +#include "rabitqlib/defines.hpp" +#include "rabitqlib/fastscan/fastscan.hpp" +#include "rabitqlib/simd/fastscan_dispatch.hpp" + +namespace rabitqlib::fastscan::simd { +void pack_lut_generic(size_t dim, const float* query, float* lut) { + for (size_t group = 0; group < dim / 4; ++group) { + lut[0] = 0; + for (size_t j = 1; j < 16; ++j) { + lut[j] = lut[j - LOWBIT(j)] + query[kPos[j]]; + } + query += 4; + lut += 16; + } +} +} // namespace rabitqlib::fastscan::simd diff --git a/src/simd/matrix_avx2.cpp b/src/simd/matrix_avx2.cpp new file mode 100644 index 0000000..00410c0 --- /dev/null +++ b/src/simd/matrix_avx2.cpp @@ -0,0 +1,45 @@ +#include + +#include "rabitqlib/simd/matrix_dispatch.hpp" + +#define RABITQ_MATRIX_EIGEN_NAMESPACE rabitqlib_eigen_avx2 +#include "matrix_kernels.hpp" + +namespace rabitqlib::simd { +void matrix_product_avx2( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_impl(left, right, result, rows, inner, cols); +} + +void matrix_product_transposed_avx2( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_transposed_impl(left, right, result, rows, inner, cols); +} + +void row_norms_avx2(const float* data, float* result, size_t rows, size_t dim) { + row_norms_impl(data, result, rows, dim); +} + +void pairwise_distances_lower_avx2( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +) { + pairwise_distances_lower_impl(data, result, norms_data, size, dim, inner_product); +} +} // namespace rabitqlib::simd diff --git a/src/simd/matrix_avx512.cpp b/src/simd/matrix_avx512.cpp new file mode 100644 index 0000000..f35c028 --- /dev/null +++ b/src/simd/matrix_avx512.cpp @@ -0,0 +1,45 @@ +#include + +#include "rabitqlib/simd/matrix_dispatch.hpp" + +#define RABITQ_MATRIX_EIGEN_NAMESPACE rabitqlib_eigen_avx512 +#include "matrix_kernels.hpp" + +namespace rabitqlib::simd { +void matrix_product_avx512( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_impl(left, right, result, rows, inner, cols); +} + +void matrix_product_transposed_avx512( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_transposed_impl(left, right, result, rows, inner, cols); +} + +void row_norms_avx512(const float* data, float* result, size_t rows, size_t dim) { + row_norms_impl(data, result, rows, dim); +} + +void pairwise_distances_lower_avx512( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +) { + pairwise_distances_lower_impl(data, result, norms_data, size, dim, inner_product); +} +} // namespace rabitqlib::simd diff --git a/src/simd/matrix_generic.cpp b/src/simd/matrix_generic.cpp new file mode 100644 index 0000000..368707e --- /dev/null +++ b/src/simd/matrix_generic.cpp @@ -0,0 +1,45 @@ +#include + +#include "rabitqlib/simd/matrix_dispatch.hpp" + +#define RABITQ_MATRIX_EIGEN_NAMESPACE rabitqlib_eigen_generic +#include "matrix_kernels.hpp" + +namespace rabitqlib::simd { +void matrix_product_generic( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_impl(left, right, result, rows, inner, cols); +} + +void matrix_product_transposed_generic( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + matrix_product_transposed_impl(left, right, result, rows, inner, cols); +} + +void row_norms_generic(const float* data, float* result, size_t rows, size_t dim) { + row_norms_impl(data, result, rows, dim); +} + +void pairwise_distances_lower_generic( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +) { + pairwise_distances_lower_impl(data, result, norms_data, size, dim, inner_product); +} +} // namespace rabitqlib::simd diff --git a/src/simd/matrix_kernels.hpp b/src/simd/matrix_kernels.hpp new file mode 100644 index 0000000..d8cac1e --- /dev/null +++ b/src/simd/matrix_kernels.hpp @@ -0,0 +1,85 @@ +#pragma once + +#include +#include + +// Keep Eigen's ISA-dependent template helpers private to each backend. The +// default also lets this header compile on its own for static analysis. +#ifndef RABITQ_MATRIX_EIGEN_NAMESPACE +#define RABITQ_MATRIX_EIGEN_NAMESPACE rabitqlib_eigen_generic +#endif +#define Eigen RABITQ_MATRIX_EIGEN_NAMESPACE +#include "rabitqlib/third/Eigen/Dense" +#undef Eigen + +namespace rabitqlib::simd { +namespace { +namespace kernel_eigen = RABITQ_MATRIX_EIGEN_NAMESPACE; +#undef RABITQ_MATRIX_EIGEN_NAMESPACE + +using Index = kernel_eigen::Index; +using Matrix = kernel_eigen:: + Matrix; +using ConstRowMajorMatrixMap = kernel_eigen::Map; +using RowMajorMatrixMap = kernel_eigen::Map; +using VectorMap = kernel_eigen::Map>; + +inline void matrix_product_impl( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + ConstRowMajorMatrixMap a(left, static_cast(rows), static_cast(inner)); + ConstRowMajorMatrixMap b(right, static_cast(inner), static_cast(cols)); + RowMajorMatrixMap out(result, static_cast(rows), static_cast(cols)); + out.noalias() = a * b; +} + +inline void matrix_product_transposed_impl( + const float* left, + const float* right, + float* result, + size_t rows, + size_t inner, + size_t cols +) { + ConstRowMajorMatrixMap a(left, static_cast(rows), static_cast(inner)); + ConstRowMajorMatrixMap b(right, static_cast(cols), static_cast(inner)); + RowMajorMatrixMap out(result, static_cast(rows), static_cast(cols)); + out.noalias() = a * b.transpose(); +} + +inline void row_norms_impl(const float* data, float* result, size_t rows, size_t dim) { + ConstRowMajorMatrixMap points(data, static_cast(rows), static_cast(dim)); + VectorMap norms(result, static_cast(rows)); + norms = points.rowwise().squaredNorm(); +} + +inline void pairwise_distances_lower_impl( + const float* data, + float* result, + float* norms_data, + size_t size, + size_t dim, + bool inner_product +) { + const auto count = static_cast(size); + ConstRowMajorMatrixMap points(data, count, static_cast(dim)); + RowMajorMatrixMap distances(result, count, count); + VectorMap norms(norms_data, count); + distances.setZero(); + distances.selfadjointView().rankUpdate(points); + norms = points.rowwise().squaredNorm(); + for (kernel_eigen::Index i = 0; i < count; ++i) { + for (kernel_eigen::Index j = 0; j <= i; ++j) { + const float dot = distances(i, j); + distances(i, j) = + !inner_product ? std::max(0.0F, norms[i] + norms[j] - 2 * dot) : -dot; + } + } +} +} // namespace +} // namespace rabitqlib::simd diff --git a/src/simd/quantization_generic.cpp b/src/simd/quantization_generic.cpp new file mode 100644 index 0000000..bcaeafe --- /dev/null +++ b/src/simd/quantization_generic.cpp @@ -0,0 +1,24 @@ +#include + +#include "rabitqlib/simd/quantization_dispatch.hpp" +#include "rescale_search.hpp" + +namespace rabitqlib::simd { +// Request the existing scalar event sweep; it certifies ambiguous SIMD searches. +double best_rescale_factor_generic(const float*, size_t, int, double, double) { return -1; } +namespace detail { + +RescaleScratch& get_thread_local_rescale_scratch(size_t dim) { + // Search evaluation is synchronous and does not re-enter the quantizer. + // Retain the largest buffers on each worker; smaller vectors overwrite only + // their active prefix without shrinking or zero-initializing the storage. + thread_local RescaleScratch scratch; + if (scratch.magnitudes.size() < dim) + scratch.magnitudes.resize(dim); + if (scratch.reciprocals.size() < dim) + scratch.reciprocals.resize(dim); + return scratch; +} + +} // namespace detail +} // namespace rabitqlib::simd diff --git a/src/simd/rotator_avx2.cpp b/src/simd/rotator_avx2.cpp index a8902d4..5b9be3d 100644 --- a/src/simd/rotator_avx2.cpp +++ b/src/simd/rotator_avx2.cpp @@ -5,6 +5,7 @@ #include #include "rabitqlib/simd/rotator_dispatch.hpp" +#include "rotator_kernels.hpp" namespace rabitqlib::simd { @@ -52,4 +53,18 @@ void kacs_walk_avx2(float* data, size_t len) { } } +void fht_rotate_avx2( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +) { + fht_rotate_impl( + data, rotated_vec, dim, padded_dim, trunc_dim, fac, flip + ); +} + } // namespace rabitqlib::simd diff --git a/src/simd/rotator_avx512.cpp b/src/simd/rotator_avx512.cpp index b2feac5..2b34d5a 100644 --- a/src/simd/rotator_avx512.cpp +++ b/src/simd/rotator_avx512.cpp @@ -5,6 +5,7 @@ #include #include "rabitqlib/simd/rotator_dispatch.hpp" +#include "rotator_kernels.hpp" namespace rabitqlib::simd { @@ -67,4 +68,18 @@ void kacs_walk_avx512(float* data, size_t len) { } } +void fht_rotate_avx512( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +) { + fht_rotate_impl( + data, rotated_vec, dim, padded_dim, trunc_dim, fac, flip + ); +} + } // namespace rabitqlib::simd diff --git a/src/simd/rotator_kernels.hpp b/src/simd/rotator_kernels.hpp new file mode 100644 index 0000000..fa5bd98 --- /dev/null +++ b/src/simd/rotator_kernels.hpp @@ -0,0 +1,104 @@ +#pragma once +#include +#include +#include +#include +#include + +namespace rabitqlib::simd { +namespace { +// The imported header has no includes. Keep its inline helpers private so the +// linker cannot merge compiler-generated code from different ISA backends. +#include "rabitqlib/utils/fht_avx.hpp" +inline void rescale(float* data, size_t dim, float factor) { + for (size_t i = 0; i < dim; ++i) { + data[i] *= factor; + } +} +// FFHT retains its imported AVX butterflies in both ISA backends. Surrounding +// C++ is compiled at the selected ISA; no vendor implementation is rewritten. +template +void fht_rotate_impl( + const float* data, + float* rotated_vec, + size_t dim, + size_t padded_dim, + size_t trunc_dim, + float fac, + const uint8_t* flip +) { + void (*fht)(float*) = nullptr; + switch (trunc_dim) { + case 64: + fht = helper_float_6; + break; + case 128: + fht = helper_float_7; + break; + case 256: + fht = helper_float_8; + break; + case 512: + fht = helper_float_9; + break; + case 1024: + fht = helper_float_10; + break; + case 2048: + fht = helper_float_11; + break; + default: + throw std::invalid_argument("Unsupported dimension for FhtKacRotator"); + } + + std::memcpy(rotated_vec, data, sizeof(float) * dim); + std::fill(rotated_vec + dim, rotated_vec + padded_dim, 0); + + if (trunc_dim == padded_dim) { + flip_sign(flip, rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + + flip_sign(flip + (padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + + flip_sign(flip + (2 * padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + + flip_sign(flip + (3 * padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + + return; + } + + size_t start = padded_dim - trunc_dim; + + flip_sign(flip, rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + kacs_walk(rotated_vec, padded_dim); + + flip_sign(flip + (padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec + start); + rescale(rotated_vec + start, trunc_dim, fac); + kacs_walk(rotated_vec, padded_dim); + + flip_sign(flip + (2 * padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec); + rescale(rotated_vec, trunc_dim, fac); + kacs_walk(rotated_vec, padded_dim); + + flip_sign(flip + (3 * padded_dim / 8), rotated_vec, padded_dim); + fht(rotated_vec + start); + rescale(rotated_vec + start, trunc_dim, fac); + kacs_walk(rotated_vec, padded_dim); + + // This can be removed if we don't care about the absolute value of + // similarities. + rescale(rotated_vec, padded_dim, 0.25F); +} +} // namespace +} // namespace rabitqlib::simd diff --git a/src/simd/space_float.cpp b/src/simd/space_generic.cpp similarity index 100% rename from src/simd/space_float.cpp rename to src/simd/space_generic.cpp diff --git a/tests/python/test_ivf.py b/tests/python/test_ivf.py index e541935..5addb0d 100644 --- a/tests/python/test_ivf.py +++ b/tests/python/test_ivf.py @@ -388,3 +388,26 @@ def test_automatic_high_accuracy(nbits, tmp_path): ): np.testing.assert_array_equal(actual[0], expected[0]) np.testing.assert_array_equal(actual[1], expected[1]) + + +def test_hnsw_centroid_routing_roundtrip(tmp_path): + # 20,000 clusters selects HNSW centroid routing rather than the flat router. + # One vector per centroid makes nearest-centroid IDs and distances exact. + rng = np.random.default_rng(314) + count, dim = 20000, 64 + centroids = rng.standard_normal((count, dim)).astype(np.float32) + labels = np.arange(count, dtype=np.uint32) + idx = IvfIndex(dim, count, count, nbits=1) + idx.build(centroids, centroids, labels, num_threads=1) + selected = np.array([0, 17, 1023, count - 1]) + queries = centroids[selected] + ids, distances = idx.search(queries, k=1, nprobe=4, num_threads=1) + np.testing.assert_array_equal(ids[:, 0], selected) + np.testing.assert_allclose(distances, 0, rtol=0, atol=1e-5) + + path = str(tmp_path / "large_centroid.index") + idx.save(path) + loaded = IvfIndex.load(path) + loaded_ids, loaded_distances = loaded.search(queries, k=1, nprobe=4, num_threads=1) + np.testing.assert_array_equal(loaded_ids, ids) + np.testing.assert_array_equal(loaded_distances, distances) diff --git a/tests/unit/rabitqlib/fastscan/fastscan_test.cpp b/tests/unit/rabitqlib/fastscan/fastscan_test.cpp index 3ba14d6..9081495 100644 --- a/tests/unit/rabitqlib/fastscan/fastscan_test.cpp +++ b/tests/unit/rabitqlib/fastscan/fastscan_test.cpp @@ -17,6 +17,7 @@ #include "rabitqlib/index/lut.hpp" #include "rabitqlib/index/query.hpp" #include "rabitqlib/quantization/data_layout.hpp" +#include "rabitqlib/simd/estimator_dispatch.hpp" #include "rabitqlib/simd/fastscan_dispatch.hpp" #include "rabitqlib/utils/cpu_features.hpp" @@ -234,6 +235,101 @@ TEST(FastScanHighAccuracyTest, Avx512TransferAcceptsUnalignedOutput) { EXPECT_TRUE(std::equal(expected.begin(), expected.end(), storage.begin() + 1)); } +TEST(BatchEstimatorTest, BackendsMatchScalarCorrectionAcrossChunksAndTailBatches) { + if (!cpu::has_avx2() && !cpu::has_avx512_core()) { + GTEST_SKIP() << "FastScan requires AVX2/FMA or AVX512"; + } + using Estimator = decltype(&rabitqlib::simd::split_batch_estdist); + std::vector backends{ + rabitqlib::simd::split_batch_estdist, rabitqlib::simd::split_batch_estdist_generic}; + if (cpu::has_avx2()) { + backends.push_back(rabitqlib::simd::split_batch_estdist_avx2); + } + if (cpu::has_avx512_core()) { + backends.push_back(rabitqlib::simd::split_batch_estdist_avx512); + } + for (size_t dim : {16U, 64U, 128U, 1024U, 1040U, 4096U}) { + for (bool hacc : {false, true}) { + for (auto metric : {METRIC_L2, METRIC_IP}) { + SCOPED_TRACE(::testing::Message() << dim << " " << hacc << " " << metric); + std::vector query(dim, 1.0F); + SplitBatchQuery q(query.data(), dim, 3, metric, hacc); + q.set_g_add(3.0F, 2.0F); + // An unaligned batch with only 17 logical rows still has 32 physical lanes. + std::vector storage(BatchDataMap::data_bytes(dim) + 1); + BatchDataMap batch(storage.data() + 1, dim); + std::vector codes(17 * dim / 8); + for (size_t i = 0; i < codes.size(); ++i) { + codes[i] = static_cast(i * 73 + i / 7); + } + pack_codes(dim, codes.data(), 17, batch.bin_code()); + for (size_t lane = 0; lane < kBatchSize; ++lane) { + batch.f_add()[lane] = static_cast(lane) + 10.0F; + batch.f_rescale()[lane] = lane % 2 == 0 ? 0.125F : -0.25F; + batch.f_error()[lane] = static_cast(lane) / 32.0F; + } + // With a constant-one query, LUT entries depend only on sign-bit count. + std::array lut{}; + const float inverse_delta = 1.0F / q.delta(); + for (size_t count = 0; count < lut.size(); ++count) { + lut[count] = static_cast( + std::nearbyint(static_cast(count) * inverse_delta) + ); + } + std::array expected_ip{}, expected_dist{}, + expected_low{}; + for (size_t lane = 0; lane < kBatchSize; ++lane) { + int32_t sum = 0; + for (size_t group = 0; group < dim / 4; ++group) { + const uint8_t byte = + lane < 17 ? codes[lane * dim / 8 + group / 2] : 0; + const unsigned code = (byte >> (group % 2 == 0 ? 4 : 0)) & 15; + unsigned count = 0; + for (unsigned bit = 0; bit < 4; ++bit) { + count += (code >> bit) & 1U; + } + sum += lut[count]; + } + expected_ip[lane] = + static_cast(q.delta()) * sum + q.sum_vl_lut(); + expected_dist[lane] = static_cast(batch.f_add()[lane]) + + q.g_add() + + static_cast(batch.f_rescale()[lane]) * + (expected_ip[lane] + q.k1xsumq()); + expected_low[lane] = + expected_dist[lane] - + static_cast(batch.f_error()[lane]) * q.g_error(); + } + for (auto backend : backends) { + std::array dist, low, ip; + dist.fill(12345.0F); + low.fill(12345.0F); + ip.fill(12345.0F); + backend( + storage.data() + 1, + q, + dim, + dist.data() + 1, + low.data() + 1, + ip.data() + 1, + hacc + ); + for (size_t lane = 0; lane < kBatchSize; ++lane) { + const double tolerance = 2e-6 * static_cast(dim); + EXPECT_NEAR(ip[lane + 1], expected_ip[lane], tolerance); + EXPECT_NEAR(dist[lane + 1], expected_dist[lane], tolerance); + EXPECT_NEAR(low[lane + 1], expected_low[lane], tolerance); + } + for (const auto* output : {&dist, &low, &ip}) { + EXPECT_EQ(output->front(), 12345.0F); + EXPECT_EQ(output->back(), 12345.0F); + } + } + } + } + } +} + TEST(FastScanHighAccuracyTest, AccumulatesAcrossChunks) { if (!cpu::has_avx2()) { GTEST_SKIP() << "AVX2 is not supported on this CPU"; diff --git a/tests/unit/rabitqlib/index/initializer_test.cpp b/tests/unit/rabitqlib/index/initializer_test.cpp index 8ce0b34..85d15fc 100644 --- a/tests/unit/rabitqlib/index/initializer_test.cpp +++ b/tests/unit/rabitqlib/index/initializer_test.cpp @@ -2,13 +2,56 @@ #include +#include #include +#include #include #include +#include namespace rabitqlib::ivf { namespace { +TEST(CentroidL2SpaceTest, MatchesDispatchedDistanceForUnalignedInputsAndTails) { + for (size_t dim : {1U, 7U, 8U, 15U, 16U, 17U, 63U, 64U, 129U}) { + CentroidL2Space space(dim); + std::vector a(dim + 1), b(dim + 1); + double expected = 0; + for (size_t i = 1; i <= dim; ++i) { + a[i] = static_cast(i) / 8.0F; + b[i] = static_cast(i % 7) / 4.0F; + const double delta = static_cast(a[i]) - b[i]; + expected += delta * delta; + } + EXPECT_EQ(space.get_data_size(), dim * sizeof(float)); + const float actual = + space.get_dist_func()(a.data() + 1, b.data() + 1, space.get_dist_func_param()); + EXPECT_FLOAT_EQ(actual, euclidean_sqr(a.data() + 1, b.data() + 1, dim)); + EXPECT_NEAR(actual, expected, 2e-6 * expected); + } +} + +TEST(HNSWInitializerTest, RoutesToExactCentroidsWithEuclideanDistances) { + constexpr size_t kDim = 17; + constexpr size_t kCount = 32; + std::vector centroids(kCount * kDim); + for (size_t i = 0; i < kCount; ++i) { + centroids[i * kDim] = static_cast(i) * 2.0F; + } + HNSWInitializer initializer(kDim, kCount); + initializer.add_vectors(centroids.data(), 1); + std::array query{}; + query[0] = 20.25F; + std::vector> candidates(3); + initializer.centroids_distances(query.data(), candidates.size(), candidates); + for (const auto& candidate : candidates) { + EXPECT_TRUE(candidate.id == 9 || candidate.id == 10 || candidate.id == 11); + EXPECT_FLOAT_EQ( + candidate.distance, std::abs(query[0] - static_cast(candidate.id) * 2.0F) + ); + } +} + TEST(ParallelForTest, AutomaticThreadCountProcessesEveryItem) { std::atomic calls{0}; parallel_for(0, 100, 0, [&](size_t, size_t) { ++calls; }); diff --git a/tests/unit/rabitqlib/index/ivf_test.cpp b/tests/unit/rabitqlib/index/ivf_test.cpp index d98dc63..36c49cd 100644 --- a/tests/unit/rabitqlib/index/ivf_test.cpp +++ b/tests/unit/rabitqlib/index/ivf_test.cpp @@ -22,6 +22,20 @@ namespace rabitqlib::ivf { namespace { +TEST(IvfSearchTest, BatchCandidatesPreserveTiesAndTailCount) { + buffer::SearchBuffer knns(3); + const std::array ids{0, 1, 2, 3, 4, 5}; + const std::array distances{3, 1, 2, 1, 4, -100}; + detail::insert_candidates(knns, ids.data(), distances.data(), 3); + detail::insert_candidates(knns, ids.data() + 3, distances.data() + 3, 2); + detail::insert_candidates(knns, ids.data(), distances.data(), 0); + std::array results{}; + std::array result_distances{}; + knns.copy_results(results.data(), result_distances.data()); + EXPECT_EQ(results, (std::array{3, 1, 2})); + EXPECT_EQ(result_distances, (std::array{1, 1, 2})); +} + TEST(IvfConfigurationTest, RejectsUnsupportedMetric) { EXPECT_THROW( (IVF(8, 64, 1, 1, static_cast(255), RotatorType::MatrixRotator)), diff --git a/tests/unit/rabitqlib/index/qg_test.cpp b/tests/unit/rabitqlib/index/qg_test.cpp index c47abe3..be038b3 100644 --- a/tests/unit/rabitqlib/index/qg_test.cpp +++ b/tests/unit/rabitqlib/index/qg_test.cpp @@ -35,7 +35,7 @@ namespace rabitqlib::symqg { struct QGConstructionTestAccess { - static void check_contiguous_rows(QuantizedGraph& graph) { + static void check_contiguous_rows(QuantizedGraph& graph) { const size_t vector_bytes = graph.is_quantized() ? ExDataMap::data_bytes(graph.padded_dim_, graph.quantization_bits_) @@ -70,7 +70,7 @@ struct QGConstructionTestAccess { } static QGBuilder from_graph( - QuantizedGraph& graph, + QuantizedGraph& graph, uint32_t ef, const float* data, const std::vector& offsets, @@ -81,7 +81,7 @@ struct QGConstructionTestAccess { builder.initialize_seed(data, offsets, neighbors); return builder; } - static auto codes(const QuantizedGraph& graph) { + static auto codes(const QuantizedGraph& graph) { std::vector result; for (PID id = 0; id < graph.num_points_; ++id) { const char* code = graph.is_quantized() @@ -99,7 +99,7 @@ struct QGConstructionTestAccess { auto& graph = builder.qg_; for (PID i = 0; i < builder.num_nodes_; ++i) { std::vector reconstructed; - std::optional> prepared; + std::optional prepared; const float* source = graph.prepare_build_query(i, reconstructed, prepared); for (size_t j = 0; j < builder.degrees_[i]; ++j) { const PID id = graph.get_neighbors(i)[j]; @@ -110,39 +110,35 @@ struct QGConstructionTestAccess { } } static void search(QGBuilder& builder) { builder.search_new_neighbors(false); } - static auto 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 + 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 - ) { + 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 + QuantizedGraph& graph, PID source, const std::vector>& neighbors ) { graph.update_qg(source, neighbors); } - static void fill_batch_factors(QuantizedGraph& graph, PID source, float value) { + 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 + 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 + const QuantizedGraph& graph, PID source, PID target ) { std::vector query(graph.padded_dim_); graph.reconstruct_quantized_vector(source, query.data()); @@ -180,9 +176,7 @@ namespace { TEST(QuantizedGraphLayoutTest, KeepsVectorsCodesAndNeighborsInContiguousRows) { for (size_t bits : {0U, 4U, 8U}) { SCOPED_TRACE(bits); - QuantizedGraph graph( - 33, 65, 32, METRIC_L2, RotatorType::FhtKacRotator, bits - ); + QuantizedGraph graph(33, 65, 32, METRIC_L2, RotatorType::FhtKacRotator, bits); QGConstructionTestAccess::check_contiguous_rows(graph); } } @@ -193,7 +187,7 @@ static_assert(std::is_move_constructible_v>); static_assert( !std::is_constructible_v< QGBuilder, - QuantizedGraph&, + QuantizedGraph&, uint32_t, const float*, const std::vector&, @@ -231,6 +225,22 @@ TEST(QGEstimatorTest, MatchesExactDistancesForCollinearResiduals) { BatchQuery q_obj(query.data(), dim, metric); q_obj.set_g_add(vertex_distance); std::array estimates{}; + std::vector backends{ + simd::qg_batch_estdist_generic, simd::qg_batch_estdist}; + if (cpu::has_avx2()) + backends.push_back(simd::qg_batch_estdist_avx2); + if (cpu::has_avx512_core()) + backends.push_back(simd::qg_batch_estdist_avx512); + for (auto backend : backends) { + backend(batch.data(), q_obj, dim, estimates.data()); + for (size_t i = 0; i < fastscan::kBatchSize; ++i) { + EXPECT_NEAR( + estimates[i], + distance(query.data(), data.data() + i * dim, dim), + 1e-5F + ); + } + } qg_batch_estdist(batch.data(), q_obj, dim, estimates.data()); EXPECT_FLOAT_EQ( q_obj.g_add(), metric == METRIC_IP ? vertex_distance - 1 : vertex_distance @@ -244,6 +254,22 @@ TEST(QGEstimatorTest, MatchesExactDistancesForCollinearResiduals) { } } +TEST(QGEstimatorTest, CandidateMaskBackendsMatchScalarOrdering) { + std::array distances{}; + for (size_t lane = 0; lane < distances.size(); ++lane) { + distances[lane] = static_cast(lane % 7) - 2.0F; + } + distances[3] = std::numeric_limits::quiet_NaN(); + distances[17] = std::numeric_limits::infinity(); + distances[29] = -std::numeric_limits::infinity(); + + constexpr float kThreshold = 2.0F; + const uint32_t expected = simd::qg_candidate_mask_generic(distances.data(), kThreshold); + if (cpu::has_avx2()) { + EXPECT_EQ(simd::qg_candidate_mask_avx2(distances.data(), kThreshold), expected); + } +} + TEST(QGConstructionTest, BuildsAndPrunesAfterInputReleaseUsingExistingCodes) { constexpr size_t kCount = 65, kDim = 65; for (auto metric : {METRIC_L2, METRIC_IP}) { @@ -259,7 +285,7 @@ TEST(QGConstructionTest, BuildsAndPrunesAfterInputReleaseUsingExistingCodes) { static_cast(1 + (i / kDim) % 4); } } - QuantizedGraph graph( + QuantizedGraph graph( kCount, kDim, 32, metric, RotatorType::FhtKacRotator, bits ); QGBuilder builder(graph, 64, data.data(), 1); @@ -308,7 +334,7 @@ TEST(QGConstructionTest, EncodedSeedSearchMatchesRetainedScores) { for (auto metric : {METRIC_L2, METRIC_IP}) { for (size_t bits : {0U, 4U, 8U}) { SCOPED_TRACE(::testing::Message() << metric << "/" << bits); - QuantizedGraph graph( + QuantizedGraph graph( kCount, kDim, kDegree, metric, RotatorType::FhtKacRotator, bits ); // Both builders search the same immutable encoded graph/rotation. @@ -348,7 +374,7 @@ TEST(QGConstructionTest, DefaultsToPipnnInitialization) { detail::build_initial_graph(data.data(), kCount, kDim, kDegree, metric, 1); for (size_t bits : {0U, 4U, 8U}) { SCOPED_TRACE(::testing::Message() << metric << "/" << bits); - QuantizedGraph graph( + QuantizedGraph graph( kCount, kDim, kDegree, metric, RotatorType::FhtKacRotator, bits ); QGBuilder builder(graph, 64, data.data(), 1); @@ -377,7 +403,7 @@ TEST(QGConstructionTest, SupportsExplicitRandomInitialization) { GTEST_SKIP() << "FastScan requires AVX2/FMA"; } std::vector data(65 * 65, 0.5F); - QuantizedGraph graph(65, 65, 32); + QuantizedGraph graph(65, 65, 32); QGBuilder builder(graph, 64, data.data(), 1, QGInitialization::Random); for (const auto& row : QGConstructionTestAccess::neighbors(builder)) { EXPECT_EQ(row.size(), 32U); @@ -394,7 +420,7 @@ TEST(QGConstructionTest, SupportsExplicitRandomInitialization) { } TEST(QGConstructionTest, RejectsNullDataForEveryInitializationMode) { - QuantizedGraph graph(33, 64, 32); + QuantizedGraph graph(33, 64, 32); EXPECT_THROW( (QGBuilder(graph, 32, nullptr, 1, QGInitialization::PiPNN)), std::invalid_argument ); @@ -411,8 +437,8 @@ TEST(QGConstructionTest, UsesExplicitThreadCountsWithoutChangingCallerState) { 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); + 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); @@ -432,7 +458,7 @@ TEST(QGConstructionTest, InitializesUnusedPartialBatchFactors) { for (size_t i = 0; i < data.size(); ++i) { data[i] = std::sin(static_cast(i) * 0.17F); } - QuantizedGraph graph(kCount, kDim, 32); + QuantizedGraph graph(kCount, kDim, 32); QGConstructionTestAccess::copy_vectors(graph, data.data(), 1); QGConstructionTestAccess::fill_batch_factors( graph, 0, std::numeric_limits::quiet_NaN() @@ -474,7 +500,7 @@ TEST(QGConstructionTest, RefinesPartialSeedOnceAfterReleasingInputs) { } offsets.push_back(edges.size()); } - QuantizedGraph graph( + QuantizedGraph graph( kCount, kDim, kDegree, metric, RotatorType::FhtKacRotator, bits ); auto builder = QGConstructionTestAccess::from_graph( @@ -534,7 +560,7 @@ TEST(QGConstructionTest, RefinesPartialSeedOnceAfterReleasingInputs) { } const std::string path = ::testing::TempDir() + "rabitq_qg_seeded.index"; graph.save(path.c_str()); - QuantizedGraph loaded; + QuantizedGraph loaded; loaded.load(path.c_str()); loaded.set_ef(96); loaded.search(query.data(), 10, loaded_ids.data(), loaded_distances.data()); @@ -549,7 +575,7 @@ TEST(QGConstructionTest, RefinesPartialSeedOnceAfterReleasingInputs) { TEST(QGConstructionTest, RejectsInvalidSeedStructureAndIds) { constexpr size_t kCount = 33, kDim = 64; std::vector data(kCount * kDim, 0.1F); - QuantizedGraph graph(kCount, kDim, 32); + QuantizedGraph graph(kCount, kDim, 32); const auto reject = [&](const std::vector& offsets, const std::vector& edges, const char* message) { @@ -582,37 +608,31 @@ TEST(QGConstructionTest, RejectsInvalidSeedStructureAndIds) { TEST(QuantizedGraphConfigurationTest, RejectsDegreeNotAlignedForFastScan) { EXPECT_THROW( - (QuantizedGraph(64, 64, 16, METRIC_L2, RotatorType::MatrixRotator)), + (QuantizedGraph(64, 64, 16, METRIC_L2, RotatorType::MatrixRotator)), std::invalid_argument ); } TEST(QuantizedGraphConfigurationTest, RejectsDegreeThatCannotExcludeSelf) { EXPECT_THROW( - (QuantizedGraph(32, 64, 32, METRIC_L2, RotatorType::MatrixRotator)), + (QuantizedGraph(32, 64, 32, METRIC_L2, RotatorType::MatrixRotator)), std::invalid_argument ); } 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_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)), + (QuantizedGraph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator, 6)), std::invalid_argument ); } TEST(QuantizedGraphConfigurationTest, RejectsUnsupportedMetric) { EXPECT_THROW( - (QuantizedGraph( + (QuantizedGraph( 33, 64, 32, static_cast(255), RotatorType::MatrixRotator )), std::invalid_argument @@ -621,13 +641,13 @@ TEST(QuantizedGraphConfigurationTest, RejectsUnsupportedMetric) { TEST(QuantizedGraphConfigurationTest, RejectsZeroDimension) { EXPECT_THROW( - (QuantizedGraph(33, 0, 32, METRIC_L2, RotatorType::MatrixRotator)), + (QuantizedGraph(33, 0, 32, METRIC_L2, RotatorType::MatrixRotator)), std::invalid_argument ); } TEST(QuantizedGraphConfigurationTest, RejectsOutOfRangeEntryPoint) { - QuantizedGraph graph(33, 64, 32); + 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); @@ -638,14 +658,14 @@ TEST(QuantizedGraphPersistenceTest, RejectsMalformedPayloadWithoutChangingTarget constexpr size_t kDim = 64; constexpr size_t kDegree = 32; std::vector data(kNumPoints * kDim, 0.25F); - QuantizedGraph source( + 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); + QuantizedGraph target(65, kDim, kDegree, METRIC_IP); source.save(path.c_str()); std::filesystem::resize_file(path, std::filesystem::file_size(path) - 1); @@ -700,7 +720,7 @@ TEST(QuantizedGraphPersistenceTest, RejectsMalformedPayloadWithoutChangingTarget } TEST(QuantizedGraphLifecycleTest, DestroysConcreteRotatorThroughBasePointer) { - QuantizedGraph graph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator); + QuantizedGraph graph(33, 64, 32, METRIC_L2, RotatorType::MatrixRotator); EXPECT_EQ(graph.num_vertices(), 33U); } @@ -711,13 +731,13 @@ TEST(QuantizedGraphLifecycleTest, RejectsSearchAndSaveBeforeBuild) { std::array ids{}; std::array distances{}; - QuantizedGraph empty; + 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); + QuantizedGraph configured(33, 64, 32); configured.set_ef(1); EXPECT_THROW(configured.save(path.c_str()), std::logic_error); EXPECT_THROW( @@ -743,7 +763,7 @@ TEST(QuantizedGraphLifecycleTest, BuilderDoesNotPublishPartialGraph) { 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); + 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); @@ -773,9 +793,7 @@ TEST(QGBuilderMetricTest, UsesInnerProductDistanceToChooseEntryPoint) { ASSERT_EQ(expected, 0U); ASSERT_NE(euclidean_entry, expected); - QuantizedGraph graph( - kNumPoints, kDim, kDegree, METRIC_IP, RotatorType::MatrixRotator - ); + QuantizedGraph graph(kNumPoints, kDim, kDegree, METRIC_IP, RotatorType::MatrixRotator); QGBuilder builder(graph, kDegree, data.data(), 1); EXPECT_EQ(graph.entry_point(), expected); @@ -830,7 +848,7 @@ TEST(QGQuantTest, SearchesAndRoundTripsFourAndEightBitIndexes) { for (size_t bits : {4U, 8U}) { SCOPED_TRACE(bits); - QuantizedGraph graph( + QuantizedGraph graph( kNumPoints, kDim, kDegree, METRIC_L2, RotatorType::MatrixRotator, bits ); { @@ -850,7 +868,7 @@ TEST(QGQuantTest, SearchesAndRoundTripsFourAndEightBitIndexes) { const std::string path = ::testing::TempDir() + "rabitq_qg_quant_" + std::to_string(bits) + ".index"; graph.save(path.c_str()); - QuantizedGraph loaded; + QuantizedGraph loaded; loaded.set_ef(kNumPoints); loaded.load(path.c_str()); EXPECT_TRUE(loaded.is_quantized()); @@ -876,7 +894,7 @@ TEST(QGSearchTest, RejectsInvalidKAndEfInsteadOfReturningPartialResults) { for (size_t i = 0; i < data.size(); ++i) { data[i] = std::sin(static_cast(i) * 0.11F); } - QuantizedGraph graph(kNumPoints, kDim, kDegree); + QuantizedGraph graph(kNumPoints, kDim, kDegree); QGBuilder builder(graph, kDegree, data.data(), 1, QGInitialization::Random); builder.build(2); @@ -945,7 +963,7 @@ TEST(QGSearchTest, ParallelQueriesMatchSerialAcrossIndexesAndSettings) { for (size_t i = 0; i < queries.size(); ++i) { queries[i] = std::cos(static_cast(i) * 0.07F); } - QuantizedGraph graph( + QuantizedGraph graph( num_points, dim, 32, metric, RotatorType::FhtKacRotator, bits ); { diff --git a/tests/unit/rabitqlib/utils/matrix_dispatch_test.cpp b/tests/unit/rabitqlib/utils/matrix_dispatch_test.cpp new file mode 100644 index 0000000..c26f0ec --- /dev/null +++ b/tests/unit/rabitqlib/utils/matrix_dispatch_test.cpp @@ -0,0 +1,109 @@ +#include "rabitqlib/simd/matrix_dispatch.hpp" + +#include + +#include +#include +#include +#include + +#include "rabitqlib/utils/cpu_features.hpp" + +namespace rabitqlib::simd { +namespace { +struct MatrixBackend { + decltype(&matrix_product) product; + decltype(&matrix_product_transposed) transposed; + decltype(&row_norms) norms; + decltype(&pairwise_distances_lower) pairwise; +}; + +TEST(MatrixDispatchTest, AllBackendsMatchScalarForRectangularUnalignedInputs) { + std::vector backends{ + {matrix_product, matrix_product_transposed, row_norms, pairwise_distances_lower}, + {matrix_product_generic, + matrix_product_transposed_generic, + row_norms_generic, + pairwise_distances_lower_generic}}; + if (cpu::has_avx2()) { + backends.push_back( + {matrix_product_avx2, + matrix_product_transposed_avx2, + row_norms_avx2, + pairwise_distances_lower_avx2} + ); + } + if (cpu::has_avx512_core()) { + backends.push_back( + {matrix_product_avx512, + matrix_product_transposed_avx512, + row_norms_avx512, + pairwise_distances_lower_avx512} + ); + } + for (size_t dim : {1U, 7U, 16U, 17U, 65U, 128U}) { + for (size_t rows : {1U, 33U}) { + constexpr size_t kCols = 19; + std::vector a(rows * dim + 1), b(dim * kCols + 1), bt(kCols * dim + 1); + for (size_t i = 0; i < rows * dim; ++i) { + a[i + 1] = std::sin(static_cast(i) * 0.13F); + } + for (size_t i = 0; i < dim; ++i) { + for (size_t j = 0; j < kCols; ++j) { + b[1 + i * kCols + j] = std::cos(static_cast(i + j * 3) * 0.27F); + bt[1 + j * dim + i] = b[1 + i * kCols + j]; + } + } + for (const auto& backend : backends) { + std::vector product(rows * kCols + 2, 12345), transposed(product), + norms(rows + 2, 12345), pairwise(rows * rows + 2, 12345); + backend.product( + a.data() + 1, b.data() + 1, product.data() + 1, rows, dim, kCols + ); + backend.transposed( + a.data() + 1, bt.data() + 1, transposed.data() + 1, rows, dim, kCols + ); + backend.norms(a.data() + 1, norms.data() + 1, rows, dim); + const double tolerance = 2e-6 * static_cast(dim); + for (size_t i = 0; i < rows; ++i) { + double norm = 0; + for (size_t k = 0; k < dim; ++k) { + norm += + static_cast(a[1 + i * dim + k]) * a[1 + i * dim + k]; + } + EXPECT_NEAR(norms[i + 1], norm, tolerance); + for (size_t j = 0; j < kCols; ++j) { + double expected = 0; + for (size_t k = 0; k < dim; ++k) { + expected += static_cast(a[1 + i * dim + k]) * + b[1 + k * kCols + j]; + } + EXPECT_NEAR(product[1 + i * kCols + j], expected, tolerance); + EXPECT_NEAR(transposed[1 + i * kCols + j], expected, tolerance); + } + } + for (bool ip : {false, true}) { + backend.pairwise( + a.data() + 1, pairwise.data() + 1, norms.data() + 1, rows, dim, ip + ); + for (size_t i = 0; i < rows; ++i) { + for (size_t j = 0; j <= i; ++j) { + double expected = 0; + for (size_t k = 0; k < dim; ++k) { + const double x = a[1 + i * dim + k], y = a[1 + j * dim + k]; + expected += ip ? -x * y : (x - y) * (x - y); + } + EXPECT_NEAR(pairwise[1 + i * rows + j], expected, tolerance); + } + } + } + for (const auto* output : {&product, &transposed, &norms, &pairwise}) { + EXPECT_EQ(output->front(), 12345); + EXPECT_EQ(output->back(), 12345); + } + } + } + } +} +} // namespace +} // namespace rabitqlib::simd diff --git a/tests/unit/rabitqlib/utils/rotator_test.cpp b/tests/unit/rabitqlib/utils/rotator_test.cpp index 9662377..f5c47c2 100644 --- a/tests/unit/rabitqlib/utils/rotator_test.cpp +++ b/tests/unit/rabitqlib/utils/rotator_test.cpp @@ -10,6 +10,7 @@ #include #include +#include "rabitqlib/utils/cpu_features.hpp" #include "test_data.hpp" using namespace rabitqlib; @@ -98,3 +99,92 @@ TEST(FlipSignTest, FlipWorks) { } } } + +TEST(FhtDispatchTest, BackendsMatchScalarButterfliesAndPadding) { + std::vector backends; + if (cpu::has_avx2()) { + backends.push_back(simd::fht_rotate_avx2); + } + if (cpu::has_avx512_core()) { + backends.push_back(simd::fht_rotate_avx512); + } + if (backends.empty()) { + GTEST_SKIP() << "FHT rotation requires AVX2/FMA or AVX512"; + } + backends.push_back(simd::fht_rotate); + for (size_t dim : {64U, 65U, 127U, 128U, 129U, 192U, 256U, 512U, 1024U, 2048U, 2049U}) { + SCOPED_TRACE(dim); + const size_t padded = (dim + 63) / 64 * 64; + size_t trunc = 1; + while (trunc * 2 <= dim) { + trunc *= 2; + } + const float fac = 1.0F / std::sqrt(static_cast(trunc)); + std::vector flip(4 * padded / 8); + for (size_t i = 0; i < flip.size(); ++i) { + flip[i] = static_cast(i * 73 + 13); + } + std::vector data(dim + 1), expected(padded, 0); + for (size_t i = 0; i < dim; ++i) { + data[i + 1] = std::sin(static_cast(i) * 0.13F); + expected[i] = data[i + 1]; + } + for (size_t pass = 0; pass < 4; ++pass) { + for (size_t i = 0; i < padded; ++i) { + if ((flip[pass * padded / 8 + i / 8] >> (i % 8)) & 1) { + expected[i] = -expected[i]; + } + } + const size_t start = pass % 2 == 0 ? 0 : padded - trunc; + for (size_t width = 1; width < trunc; width *= 2) { + for (size_t block = 0; block < trunc; block += width * 2) { + for (size_t i = 0; i < width; ++i) { + const size_t x = start + block + i, y = x + width; + const float a = expected[x], b = expected[y]; + expected[x] = a + b; + expected[y] = a - b; + } + } + } + for (size_t i = start; i < start + trunc; ++i) { + expected[i] *= fac; + } + if (padded != trunc) { + for (size_t i = 0; i < padded / 2; ++i) { + const float a = expected[i], b = expected[i + padded / 2]; + expected[i] = a + b; + expected[i + padded / 2] = a - b; + } + } + } + if (padded != trunc) { + for (float& value : expected) { + value *= 0.25F; + } + } + for (auto backend : backends) { + std::vector actual(padded + 2, 12345); + backend( + data.data() + 1, actual.data() + 1, dim, padded, trunc, fac, flip.data() + ); + for (size_t i = 0; i < padded; ++i) { + EXPECT_NEAR(actual[i + 1], expected[i], 2e-5F); + } + EXPECT_EQ(actual.front(), 12345); + EXPECT_EQ(actual.back(), 12345); + } + } +} + +TEST(MatrixRotatorTest, PreservesOverlappingInputAndOutput) { + rotator_impl::MatrixRotator rotator(3, 3); + const float matrix[] = {1, 2, 3, 4, 5, 6, 7, 8, 9}; + rotator.load(reinterpret_cast(matrix)); + for (size_t offset : {0U, 1U}) { + float data[] = {1, 2, 3, 0}; + rotator.rotate(data, data + offset); + EXPECT_FLOAT_EQ(data[offset], 30); + EXPECT_FLOAT_EQ(data[offset + 1], 36); + EXPECT_FLOAT_EQ(data[offset + 2], 42); + } +}