mirror of
https://github.com/facebookresearch/faiss.git
synced 2026-10-11 22:50:00 +00:00
feat: implement RVV optimized scalar quantizer codecs and distance computation (#5535)
Summary:
This PR adds RISC-V Vector Extension SIMD support for Faiss scalar quantizer.
The implementation introduces runtime vector-length aware RVV kernels using the m8 vector configuration, instead of assuming a fixed vector width. This matches the variable-length nature of RVV hardware and allows the implementation to scale across different RISC-V vector implementations.
## Changes
- Add RVV implementations for scalar quantizer codecs:
- 8-bit codec decoding
- 4-bit codec decoding
- (6-bit stays scalar: the RVV gather decoder measured 12-22% slower
than the scalar path, so QT_6bit falls back until a faster decoder
exists)
- Add RVV quantizer reconstruction support for:
- Uniform quantizers
- Non-uniform quantizers
- FP16 (Zvfhmin is the enforced minimum ISA of the RISCV_RVV level; a
build without it fails compilation rather than silently falling back)
- BF16
- Direct 8-bit quantization
- Signed 8-bit quantization
- LloydMax quantizers (1/2/3/4/8-bit)
- Add RVV optimized similarity implementations:
- L2 distance
- Inner product
Pull Request resolved: https://github.com/facebookresearch/faiss/pull/5535
Reviewed By: mnorris11
Differential Revision: D116996553
Pulled By: juancarpio27
fbshipit-source-id: dc726948fa602a04a757d8f9cab8eea41ac1979a
This commit is contained in:
committed by
meta-codesync[bot]
parent
02cacc5293
commit
5e74b1bdde
@@ -469,6 +469,22 @@ jobs:
|
||||
build/tests/faiss_test \
|
||||
--gtest_filter="*SQRVV*:*D5Overflow*:*D5bLargeDim*:*D10FloatQuery*"
|
||||
|
||||
- name: Run scalar quantizer RVV parity tests (via QEMU)
|
||||
run: |
|
||||
qemu-riscv64-static \
|
||||
-cpu rv64,v=true,x-zvfhmin=true \
|
||||
-L /usr/riscv64-linux-gnu \
|
||||
build/tests/faiss_test \
|
||||
--gtest_filter="ScalarQuantizer.*DistancePathParity*"
|
||||
|
||||
- name: Run scalar quantizer RVV parity tests (VLEN=1024, via QEMU)
|
||||
run: |
|
||||
qemu-riscv64-static \
|
||||
-cpu rv64,v=true,vlen=1024,x-zvfhmin=true \
|
||||
-L /usr/riscv64-linux-gnu \
|
||||
build/tests/faiss_test \
|
||||
--gtest_filter="ScalarQuantizer.*DistancePathParity*"
|
||||
|
||||
index-io-backward-compatibility:
|
||||
needs: linux-x86_64-cmake
|
||||
name: Index serialization backward compatibility
|
||||
|
||||
@@ -3,18 +3,23 @@
|
||||
# This source code is licensed under the MIT license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
|
||||
# Cross-compilation toolchain for RISC-V 64-bit (lp64d ABI) on Ubuntu/Debian.
|
||||
# Cross-compilation toolchain for RISC-V 64-bit (lp64d ABI) on Ubuntu 24.04.
|
||||
#
|
||||
# Requires these packages installed on the build host:
|
||||
# gcc-riscv64-linux-gnu g++-riscv64-linux-gnu
|
||||
# Requires:
|
||||
# gcc-14-riscv64-linux-gnu
|
||||
# g++-14-riscv64-linux-gnu
|
||||
#
|
||||
# Target libraries (e.g. libopenblas-dev:riscv64) are installed via apt
|
||||
# multiarch to /usr/lib/riscv64-linux-gnu/. CMake's ONLY find-root mode
|
||||
# searches ${CMAKE_FIND_ROOT_PATH}/usr/lib/riscv64-linux-gnu/ (via the
|
||||
# compiler's multiarch tuple), so the CI script creates a symlink:
|
||||
# Target libraries:
|
||||
# libopenblas-dev:riscv64
|
||||
#
|
||||
# Ubuntu multiarch installs target libraries into:
|
||||
# /usr/lib/riscv64-linux-gnu
|
||||
#
|
||||
# The CI creates:
|
||||
# /usr/riscv64-linux-gnu/usr/lib/riscv64-linux-gnu
|
||||
# -> /usr/lib/riscv64-linux-gnu
|
||||
# before invoking cmake.
|
||||
# -> /usr/lib/riscv64-linux-gnu
|
||||
#
|
||||
# so CMake can locate target libraries while using ONLY root mode.
|
||||
|
||||
set(CMAKE_SYSTEM_NAME Linux)
|
||||
set(CMAKE_SYSTEM_PROCESSOR riscv64)
|
||||
@@ -24,12 +29,20 @@ set(CMAKE_SYSTEM_PROCESSOR riscv64)
|
||||
set(CMAKE_C_COMPILER riscv64-linux-gnu-gcc-14)
|
||||
set(CMAKE_CXX_COMPILER riscv64-linux-gnu-g++-14)
|
||||
|
||||
# Cross-compiler sysroot provided by gcc-riscv64-linux-gnu.
|
||||
# Do NOT set CMAKE_SYSROOT here. Ubuntu's cross toolchain resolves target
|
||||
# libraries through its built-in paths (/usr/riscv64-linux-gnu/lib). Passing
|
||||
# --sysroot=/usr/riscv64-linux-gnu makes ld prepend the sysroot to the
|
||||
# absolute paths inside libc6-dev-riscv64-cross's libc.so linker script,
|
||||
# producing doubled paths such as
|
||||
# /usr/riscv64-linux-gnu/usr/riscv64-linux-gnu/lib/libc.so.6
|
||||
# and failing the CMake compiler check at link time.
|
||||
|
||||
set(CMAKE_FIND_ROOT_PATH /usr/riscv64-linux-gnu)
|
||||
|
||||
# Never look for host-side tools (cmake, python, …) inside the sysroot.
|
||||
# Never search host binaries inside sysroot.
|
||||
set(CMAKE_FIND_ROOT_PATH_MODE_PROGRAM NEVER)
|
||||
# Look for target libraries/headers/packages only inside the sysroot.
|
||||
|
||||
# Target libraries/headers/packages only.
|
||||
set(CMAKE_FIND_ROOT_PATH_MODE_LIBRARY ONLY)
|
||||
set(CMAKE_FIND_ROOT_PATH_MODE_INCLUDE ONLY)
|
||||
set(CMAKE_FIND_ROOT_PATH_MODE_PACKAGE ONLY)
|
||||
|
||||
@@ -635,7 +635,9 @@ if(CMAKE_SYSTEM_PROCESSOR MATCHES "(aarch64|arm64|ARM64)" AND NOT MSVC)
|
||||
target_sources(faiss PRIVATE ${FAISS_SIMD_NEON_SRC})
|
||||
endif()
|
||||
|
||||
# RVV is the baseline SIMD on rv64 builds compiled with rv64gcv. Compile RVV
|
||||
# RVV is the baseline SIMD on rv64 builds; the RVV translation units require
|
||||
# rv64gcv + Zvfhmin (see the RISCV_RVV contract in utils/simd_levels.h), so
|
||||
# they are compiled with -march=rv64gcv_zvfhmin. Compile RVV
|
||||
# sources into the main faiss target, mirroring the ARM NEON story on aarch64.
|
||||
# Guard against DD mode: in dd builds, FAISS_SIMD_SRC already contains
|
||||
# FAISS_SIMD_RVV_SRC and is added via target_sources above, so we must not
|
||||
|
||||
@@ -18,6 +18,10 @@
|
||||
#include <cmath>
|
||||
#include <cstring>
|
||||
|
||||
#if !defined(__riscv_zvfhmin) && !defined(__riscv_zvfh)
|
||||
#error "RISCV_RVV scalar quantizers require Zvfhmin; compile this file with -march including _zvfhmin (e.g. rv64gcv_zvfhmin)"
|
||||
#endif
|
||||
|
||||
namespace faiss {
|
||||
|
||||
namespace scalar_quantizer {
|
||||
@@ -109,14 +113,263 @@ struct Quantizer8bitDirectSigned<SIMDLevel::RISCV_RVV>
|
||||
: Quantizer8bitDirectSigned<SIMDLevel::NONE>(d, trained) {}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct QuantizerLloydMax<1, SIMDLevel::RISCV_RVV>
|
||||
: QuantizerLloydMax<1, SIMDLevel::NONE> {
|
||||
using Base = QuantizerLloydMax<1, SIMDLevel::NONE>;
|
||||
|
||||
QuantizerLloydMax(size_t d, const std::vector<float>& trained)
|
||||
: Base(d, trained) {}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
reconstruct_m8_components(const uint8_t* code, size_t i, size_t vl) const {
|
||||
const size_t byte_base = i >> 3;
|
||||
const size_t byte_vl = ((i + vl - 1) >> 3) - byte_base + 1;
|
||||
vuint8m2_t packed = __riscv_vle8_v_u8m2(code + byte_base, byte_vl);
|
||||
vuint32m8_t packed32 = __riscv_vzext_vf4_u32m8(packed, byte_vl);
|
||||
vuint32m8_t comp =
|
||||
__riscv_vadd_vx_u32m8(__riscv_vid_v_u32m8(vl), i, vl);
|
||||
vuint32m8_t rel = __riscv_vsub_vx_u32m8(
|
||||
__riscv_vsrl_vx_u32m8(comp, 3, vl), byte_base, vl);
|
||||
vuint32m8_t bytes = __riscv_vrgather_vv_u32m8(packed32, rel, vl);
|
||||
vuint32m8_t shift = __riscv_vand_vx_u32m8(comp, 7, vl);
|
||||
vuint32m8_t idx = __riscv_vand_vx_u32m8(
|
||||
__riscv_vsrl_vv_u32m8(bytes, shift, vl), 1, vl);
|
||||
vuint32m8_t off = __riscv_vsll_vx_u32m8(idx, 2, vl);
|
||||
return __riscv_vluxei32_v_f32m8(this->centroids, off, vl);
|
||||
}
|
||||
|
||||
void decode_vector(const uint8_t* code, float* x) const final {
|
||||
size_t i = 0;
|
||||
while (i < this->d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(this->d - i);
|
||||
vfloat32m8_t v = reconstruct_m8_components(code, i, vl);
|
||||
__riscv_vse32_v_f32m8(x + i, v, vl);
|
||||
i += vl;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct QuantizerLloydMax<2, SIMDLevel::RISCV_RVV>
|
||||
: QuantizerLloydMax<2, SIMDLevel::NONE> {
|
||||
using Base = QuantizerLloydMax<2, SIMDLevel::NONE>;
|
||||
|
||||
QuantizerLloydMax(size_t d, const std::vector<float>& trained)
|
||||
: Base(d, trained) {}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
reconstruct_m8_components(const uint8_t* code, size_t i, size_t vl) const {
|
||||
const size_t byte_base = i >> 2;
|
||||
const size_t byte_vl = ((i + vl - 1) >> 2) - byte_base + 1;
|
||||
vuint8m2_t packed = __riscv_vle8_v_u8m2(code + byte_base, byte_vl);
|
||||
vuint32m8_t packed32 = __riscv_vzext_vf4_u32m8(packed, byte_vl);
|
||||
vuint32m8_t comp =
|
||||
__riscv_vadd_vx_u32m8(__riscv_vid_v_u32m8(vl), i, vl);
|
||||
vuint32m8_t rel = __riscv_vsub_vx_u32m8(
|
||||
__riscv_vsrl_vx_u32m8(comp, 2, vl), byte_base, vl);
|
||||
vuint32m8_t bytes = __riscv_vrgather_vv_u32m8(packed32, rel, vl);
|
||||
vuint32m8_t shift = __riscv_vsll_vx_u32m8(
|
||||
__riscv_vand_vx_u32m8(comp, 3, vl), 1, vl);
|
||||
vuint32m8_t idx = __riscv_vand_vx_u32m8(
|
||||
__riscv_vsrl_vv_u32m8(bytes, shift, vl), 3, vl);
|
||||
vuint32m8_t off = __riscv_vsll_vx_u32m8(idx, 2, vl);
|
||||
return __riscv_vluxei32_v_f32m8(this->centroids, off, vl);
|
||||
}
|
||||
|
||||
void decode_vector(const uint8_t* code, float* x) const final {
|
||||
size_t i = 0;
|
||||
while (i < this->d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(this->d - i);
|
||||
vfloat32m8_t v = reconstruct_m8_components(code, i, vl);
|
||||
__riscv_vse32_v_f32m8(x + i, v, vl);
|
||||
i += vl;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct QuantizerLloydMax<3, SIMDLevel::RISCV_RVV>
|
||||
: QuantizerLloydMax<3, SIMDLevel::NONE> {
|
||||
using Base = QuantizerLloydMax<3, SIMDLevel::NONE>;
|
||||
|
||||
QuantizerLloydMax(size_t d, const std::vector<float>& trained)
|
||||
: Base(d, trained) {}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
reconstruct_m8_components(const uint8_t* code, size_t i, size_t vl) const {
|
||||
vuint32m8_t idx0 =
|
||||
__riscv_vadd_vx_u32m8(__riscv_vid_v_u32m8(vl), i, vl);
|
||||
vuint32m8_t bitpos = __riscv_vmul_vx_u32m8(idx0, 3, vl);
|
||||
vuint32m8_t byteoff = __riscv_vsrl_vx_u32m8(bitpos, 3, vl);
|
||||
vuint32m8_t shift = __riscv_vand_vx_u32m8(bitpos, 7, vl);
|
||||
size_t last = i + vl - 1;
|
||||
size_t last_used = (3 * last + 2) >> 3;
|
||||
vuint8m2_t lo = __riscv_vluxei32_v_u8m2(code, byteoff, vl);
|
||||
vuint8m2_t hi = __riscv_vluxei32_v_u8m2(
|
||||
code,
|
||||
__riscv_vminu_vx_u32m8(
|
||||
__riscv_vadd_vx_u32m8(byteoff, 1, vl), last_used, vl),
|
||||
vl);
|
||||
vuint32m8_t w = __riscv_vor_vv_u32m8(
|
||||
__riscv_vzext_vf4_u32m8(lo, vl),
|
||||
__riscv_vsll_vx_u32m8(__riscv_vzext_vf4_u32m8(hi, vl), 8, vl),
|
||||
vl);
|
||||
vuint32m8_t idx = __riscv_vand_vx_u32m8(
|
||||
__riscv_vsrl_vv_u32m8(w, shift, vl), 7, vl);
|
||||
vuint32m8_t off = __riscv_vsll_vx_u32m8(idx, 2, vl);
|
||||
return __riscv_vluxei32_v_f32m8(this->centroids, off, vl);
|
||||
}
|
||||
|
||||
void decode_vector(const uint8_t* code, float* x) const final {
|
||||
size_t i = 0;
|
||||
while (i < this->d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(this->d - i);
|
||||
vfloat32m8_t v = reconstruct_m8_components(code, i, vl);
|
||||
__riscv_vse32_v_f32m8(x + i, v, vl);
|
||||
i += vl;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct QuantizerLloydMax<4, SIMDLevel::RISCV_RVV>
|
||||
: QuantizerLloydMax<4, SIMDLevel::NONE> {
|
||||
using Base = QuantizerLloydMax<4, SIMDLevel::NONE>;
|
||||
|
||||
QuantizerLloydMax(size_t d, const std::vector<float>& trained)
|
||||
: Base(d, trained) {}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
reconstruct_m8_components(const uint8_t* code, size_t i, size_t vl) const {
|
||||
const size_t byte_base = i >> 1;
|
||||
const size_t byte_vl = ((i + vl - 1) >> 1) - byte_base + 1;
|
||||
vuint8m2_t packed = __riscv_vle8_v_u8m2(code + byte_base, byte_vl);
|
||||
vuint32m8_t packed32 = __riscv_vzext_vf4_u32m8(packed, byte_vl);
|
||||
vuint32m8_t comp =
|
||||
__riscv_vadd_vx_u32m8(__riscv_vid_v_u32m8(vl), i, vl);
|
||||
vuint32m8_t rel = __riscv_vsub_vx_u32m8(
|
||||
__riscv_vsrl_vx_u32m8(comp, 1, vl), byte_base, vl);
|
||||
vuint32m8_t bytes = __riscv_vrgather_vv_u32m8(packed32, rel, vl);
|
||||
vuint32m8_t lo = __riscv_vand_vx_u32m8(bytes, 0xf, vl);
|
||||
vuint32m8_t hi = __riscv_vsrl_vx_u32m8(bytes, 4, vl);
|
||||
vbool4_t odd = __riscv_vmsne_vx_u32m8_b4(
|
||||
__riscv_vand_vx_u32m8(comp, 1, vl), 0, vl);
|
||||
vuint32m8_t idx = __riscv_vmerge_vvm_u32m8(lo, hi, odd, vl);
|
||||
vuint32m8_t off = __riscv_vsll_vx_u32m8(idx, 2, vl);
|
||||
return __riscv_vluxei32_v_f32m8(this->centroids, off, vl);
|
||||
}
|
||||
|
||||
void decode_vector(const uint8_t* code, float* x) const final {
|
||||
size_t i = 0;
|
||||
while (i < this->d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(this->d - i);
|
||||
vfloat32m8_t v = reconstruct_m8_components(code, i, vl);
|
||||
__riscv_vse32_v_f32m8(x + i, v, vl);
|
||||
i += vl;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct QuantizerLloydMax<8, SIMDLevel::RISCV_RVV>
|
||||
: QuantizerLloydMax<8, SIMDLevel::NONE> {
|
||||
using Base = QuantizerLloydMax<8, SIMDLevel::NONE>;
|
||||
|
||||
QuantizerLloydMax(size_t d, const std::vector<float>& trained)
|
||||
: Base(d, trained) {}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
reconstruct_m8_components(const uint8_t* code, size_t i, size_t vl) const {
|
||||
vuint8m2_t vb = __riscv_vle8_v_u8m2(code + i, vl);
|
||||
vuint32m8_t off =
|
||||
__riscv_vsll_vx_u32m8(__riscv_vzext_vf4_u32m8(vb, vl), 2, vl);
|
||||
return __riscv_vluxei32_v_f32m8(this->centroids, off, vl);
|
||||
}
|
||||
|
||||
void decode_vector(const uint8_t* code, float* x) const final {
|
||||
size_t i = 0;
|
||||
while (i < this->d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(this->d - i);
|
||||
vfloat32m8_t v = reconstruct_m8_components(code, i, vl);
|
||||
__riscv_vse32_v_f32m8(x + i, v, vl);
|
||||
i += vl;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimilarityL2<SIMDLevel::RISCV_RVV> : SimilarityL2<SIMDLevel::NONE> {
|
||||
using SimilarityL2<SIMDLevel::NONE>::SimilarityL2;
|
||||
|
||||
static constexpr SIMDLevel simd_level = SIMDLevel::RISCV_RVV;
|
||||
|
||||
FAISS_ALWAYS_INLINE void begin_m8() {
|
||||
yi = y;
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE vfloat32m8_t zero_m8(size_t vl) {
|
||||
return __riscv_vfmv_v_f_f32m8(0.0f, vl);
|
||||
}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
add_m8_components(vfloat32m8_t accu, vfloat32m8_t x, size_t vl) {
|
||||
vfloat32m8_t yiv = __riscv_vle32_v_f32m8(yi, vl);
|
||||
yi += vl;
|
||||
vfloat32m8_t tmp = __riscv_vfsub_vv_f32m8(yiv, x, vl);
|
||||
return __riscv_vfmacc_vv_f32m8_tu(accu, tmp, tmp, vl);
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE vfloat32m8_t add_m8_components_2(
|
||||
vfloat32m8_t accu,
|
||||
vfloat32m8_t x,
|
||||
vfloat32m8_t y_2,
|
||||
size_t vl) {
|
||||
vfloat32m8_t tmp = __riscv_vfsub_vv_f32m8(y_2, x, vl);
|
||||
return __riscv_vfmacc_vv_f32m8_tu(accu, tmp, tmp, vl);
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE float result_m8(vfloat32m8_t accu, size_t vl) {
|
||||
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
|
||||
vfloat32m1_t sum = __riscv_vfredusum_vs_f32m8_f32m1(accu, zero, vl);
|
||||
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
||||
}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct SimilarityIP<SIMDLevel::RISCV_RVV> : SimilarityIP<SIMDLevel::NONE> {
|
||||
using SimilarityIP<SIMDLevel::NONE>::SimilarityIP;
|
||||
|
||||
static constexpr SIMDLevel simd_level = SIMDLevel::RISCV_RVV;
|
||||
|
||||
FAISS_ALWAYS_INLINE void begin_m8() {
|
||||
yi = y;
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE vfloat32m8_t zero_m8(size_t vl) {
|
||||
return __riscv_vfmv_v_f_f32m8(0.0f, vl);
|
||||
}
|
||||
|
||||
FAISS_ALWAYS_INLINE vfloat32m8_t
|
||||
add_m8_components(vfloat32m8_t accu, vfloat32m8_t x, size_t vl) {
|
||||
vfloat32m8_t yiv = __riscv_vle32_v_f32m8(yi, vl);
|
||||
yi += vl;
|
||||
return __riscv_vfmacc_vv_f32m8_tu(accu, yiv, x, vl);
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE vfloat32m8_t add_m8_components_2(
|
||||
vfloat32m8_t accu,
|
||||
vfloat32m8_t x1,
|
||||
vfloat32m8_t x2,
|
||||
size_t vl) {
|
||||
return __riscv_vfmacc_vv_f32m8_tu(accu, x1, x2, vl);
|
||||
}
|
||||
|
||||
static FAISS_ALWAYS_INLINE float result_m8(vfloat32m8_t accu, size_t vl) {
|
||||
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
|
||||
vfloat32m1_t sum = __riscv_vfredusum_vs_f32m8_f32m1(accu, zero, vl);
|
||||
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
||||
}
|
||||
};
|
||||
|
||||
/*************************************************************************
|
||||
@@ -127,7 +380,14 @@ struct SimilarityIP<SIMDLevel::RISCV_RVV> : SimilarityIP<SIMDLevel::NONE> {
|
||||
* falls through to scalar code. Callers and the dispatcher don't know or care.
|
||||
************************************************************************/
|
||||
|
||||
template <class Quantizer>
|
||||
inline constexpr bool has_reconstruct_m8_v =
|
||||
requires(const Quantizer& q, const uint8_t* code, size_t i, size_t vl) {
|
||||
q.reconstruct_m8_components(code, i, vl);
|
||||
};
|
||||
|
||||
template <class Quantizer, class Similarity>
|
||||
requires(!has_reconstruct_m8_v<Quantizer>)
|
||||
struct DCTemplate<Quantizer, Similarity, SIMDLevel::RISCV_RVV>
|
||||
: DCTemplate<Quantizer, Similarity, SIMDLevel::NONE> {
|
||||
using Base = DCTemplate<Quantizer, Similarity, SIMDLevel::NONE>;
|
||||
@@ -141,6 +401,84 @@ struct DistanceComputerByte<Similarity, SIMDLevel::RISCV_RVV>
|
||||
using Base::Base;
|
||||
};
|
||||
|
||||
template <class Quantizer, class Similarity>
|
||||
requires(has_reconstruct_m8_v<Quantizer>)
|
||||
struct DCTemplate<Quantizer, Similarity, SIMDLevel::RISCV_RVV>
|
||||
: SQDistanceComputer {
|
||||
using Sim = Similarity;
|
||||
|
||||
Quantizer quant;
|
||||
|
||||
DCTemplate(size_t d, const std::vector<float>& trained)
|
||||
: quant(d, trained) {}
|
||||
|
||||
float compute_distance(const float* x, const uint8_t* code) const {
|
||||
if (quant.d == 0) {
|
||||
return 0.0f;
|
||||
}
|
||||
Similarity sim(x);
|
||||
sim.begin_m8();
|
||||
const size_t first_vl = __riscv_vsetvl_e32m8(quant.d);
|
||||
vfloat32m8_t accu = Sim::zero_m8(first_vl);
|
||||
size_t i = 0;
|
||||
while (i < quant.d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(quant.d - i);
|
||||
vfloat32m8_t xi = quant.reconstruct_m8_components(code, i, vl);
|
||||
accu = sim.add_m8_components(accu, xi, vl);
|
||||
i += vl;
|
||||
}
|
||||
return Sim::result_m8(accu, first_vl);
|
||||
}
|
||||
|
||||
float compute_code_distance(const uint8_t* code1, const uint8_t* code2)
|
||||
const {
|
||||
if (quant.d == 0) {
|
||||
return 0.0f;
|
||||
}
|
||||
Similarity sim(nullptr);
|
||||
sim.begin_m8();
|
||||
const size_t first_vl = __riscv_vsetvl_e32m8(quant.d);
|
||||
vfloat32m8_t accu = Sim::zero_m8(first_vl);
|
||||
size_t i = 0;
|
||||
while (i < quant.d) {
|
||||
size_t vl = __riscv_vsetvl_e32m8(quant.d - i);
|
||||
vfloat32m8_t x1 = quant.reconstruct_m8_components(code1, i, vl);
|
||||
vfloat32m8_t x2 = quant.reconstruct_m8_components(code2, i, vl);
|
||||
accu = Sim::add_m8_components_2(accu, x1, x2, vl);
|
||||
i += vl;
|
||||
}
|
||||
return Sim::result_m8(accu, first_vl);
|
||||
}
|
||||
|
||||
void set_query(const float* x) final {
|
||||
q = x;
|
||||
}
|
||||
|
||||
float symmetric_dis(idx_t i, idx_t j) override {
|
||||
return compute_code_distance(
|
||||
codes + i * code_size, codes + j * code_size);
|
||||
}
|
||||
|
||||
float query_to_code(const uint8_t* code) const final {
|
||||
return compute_distance(q, code);
|
||||
}
|
||||
|
||||
void query_to_codes_batch_4(
|
||||
const uint8_t* code_0,
|
||||
const uint8_t* code_1,
|
||||
const uint8_t* code_2,
|
||||
const uint8_t* code_3,
|
||||
float& dis0,
|
||||
float& dis1,
|
||||
float& dis2,
|
||||
float& dis3) const final {
|
||||
dis0 = compute_distance(q, code_0);
|
||||
dis1 = compute_distance(q, code_1);
|
||||
dis2 = compute_distance(q, code_2);
|
||||
dis3 = compute_distance(q, code_3);
|
||||
}
|
||||
};
|
||||
|
||||
// * Fast path — QT_4bit_uniform + L2 (float domain)
|
||||
// *
|
||||
// * 4-bit UNIFORM scaling: every component reconstructs as
|
||||
|
||||
@@ -342,7 +342,10 @@ SIMDLevel SIMDConfig::auto_detect_simd_level() {
|
||||
#endif
|
||||
|
||||
#if defined(__riscv) && defined(COMPILE_SIMD_RISCV_RVV)
|
||||
// RVV is always available on RISC-V builds compiled with rv64gcv.
|
||||
// RVV is always available on RISC-V builds compiled with
|
||||
// rv64gcv_zvfhmin: that ISA (including Zvfhmin, used by the QT_fp16
|
||||
// kernels) is the minimum requirement to run such binaries at all,
|
||||
// so no runtime feature check is needed beyond the build contract.
|
||||
supported_simd_levels |= (1 << static_cast<int>(SIMDLevel::RISCV_RVV));
|
||||
detected_level = SIMDLevel::RISCV_RVV;
|
||||
#endif
|
||||
|
||||
@@ -26,7 +26,11 @@ enum class SIMDLevel {
|
||||
ARM_NEON,
|
||||
ARM_SVE, // Scalable Vector Extension (ARMv8.2+)
|
||||
// riscv
|
||||
RISCV_RVV, // RISC-V Vector Extension (rv64gcv)
|
||||
// RISC-V Vector Extension. The minimum ISA for this level is rv64gcv
|
||||
// plus Zvfhmin (FP16 vector conversion), which the QT_fp16 kernels
|
||||
// require; RVV translation units are compiled with
|
||||
// -march=rv64gcv_zvfhmin.
|
||||
RISCV_RVV,
|
||||
|
||||
// Appended to preserve the numeric values of the existing public enum.
|
||||
// AVX-512 core features plus AVX512_VPOPCNTDQ and AVX512_BITALG
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
|
||||
#include <gtest/gtest.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <array>
|
||||
#include <cmath>
|
||||
#include <cstring>
|
||||
@@ -93,6 +94,24 @@ std::vector<faiss::SIMDLevel> available_lloyd_max_simd_levels() {
|
||||
return levels;
|
||||
}
|
||||
|
||||
// SIMD levels with scalar-quantizer kernels that can be checked for parity
|
||||
// against the scalar (NONE) path. Unlike dispatch tests, parity tests do not
|
||||
// require dimension restrictions, so RISCV_RVV (no alignment constraints)
|
||||
// is included here.
|
||||
std::vector<faiss::SIMDLevel> available_sq_parity_simd_levels() {
|
||||
std::vector<faiss::SIMDLevel> levels;
|
||||
for (faiss::SIMDLevel level :
|
||||
{faiss::SIMDLevel::AVX512,
|
||||
faiss::SIMDLevel::AVX2,
|
||||
faiss::SIMDLevel::ARM_NEON,
|
||||
faiss::SIMDLevel::RISCV_RVV}) {
|
||||
if (faiss::SIMDConfig::is_simd_level_available(level)) {
|
||||
levels.push_back(level);
|
||||
}
|
||||
}
|
||||
return levels;
|
||||
}
|
||||
|
||||
template <int NBits>
|
||||
void expect_lloyd_max_simd_dispatch_for_compatible_dim(
|
||||
faiss::SIMDLevel level,
|
||||
@@ -256,6 +275,156 @@ void check_lloyd_max_distance_path_parity(
|
||||
}
|
||||
}
|
||||
|
||||
// Parity between a SIMD level and the scalar (NONE) path for one quantizer
|
||||
// type, metric and dimension. Covers quantizer decoding, the
|
||||
// SQDistanceComputer APIs (query_to_code, query_to_codes_batch_4,
|
||||
// symmetric_dis) and the inverted-list scanner's distance_to_code.
|
||||
void check_sq_distance_path_parity(
|
||||
faiss::SIMDLevel level,
|
||||
faiss::ScalarQuantizer::QuantizerType qtype,
|
||||
faiss::MetricType metric,
|
||||
size_t d) {
|
||||
ScopedSIMDLevel scoped(level);
|
||||
const size_t n = 64;
|
||||
std::vector<float> xb = make_normalized_vectors(n, d);
|
||||
std::vector<float> xq = make_normalized_vectors(1, d);
|
||||
|
||||
faiss::ScalarQuantizer sq(d, qtype);
|
||||
sq.train(n, xb.data());
|
||||
|
||||
std::vector<uint8_t> codes(sq.code_size * n, 0);
|
||||
sq.compute_codes(xb.data(), codes.data(), n);
|
||||
|
||||
// fp16/bf16 codes are raw IEEE bit patterns: arbitrary bytes can decode
|
||||
// to NaN, so only codes produced by compute_codes are checked for those.
|
||||
const bool arbitrary_bytes_ok = qtype != faiss::ScalarQuantizer::QT_fp16 &&
|
||||
qtype != faiss::ScalarQuantizer::QT_bf16;
|
||||
|
||||
std::vector<uint8_t> zero_code(sq.code_size, 0);
|
||||
std::vector<uint8_t> max_code(sq.code_size, 0xff);
|
||||
std::vector<const uint8_t*> query_codes = {
|
||||
codes.data(),
|
||||
codes.data() + sq.code_size,
|
||||
codes.data() + 2 * sq.code_size,
|
||||
codes.data() + 3 * sq.code_size};
|
||||
if (arbitrary_bytes_ok) {
|
||||
query_codes.push_back(zero_code.data());
|
||||
query_codes.push_back(max_code.data());
|
||||
}
|
||||
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQuantizer> scalar_quant(
|
||||
faiss::scalar_quantizer::sq_select_quantizer<
|
||||
faiss::SIMDLevel::NONE>(qtype, d, sq.trained));
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQuantizer> simd_quant(
|
||||
sq.select_quantizer());
|
||||
ASSERT_NE(scalar_quant, nullptr);
|
||||
ASSERT_NE(simd_quant, nullptr);
|
||||
|
||||
for (const uint8_t* code : query_codes) {
|
||||
std::vector<float> ref(d), out(d);
|
||||
scalar_quant->decode_vector(code, ref.data());
|
||||
simd_quant->decode_vector(code, out.data());
|
||||
for (size_t j = 0; j < d; j++) {
|
||||
EXPECT_NEAR(
|
||||
ref[j], out[j], 1e-5 * std::max(1.0f, std::fabs(ref[j])));
|
||||
}
|
||||
}
|
||||
|
||||
auto expect_distance_near = [](float ref, float out) {
|
||||
EXPECT_NEAR(ref, out, 5e-3 * std::max(1.0f, std::fabs(ref)));
|
||||
};
|
||||
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQDistanceComputer> scalar_dc(
|
||||
faiss::scalar_quantizer::sq_select_distance_computer<
|
||||
faiss::SIMDLevel::NONE>(metric, qtype, d, sq.trained));
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQDistanceComputer> simd_dc(
|
||||
sq.get_distance_computer(metric));
|
||||
ASSERT_NE(scalar_dc, nullptr);
|
||||
ASSERT_NE(simd_dc, nullptr);
|
||||
|
||||
scalar_dc->set_query(xq.data());
|
||||
simd_dc->set_query(xq.data());
|
||||
|
||||
for (const uint8_t* code : query_codes) {
|
||||
expect_distance_near(
|
||||
scalar_dc->query_to_code(code), simd_dc->query_to_code(code));
|
||||
}
|
||||
|
||||
float scalar_dis[4], simd_dis[4];
|
||||
scalar_dc->query_to_codes_batch_4(
|
||||
query_codes[0],
|
||||
query_codes[1],
|
||||
query_codes[2],
|
||||
query_codes[3],
|
||||
scalar_dis[0],
|
||||
scalar_dis[1],
|
||||
scalar_dis[2],
|
||||
scalar_dis[3]);
|
||||
simd_dc->query_to_codes_batch_4(
|
||||
query_codes[0],
|
||||
query_codes[1],
|
||||
query_codes[2],
|
||||
query_codes[3],
|
||||
simd_dis[0],
|
||||
simd_dis[1],
|
||||
simd_dis[2],
|
||||
simd_dis[3]);
|
||||
for (int k = 0; k < 4; k++) {
|
||||
expect_distance_near(scalar_dis[k], simd_dis[k]);
|
||||
}
|
||||
|
||||
std::vector<uint8_t> bundle(4 * sq.code_size, 0);
|
||||
for (int k = 0; k < 4; k++) {
|
||||
std::memcpy(
|
||||
bundle.data() + k * sq.code_size, query_codes[k], sq.code_size);
|
||||
}
|
||||
scalar_dc->codes = bundle.data();
|
||||
scalar_dc->code_size = sq.code_size;
|
||||
simd_dc->codes = bundle.data();
|
||||
simd_dc->code_size = sq.code_size;
|
||||
|
||||
const std::array<std::pair<faiss::idx_t, faiss::idx_t>, 4> pairs = {{
|
||||
{0, 1},
|
||||
{0, 2},
|
||||
{1, 3},
|
||||
{2, 3},
|
||||
}};
|
||||
for (const auto& [lhs, rhs] : pairs) {
|
||||
expect_distance_near(
|
||||
scalar_dc->symmetric_dis(lhs, rhs),
|
||||
simd_dc->symmetric_dis(lhs, rhs));
|
||||
}
|
||||
|
||||
std::unique_ptr<faiss::InvertedListScanner> scalar_scanner(
|
||||
faiss::scalar_quantizer::sq_select_InvertedListScanner<
|
||||
faiss::SIMDLevel::NONE>(
|
||||
qtype,
|
||||
metric,
|
||||
d,
|
||||
sq.code_size,
|
||||
sq.trained,
|
||||
nullptr,
|
||||
false,
|
||||
nullptr,
|
||||
false));
|
||||
std::unique_ptr<faiss::InvertedListScanner> simd_scanner(
|
||||
sq.select_InvertedListScanner(
|
||||
metric, nullptr, false, nullptr, false));
|
||||
ASSERT_NE(scalar_scanner, nullptr);
|
||||
ASSERT_NE(simd_scanner, nullptr);
|
||||
|
||||
scalar_scanner->set_query(xq.data());
|
||||
simd_scanner->set_query(xq.data());
|
||||
scalar_scanner->set_list(0, 0.0f);
|
||||
simd_scanner->set_list(0, 0.0f);
|
||||
|
||||
for (const uint8_t* code : query_codes) {
|
||||
expect_distance_near(
|
||||
scalar_scanner->distance_to_code(code),
|
||||
simd_scanner->distance_to_code(code));
|
||||
}
|
||||
}
|
||||
|
||||
void check_tqmse_roundtrip(
|
||||
size_t d,
|
||||
faiss::ScalarQuantizer::QuantizerType qtype) {
|
||||
@@ -652,7 +821,7 @@ TEST(ScalarQuantizer, EDENSimdDispatchSelection) {
|
||||
|
||||
TEST(ScalarQuantizer, TQMSESimdDistancePathParity) {
|
||||
const std::vector<faiss::SIMDLevel> levels =
|
||||
available_lloyd_max_simd_levels();
|
||||
available_sq_parity_simd_levels();
|
||||
if (levels.empty()) {
|
||||
GTEST_SKIP() << "No SIMD level available for TurboQuant parity tests";
|
||||
}
|
||||
@@ -674,7 +843,7 @@ TEST(ScalarQuantizer, TQMSESimdDistancePathParity) {
|
||||
|
||||
TEST(ScalarQuantizer, EDENSimdDistancePathParity) {
|
||||
const std::vector<faiss::SIMDLevel> levels =
|
||||
available_lloyd_max_simd_levels();
|
||||
available_sq_parity_simd_levels();
|
||||
if (levels.empty()) {
|
||||
GTEST_SKIP() << "No SIMD level available for EDEN parity tests";
|
||||
}
|
||||
@@ -693,3 +862,225 @@ TEST(ScalarQuantizer, EDENSimdDistancePathParity) {
|
||||
level, faiss::ScalarQuantizer::QT_8bit_eden);
|
||||
}
|
||||
}
|
||||
|
||||
// RVV-versus-scalar parity for all quantizers with RVV kernels, over both
|
||||
// metrics and dimensions around VLMAX boundaries. With e32m8 vectors,
|
||||
// VLMAX = VLEN / 4: 32 lanes on QEMU's default VLEN=128 and up to 512
|
||||
// lanes on VLEN=2048 hardware, so the dimensions below straddle those
|
||||
// chunk boundaries as well as tiny tails.
|
||||
// QT_fp16 is covered by the dedicated RVVFP16DistancePathParity test below
|
||||
// so that the Zvfhmin-dependent case is visible individually in test logs.
|
||||
TEST(ScalarQuantizer, RVVDistancePathParity) {
|
||||
if (!faiss::SIMDConfig::is_simd_level_available(
|
||||
faiss::SIMDLevel::RISCV_RVV)) {
|
||||
GTEST_SKIP() << "RISCV_RVV not available on this machine";
|
||||
}
|
||||
|
||||
const std::vector<faiss::ScalarQuantizer::QuantizerType> qtypes = {
|
||||
faiss::ScalarQuantizer::QT_8bit,
|
||||
faiss::ScalarQuantizer::QT_4bit,
|
||||
faiss::ScalarQuantizer::QT_6bit,
|
||||
faiss::ScalarQuantizer::QT_8bit_uniform,
|
||||
faiss::ScalarQuantizer::QT_4bit_uniform,
|
||||
faiss::ScalarQuantizer::QT_bf16,
|
||||
faiss::ScalarQuantizer::QT_8bit_direct,
|
||||
faiss::ScalarQuantizer::QT_8bit_direct_signed,
|
||||
faiss::ScalarQuantizer::QT_1bit_eden,
|
||||
faiss::ScalarQuantizer::QT_2bit_eden,
|
||||
faiss::ScalarQuantizer::QT_3bit_eden,
|
||||
faiss::ScalarQuantizer::QT_4bit_eden,
|
||||
faiss::ScalarQuantizer::QT_8bit_eden,
|
||||
// The QT_*_tqmse types currently alias QuantizerLloydMax, so
|
||||
// they share the EDEN RVV kernels; listing them makes that
|
||||
// dispatch coverage explicit for every metric and distance API.
|
||||
faiss::ScalarQuantizer::QT_1bit_tqmse,
|
||||
faiss::ScalarQuantizer::QT_2bit_tqmse,
|
||||
faiss::ScalarQuantizer::QT_3bit_tqmse,
|
||||
faiss::ScalarQuantizer::QT_4bit_tqmse,
|
||||
faiss::ScalarQuantizer::QT_8bit_tqmse,
|
||||
};
|
||||
const std::vector<size_t> dims = {
|
||||
1,
|
||||
7,
|
||||
31,
|
||||
32,
|
||||
33,
|
||||
63,
|
||||
64,
|
||||
65,
|
||||
127,
|
||||
128,
|
||||
129,
|
||||
255,
|
||||
256,
|
||||
257,
|
||||
511,
|
||||
512,
|
||||
513};
|
||||
const std::vector<faiss::MetricType> metrics = {
|
||||
faiss::METRIC_L2, faiss::METRIC_INNER_PRODUCT};
|
||||
|
||||
for (auto qtype : qtypes) {
|
||||
for (size_t d : dims) {
|
||||
for (auto metric : metrics) {
|
||||
SCOPED_TRACE(
|
||||
testing::Message()
|
||||
<< "qtype=" << static_cast<int>(qtype) << " d=" << d
|
||||
<< " metric=" << metric);
|
||||
check_sq_distance_path_parity(
|
||||
faiss::SIMDLevel::RISCV_RVV, qtype, metric, d);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// QT_fp16 RVV-versus-scalar parity, split out of RVVDistancePathParity so
|
||||
// the Zvfhmin-dependent case shows up individually in test logs. The RVV
|
||||
// FP16 kernel emits Zvfhmin instructions, so this test only runs where the
|
||||
// CPU/emulator enables the extension (CI uses -cpu rv64,v=true,
|
||||
// x-zvfhmin=true); without it the process would die with SIGILL, which is
|
||||
// what makes this test a check that the FP16 case genuinely executes.
|
||||
TEST(ScalarQuantizer, RVVFP16DistancePathParity) {
|
||||
if (!faiss::SIMDConfig::is_simd_level_available(
|
||||
faiss::SIMDLevel::RISCV_RVV)) {
|
||||
GTEST_SKIP() << "RISCV_RVV not available on this machine";
|
||||
}
|
||||
|
||||
const std::vector<size_t> dims = {
|
||||
1,
|
||||
7,
|
||||
31,
|
||||
32,
|
||||
33,
|
||||
63,
|
||||
64,
|
||||
65,
|
||||
127,
|
||||
128,
|
||||
129,
|
||||
255,
|
||||
256,
|
||||
257,
|
||||
511,
|
||||
512,
|
||||
513};
|
||||
const std::vector<faiss::MetricType> metrics = {
|
||||
faiss::METRIC_L2, faiss::METRIC_INNER_PRODUCT};
|
||||
|
||||
for (size_t d : dims) {
|
||||
for (auto metric : metrics) {
|
||||
SCOPED_TRACE(
|
||||
testing::Message() << "d=" << d << " metric=" << metric);
|
||||
check_sq_distance_path_parity(
|
||||
faiss::SIMDLevel::RISCV_RVV,
|
||||
faiss::ScalarQuantizer::QT_fp16,
|
||||
metric,
|
||||
d);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Zero-dimensional regression for the d == 0 early returns in the RVV
|
||||
// DCTemplate distance loops: query_to_code, query_to_codes_batch_4 and
|
||||
// symmetric_dis must return exactly 0 without touching the query pointer or
|
||||
// the codes, matching the scalar (NONE) path. The parity dimension lists
|
||||
// above start at 1, so this case is covered here directly.
|
||||
TEST(ScalarQuantizer, RVVZeroDimDistancePathParity) {
|
||||
if (!faiss::SIMDConfig::is_simd_level_available(
|
||||
faiss::SIMDLevel::RISCV_RVV)) {
|
||||
GTEST_SKIP() << "RISCV_RVV not available on this machine";
|
||||
}
|
||||
ScopedSIMDLevel scoped(faiss::SIMDLevel::RISCV_RVV);
|
||||
|
||||
// trained layouts each quantizer expects at d == 0: non-uniform
|
||||
// templates take 2 * d (empty) values, uniform templates take {vmin,
|
||||
// vdiff}, LloydMax takes 2^k - 1 centroids/boundaries, and the raw
|
||||
// codecs (fp16/bf16/direct) ignore trained entirely.
|
||||
const std::vector<std::pair<
|
||||
faiss::ScalarQuantizer::QuantizerType,
|
||||
std::vector<float>>>
|
||||
cases = {
|
||||
{faiss::ScalarQuantizer::QT_8bit, {}},
|
||||
{faiss::ScalarQuantizer::QT_4bit, {}},
|
||||
{faiss::ScalarQuantizer::QT_6bit, {}},
|
||||
{faiss::ScalarQuantizer::QT_8bit_uniform, {0.0f, 1.0f}},
|
||||
{faiss::ScalarQuantizer::QT_4bit_uniform, {0.0f, 1.0f}},
|
||||
{faiss::ScalarQuantizer::QT_fp16, {}},
|
||||
{faiss::ScalarQuantizer::QT_bf16, {}},
|
||||
{faiss::ScalarQuantizer::QT_8bit_direct, {}},
|
||||
{faiss::ScalarQuantizer::QT_8bit_direct_signed, {}},
|
||||
{faiss::ScalarQuantizer::QT_1bit_eden,
|
||||
std::vector<float>(3, 0.0f)},
|
||||
{faiss::ScalarQuantizer::QT_2bit_eden,
|
||||
std::vector<float>(7, 0.0f)},
|
||||
{faiss::ScalarQuantizer::QT_3bit_eden,
|
||||
std::vector<float>(15, 0.0f)},
|
||||
{faiss::ScalarQuantizer::QT_4bit_eden,
|
||||
std::vector<float>(31, 0.0f)},
|
||||
{faiss::ScalarQuantizer::QT_8bit_eden,
|
||||
std::vector<float>(511, 0.0f)},
|
||||
};
|
||||
const std::vector<faiss::MetricType> metrics = {
|
||||
faiss::METRIC_L2, faiss::METRIC_INNER_PRODUCT};
|
||||
|
||||
uint8_t dummy_code[8] = {};
|
||||
|
||||
for (const auto& [qtype, trained] : cases) {
|
||||
for (auto metric : metrics) {
|
||||
SCOPED_TRACE(
|
||||
testing::Message() << "qtype=" << static_cast<int>(qtype)
|
||||
<< " metric=" << metric);
|
||||
|
||||
faiss::ScalarQuantizer sq(0, qtype);
|
||||
sq.trained = trained;
|
||||
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQDistanceComputer>
|
||||
scalar_dc(
|
||||
faiss::scalar_quantizer::
|
||||
sq_select_distance_computer<
|
||||
faiss::SIMDLevel::NONE>(
|
||||
metric, qtype, 0, sq.trained));
|
||||
std::unique_ptr<faiss::ScalarQuantizer::SQDistanceComputer> simd_dc(
|
||||
sq.get_distance_computer(metric));
|
||||
ASSERT_NE(scalar_dc, nullptr);
|
||||
ASSERT_NE(simd_dc, nullptr);
|
||||
|
||||
scalar_dc->set_query(nullptr);
|
||||
simd_dc->set_query(nullptr);
|
||||
|
||||
EXPECT_FLOAT_EQ(scalar_dc->query_to_code(dummy_code), 0.0f);
|
||||
EXPECT_FLOAT_EQ(simd_dc->query_to_code(dummy_code), 0.0f);
|
||||
|
||||
float scalar_dis[4], simd_dis[4];
|
||||
scalar_dc->query_to_codes_batch_4(
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
scalar_dis[0],
|
||||
scalar_dis[1],
|
||||
scalar_dis[2],
|
||||
scalar_dis[3]);
|
||||
simd_dc->query_to_codes_batch_4(
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
dummy_code,
|
||||
simd_dis[0],
|
||||
simd_dis[1],
|
||||
simd_dis[2],
|
||||
simd_dis[3]);
|
||||
for (int k = 0; k < 4; k++) {
|
||||
EXPECT_FLOAT_EQ(scalar_dis[k], 0.0f);
|
||||
EXPECT_FLOAT_EQ(simd_dis[k], 0.0f);
|
||||
}
|
||||
|
||||
scalar_dc->codes = dummy_code;
|
||||
scalar_dc->code_size = 0;
|
||||
simd_dc->codes = dummy_code;
|
||||
simd_dc->code_size = 0;
|
||||
EXPECT_FLOAT_EQ(scalar_dc->symmetric_dis(0, 1), 0.0f);
|
||||
EXPECT_FLOAT_EQ(simd_dc->symmetric_dis(0, 1), 0.0f);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user