faiss 0.6.1 → 0.6.3
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/CHANGELOG.md +8 -0
- data/lib/faiss/version.rb +1 -1
- data/vendor/faiss/faiss/AutoTune.cpp +3 -1
- data/vendor/faiss/faiss/Clustering.cpp +9 -1
- data/vendor/faiss/faiss/IVFlib.cpp +14 -3
- data/vendor/faiss/faiss/Index.h +2 -2
- data/vendor/faiss/faiss/IndexAdditiveQuantizer.cpp +9 -10
- data/vendor/faiss/faiss/IndexAdditiveQuantizerFastScan.cpp +2 -3
- data/vendor/faiss/faiss/IndexBinaryFromFloat.cpp +1 -2
- data/vendor/faiss/faiss/IndexBinaryHNSW.cpp +10 -12
- data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
- data/vendor/faiss/faiss/IndexBinaryIVF.cpp +5 -7
- data/vendor/faiss/faiss/IndexEDEN.cpp +273 -0
- data/vendor/faiss/faiss/IndexEDEN.h +57 -0
- data/vendor/faiss/faiss/IndexFastScan.cpp +15 -4
- data/vendor/faiss/faiss/IndexFlat.cpp +13 -50
- data/vendor/faiss/faiss/IndexHNSW.cpp +177 -148
- data/vendor/faiss/faiss/IndexIDMap.cpp +16 -3
- data/vendor/faiss/faiss/IndexIDMap.h +2 -0
- data/vendor/faiss/faiss/IndexIVF.cpp +19 -8
- data/vendor/faiss/faiss/IndexIVFAdditiveQuantizer.cpp +3 -3
- data/vendor/faiss/faiss/IndexIVFAdditiveQuantizerFastScan.cpp +3 -4
- data/vendor/faiss/faiss/IndexIVFEDEN.cpp +302 -0
- data/vendor/faiss/faiss/IndexIVFEDEN.h +70 -0
- data/vendor/faiss/faiss/IndexIVFFastScan.cpp +5 -6
- data/vendor/faiss/faiss/IndexIVFFlat.cpp +6 -5
- data/vendor/faiss/faiss/IndexIVFFlatPanorama.cpp +3 -3
- data/vendor/faiss/faiss/IndexIVFIndependentQuantizer.cpp +1 -1
- data/vendor/faiss/faiss/IndexIVFPQ.cpp +42 -25
- data/vendor/faiss/faiss/IndexIVFPQFastScan.cpp +0 -1
- data/vendor/faiss/faiss/IndexIVFPQR.cpp +2 -3
- data/vendor/faiss/faiss/IndexIVFRaBitQ.cpp +23 -62
- data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.cpp +180 -76
- data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.h +5 -4
- data/vendor/faiss/faiss/IndexIVFSpectralHash.cpp +8 -6
- data/vendor/faiss/faiss/IndexLSH.cpp +2 -3
- data/vendor/faiss/faiss/IndexLattice.cpp +5 -0
- data/vendor/faiss/faiss/IndexNNDescent.cpp +10 -3
- data/vendor/faiss/faiss/IndexNSG.cpp +8 -4
- data/vendor/faiss/faiss/IndexPQ.cpp +6 -8
- data/vendor/faiss/faiss/IndexPreTransform.cpp +15 -0
- data/vendor/faiss/faiss/IndexRaBitQ.cpp +2 -2
- data/vendor/faiss/faiss/IndexRaBitQFastScan.cpp +1 -2
- data/vendor/faiss/faiss/IndexRaBitQFastScan.h +5 -1
- data/vendor/faiss/faiss/IndexRefine.cpp +30 -1
- data/vendor/faiss/faiss/IndexReplicas.cpp +1 -2
- data/vendor/faiss/faiss/IndexScalarQuantizer.cpp +68 -6
- data/vendor/faiss/faiss/IndexScalarQuantizer.h +10 -0
- data/vendor/faiss/faiss/IndexShards.cpp +2 -2
- data/vendor/faiss/faiss/IndexShardsIVF.cpp +2 -2
- data/vendor/faiss/faiss/MetaIndexes.cpp +2 -4
- data/vendor/faiss/faiss/SuperKMeans.cpp +256 -240
- data/vendor/faiss/faiss/SuperKMeans.h +30 -0
- data/vendor/faiss/faiss/VectorTransform.cpp +33 -2
- data/vendor/faiss/faiss/clone_index.cpp +5 -0
- data/vendor/faiss/faiss/cppcontrib/SaDecodeKernels.h +1 -1
- data/vendor/faiss/faiss/cppcontrib/sa_decode/Level2-neon-inl.h +902 -12
- data/vendor/faiss/faiss/cppcontrib/sa_decode/PQ-neon-inl.h +702 -10
- data/vendor/faiss/faiss/factory_tools.cpp +51 -4
- data/vendor/faiss/faiss/gpu/GpuCloner.cpp +11 -11
- data/vendor/faiss/faiss/gpu/GpuIndex.h +34 -11
- data/vendor/faiss/faiss/gpu/GpuIndexCagra.h +47 -0
- data/vendor/faiss/faiss/gpu/GpuIndexIVF.h +17 -0
- data/vendor/faiss/faiss/gpu/GpuIndexIVFScalarQuantizer.h +16 -0
- data/vendor/faiss/faiss/gpu/GpuResources.h +3 -2
- data/vendor/faiss/faiss/gpu/StandardGpuResources.cpp +11 -12
- data/vendor/faiss/faiss/gpu/StandardGpuResources.h +3 -3
- data/vendor/faiss/faiss/gpu/perf/PerfClustering.cpp +1 -1
- data/vendor/faiss/faiss/gpu/perf/PerfIVFPQAdd.cpp +2 -2
- data/vendor/faiss/faiss/gpu/test/TestGpuIndexIVFScalarQuantizer.cpp +180 -0
- data/vendor/faiss/faiss/gpu_metal/MetalDistance.h +87 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndex.h +7 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFFlat.h +177 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFPQ.h +88 -0
- data/vendor/faiss/faiss/gpu_metal/MetalKernels.h +48 -3
- data/vendor/faiss/faiss/gpu_metal/MetalPythonBridge.h +45 -0
- data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFFlat.h +193 -0
- data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFPQ.h +134 -0
- data/vendor/faiss/faiss/impl/ClusteringInitialization.cpp +2 -2
- data/vendor/faiss/faiss/impl/DistanceComputer.h +34 -0
- data/vendor/faiss/faiss/impl/EDENQuantizer.h +119 -0
- data/vendor/faiss/faiss/impl/HNSW.cpp +658 -344
- data/vendor/faiss/faiss/impl/HNSW.h +51 -13
- data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
- data/vendor/faiss/faiss/impl/NSG.cpp +18 -12
- data/vendor/faiss/faiss/impl/Panorama.h +20 -7
- data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
- data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +59 -24
- data/vendor/faiss/faiss/impl/RaBitQUtils.cpp +45 -37
- data/vendor/faiss/faiss/impl/RaBitQUtils.h +35 -0
- data/vendor/faiss/faiss/impl/RaBitQuantizer.cpp +175 -68
- data/vendor/faiss/faiss/impl/RaBitQuantizer.h +19 -0
- data/vendor/faiss/faiss/impl/RaBitQuantizerMultiBit.cpp +2 -11
- data/vendor/faiss/faiss/impl/ResultHandler.h +26 -31
- data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +522 -58
- data/vendor/faiss/faiss/impl/ScalarQuantizer.h +70 -0
- data/vendor/faiss/faiss/impl/ThreadedIndex-inl.h +2 -2
- data/vendor/faiss/faiss/impl/VisitedTable.cpp +33 -13
- data/vendor/faiss/faiss/impl/VisitedTable.h +88 -33
- data/vendor/faiss/faiss/impl/binary_hamming/IndexBinaryIVF_impl.h +1 -1
- data/vendor/faiss/faiss/impl/binary_hamming/avx2.cpp +4 -4
- data/vendor/faiss/faiss/impl/fast_scan/dispatching.h +38 -3
- data/vendor/faiss/faiss/impl/hnsw/LockVector.cpp +1 -1
- data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.cpp +35 -43
- data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.h +64 -15
- data/vendor/faiss/faiss/impl/hnsw/avx2.cpp +86 -40
- data/vendor/faiss/faiss/impl/hnsw/avx512.cpp +81 -50
- data/vendor/faiss/faiss/impl/index_read.cpp +476 -75
- data/vendor/faiss/faiss/impl/index_write.cpp +56 -4
- data/vendor/faiss/faiss/impl/io_macros.h +25 -0
- data/vendor/faiss/faiss/impl/lattice_Zn.cpp +8 -9
- data/vendor/faiss/faiss/impl/platform_macros.h +15 -9
- data/vendor/faiss/faiss/impl/polysemous_training/avx512.cpp +284 -0
- data/vendor/faiss/faiss/impl/polysemous_training/dispatch.h +115 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/IVFPQ_QueryTables.cpp +0 -1
- data/vendor/faiss/faiss/impl/pq_code_distance/PQDistanceComputer_impl.h +26 -15
- data/vendor/faiss/faiss/impl/pq_code_distance/avx2.cpp +6 -4
- data/vendor/faiss/faiss/impl/pq_code_distance/avx512.cpp +2 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/neon.cpp +2 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-generic.cpp +20 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-inl.h +36 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-sve.cpp +5 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_scan_impl.h +105 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/rvv.cpp +2 -0
- data/vendor/faiss/faiss/impl/result_handler/ResultHandler.cpp +195 -0
- data/vendor/faiss/faiss/impl/result_handler/avx2.cpp +133 -0
- data/vendor/faiss/faiss/impl/result_handler/avx512.cpp +281 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx2.cpp +72 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx512.cpp +228 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer.cpp +882 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/distance_computers.h +6 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/quantizers.h +336 -26
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +331 -32
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-impl.h +553 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-spr.cpp +558 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512.cpp +284 -45
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-dispatch.h +502 -3
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +157 -32
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-rvv.cpp +26 -0
- data/vendor/faiss/faiss/impl/simd_dispatch.h +86 -8
- data/vendor/faiss/faiss/index_factory.cpp +37 -7
- data/vendor/faiss/faiss/index_io.h +16 -0
- data/vendor/faiss/faiss/invlists/DirectMap.cpp +5 -2
- data/vendor/faiss/faiss/invlists/InvertedLists.cpp +15 -15
- data/vendor/faiss/faiss/invlists/InvertedLists.h +2 -2
- data/vendor/faiss/faiss/invlists/OnDiskInvertedLists.cpp +19 -4
- data/vendor/faiss/faiss/python/python_callbacks.cpp +3 -1
- data/vendor/faiss/faiss/svs/IndexSVSFaissUtils.h +60 -0
- data/vendor/faiss/faiss/svs/IndexSVSFlat.cpp +26 -1
- data/vendor/faiss/faiss/svs/IndexSVSFlat.h +13 -0
- data/vendor/faiss/faiss/svs/IndexSVSIVF.cpp +1 -1
- data/vendor/faiss/faiss/svs/IndexSVSIVFLeanVec.cpp +1 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamana.cpp +150 -23
- data/vendor/faiss/faiss/svs/IndexSVSVamana.h +30 -7
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.cpp +3 -2
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.h +2 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +65 -25
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.h +3 -2
- data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
- data/vendor/faiss/faiss/utils/bf16.h +34 -0
- data/vendor/faiss/faiss/utils/distances.cpp +14 -2
- data/vendor/faiss/faiss/utils/distances_simd.cpp +4 -4
- data/vendor/faiss/faiss/utils/extra_distances.cpp +4 -14
- data/vendor/faiss/faiss/utils/extra_distances.h +1 -2
- data/vendor/faiss/faiss/utils/hamming.cpp +9 -9
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx2.cpp +2 -1
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_spr.cpp +15 -0
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512.h +6 -30
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512_spr.h +171 -0
- data/vendor/faiss/faiss/utils/partitioning.cpp +0 -2
- data/vendor/faiss/faiss/utils/quantize_lut.cpp +29 -8
- data/vendor/faiss/faiss/utils/rabitq_simd.h +202 -0
- data/vendor/faiss/faiss/utils/simd_impl/distances_avx2.cpp +0 -1
- data/vendor/faiss/faiss/utils/simd_impl/distances_avx512.cpp +263 -15
- data/vendor/faiss/faiss/utils/simd_impl/distances_rvv.cpp +160 -18
- data/vendor/faiss/faiss/utils/simd_impl/partitioning_simdlib256.h +14 -68
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx2.cpp +245 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512.cpp +273 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512_spr.cpp +435 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_neon.cpp +11 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_rvv.cpp +143 -6
- data/vendor/faiss/faiss/utils/simd_levels.cpp +56 -2
- data/vendor/faiss/faiss/utils/simd_levels.h +14 -0
- data/vendor/faiss/faiss/utils/utils.cpp +9 -27
- metadata +27 -2
|
@@ -12,12 +12,24 @@
|
|
|
12
12
|
#ifdef COMPILE_SIMD_RISCV_RVV
|
|
13
13
|
|
|
14
14
|
#include <faiss/utils/extra_distances.h>
|
|
15
|
+
#include <riscv_vector.h>
|
|
15
16
|
|
|
16
17
|
namespace faiss {
|
|
17
18
|
|
|
18
19
|
template <>
|
|
19
20
|
float fvec_norm_L2sqr<SIMDLevel::RISCV_RVV>(const float* x, size_t d) {
|
|
20
|
-
|
|
21
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
22
|
+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
23
|
+
size_t i = 0;
|
|
24
|
+
while (i < d) {
|
|
25
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
26
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
27
|
+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, vx, vx, vl);
|
|
28
|
+
i += vl;
|
|
29
|
+
}
|
|
30
|
+
vfloat32m1_t sum = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
31
|
+
sum = __riscv_vfredusum_vs_f32m8_f32m1(acc, sum, vlmax);
|
|
32
|
+
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
|
21
33
|
}
|
|
22
34
|
|
|
23
35
|
template <>
|
|
@@ -25,7 +37,20 @@ float fvec_L2sqr<SIMDLevel::RISCV_RVV>(
|
|
|
25
37
|
const float* x,
|
|
26
38
|
const float* y,
|
|
27
39
|
size_t d) {
|
|
28
|
-
|
|
40
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
41
|
+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
42
|
+
size_t i = 0;
|
|
43
|
+
while (i < d) {
|
|
44
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
45
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
46
|
+
vfloat32m8_t vy = __riscv_vle32_v_f32m8(y + i, vl);
|
|
47
|
+
vx = __riscv_vfsub_vv_f32m8(vx, vy, vl);
|
|
48
|
+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, vx, vx, vl);
|
|
49
|
+
i += vl;
|
|
50
|
+
}
|
|
51
|
+
vfloat32m1_t sum = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
52
|
+
sum = __riscv_vfredusum_vs_f32m8_f32m1(acc, sum, vlmax);
|
|
53
|
+
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
|
29
54
|
}
|
|
30
55
|
|
|
31
56
|
template <>
|
|
@@ -33,12 +58,38 @@ float fvec_inner_product<SIMDLevel::RISCV_RVV>(
|
|
|
33
58
|
const float* x,
|
|
34
59
|
const float* y,
|
|
35
60
|
size_t d) {
|
|
36
|
-
|
|
61
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
62
|
+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
63
|
+
size_t i = 0;
|
|
64
|
+
while (i < d) {
|
|
65
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
66
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
67
|
+
vfloat32m8_t vy = __riscv_vle32_v_f32m8(y + i, vl);
|
|
68
|
+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, vx, vy, vl);
|
|
69
|
+
i += vl;
|
|
70
|
+
}
|
|
71
|
+
vfloat32m1_t sum = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
72
|
+
sum = __riscv_vfredusum_vs_f32m8_f32m1(acc, sum, vlmax);
|
|
73
|
+
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
|
37
74
|
}
|
|
38
75
|
|
|
39
76
|
template <>
|
|
40
77
|
float fvec_L1<SIMDLevel::RISCV_RVV>(const float* x, const float* y, size_t d) {
|
|
41
|
-
|
|
78
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
79
|
+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
80
|
+
size_t i = 0;
|
|
81
|
+
while (i < d) {
|
|
82
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
83
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
84
|
+
vfloat32m8_t vy = __riscv_vle32_v_f32m8(y + i, vl);
|
|
85
|
+
vx = __riscv_vfsub_vv_f32m8(vx, vy, vl);
|
|
86
|
+
vx = __riscv_vfsgnjx_vv_f32m8(vx, vx, vl);
|
|
87
|
+
acc = __riscv_vfadd_vv_f32m8_tu(acc, acc, vx, vl);
|
|
88
|
+
i += vl;
|
|
89
|
+
}
|
|
90
|
+
vfloat32m1_t sum = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
91
|
+
sum = __riscv_vfredusum_vs_f32m8_f32m1(acc, sum, vlmax);
|
|
92
|
+
return __riscv_vfmv_f_s_f32m1_f32(sum);
|
|
42
93
|
}
|
|
43
94
|
|
|
44
95
|
template <>
|
|
@@ -46,7 +97,21 @@ float fvec_Linf<SIMDLevel::RISCV_RVV>(
|
|
|
46
97
|
const float* x,
|
|
47
98
|
const float* y,
|
|
48
99
|
size_t d) {
|
|
49
|
-
|
|
100
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
101
|
+
vfloat32m8_t vmax = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
102
|
+
size_t i = 0;
|
|
103
|
+
while (i < d) {
|
|
104
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
105
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
106
|
+
vfloat32m8_t vy = __riscv_vle32_v_f32m8(y + i, vl);
|
|
107
|
+
vx = __riscv_vfsub_vv_f32m8(vx, vy, vl);
|
|
108
|
+
vx = __riscv_vfsgnjx_vv_f32m8(vx, vx, vl);
|
|
109
|
+
vmax = __riscv_vfmax_vv_f32m8_tu(vmax, vmax, vx, vl);
|
|
110
|
+
i += vl;
|
|
111
|
+
}
|
|
112
|
+
vfloat32m1_t max = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
113
|
+
max = __riscv_vfredmax_vs_f32m8_f32m1(vmax, max, vlmax);
|
|
114
|
+
return __riscv_vfmv_f_s_f32m1_f32(max);
|
|
50
115
|
}
|
|
51
116
|
|
|
52
117
|
template <>
|
|
@@ -61,8 +126,10 @@ void fvec_inner_product_batch_4<SIMDLevel::RISCV_RVV>(
|
|
|
61
126
|
float& dis1,
|
|
62
127
|
float& dis2,
|
|
63
128
|
float& dis3) {
|
|
64
|
-
|
|
65
|
-
|
|
129
|
+
dis0 = fvec_inner_product<SIMDLevel::RISCV_RVV>(x, y0, d);
|
|
130
|
+
dis1 = fvec_inner_product<SIMDLevel::RISCV_RVV>(x, y1, d);
|
|
131
|
+
dis2 = fvec_inner_product<SIMDLevel::RISCV_RVV>(x, y2, d);
|
|
132
|
+
dis3 = fvec_inner_product<SIMDLevel::RISCV_RVV>(x, y3, d);
|
|
66
133
|
}
|
|
67
134
|
|
|
68
135
|
template <>
|
|
@@ -77,8 +144,10 @@ void fvec_L2sqr_batch_4<SIMDLevel::RISCV_RVV>(
|
|
|
77
144
|
float& dis1,
|
|
78
145
|
float& dis2,
|
|
79
146
|
float& dis3) {
|
|
80
|
-
|
|
81
|
-
|
|
147
|
+
dis0 = fvec_L2sqr<SIMDLevel::RISCV_RVV>(x, y0, d);
|
|
148
|
+
dis1 = fvec_L2sqr<SIMDLevel::RISCV_RVV>(x, y1, d);
|
|
149
|
+
dis2 = fvec_L2sqr<SIMDLevel::RISCV_RVV>(x, y2, d);
|
|
150
|
+
dis3 = fvec_L2sqr<SIMDLevel::RISCV_RVV>(x, y3, d);
|
|
82
151
|
}
|
|
83
152
|
|
|
84
153
|
template <>
|
|
@@ -90,8 +159,33 @@ void fvec_L2sqr_ny_transposed<SIMDLevel::RISCV_RVV>(
|
|
|
90
159
|
size_t d,
|
|
91
160
|
size_t d_offset,
|
|
92
161
|
size_t ny) {
|
|
93
|
-
|
|
94
|
-
|
|
162
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
163
|
+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
|
|
164
|
+
size_t i = 0;
|
|
165
|
+
while (i < d) {
|
|
166
|
+
size_t vl = __riscv_vsetvl_e32m8(d - i);
|
|
167
|
+
vfloat32m8_t vx = __riscv_vle32_v_f32m8(x + i, vl);
|
|
168
|
+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, vx, vx, vl);
|
|
169
|
+
i += vl;
|
|
170
|
+
}
|
|
171
|
+
vfloat32m1_t sum = __riscv_vfmv_s_f_f32m1(0.0f, 1);
|
|
172
|
+
sum = __riscv_vfredusum_vs_f32m8_f32m1(acc, sum, vlmax);
|
|
173
|
+
float x_sqlen = __riscv_vfmv_f_s_f32m1_f32(sum);
|
|
174
|
+
i = 0;
|
|
175
|
+
while (i < ny) {
|
|
176
|
+
size_t vl = __riscv_vsetvl_e32m8(ny - i);
|
|
177
|
+
acc = __riscv_vfmv_v_f_f32m8(0.0f, vl);
|
|
178
|
+
for (size_t j = 0; j < d; j++) {
|
|
179
|
+
vfloat32m8_t vy = __riscv_vle32_v_f32m8(y + j * d_offset + i, vl);
|
|
180
|
+
acc = __riscv_vfmacc_vf_f32m8(acc, x[j], vy, vl);
|
|
181
|
+
}
|
|
182
|
+
vfloat32m8_t vres = __riscv_vle32_v_f32m8(y_sqlen + i, vl);
|
|
183
|
+
vres = __riscv_vfadd_vf_f32m8(vres, x_sqlen, vl);
|
|
184
|
+
acc = __riscv_vfmul_vf_f32m8(acc, 2.0f, vl);
|
|
185
|
+
vres = __riscv_vfsub_vv_f32m8(vres, acc, vl);
|
|
186
|
+
__riscv_vse32_v_f32m8(dis + i, vres, vl);
|
|
187
|
+
i += vl;
|
|
188
|
+
}
|
|
95
189
|
}
|
|
96
190
|
|
|
97
191
|
template <>
|
|
@@ -101,7 +195,10 @@ void fvec_inner_products_ny<SIMDLevel::RISCV_RVV>(
|
|
|
101
195
|
const float* y,
|
|
102
196
|
size_t d,
|
|
103
197
|
size_t ny) {
|
|
104
|
-
|
|
198
|
+
for (size_t i = 0; i < ny; i++) {
|
|
199
|
+
ip[i] = fvec_inner_product<SIMDLevel::RISCV_RVV>(x, y, d);
|
|
200
|
+
y += d;
|
|
201
|
+
}
|
|
105
202
|
}
|
|
106
203
|
|
|
107
204
|
template <>
|
|
@@ -111,7 +208,39 @@ void fvec_L2sqr_ny<SIMDLevel::RISCV_RVV>(
|
|
|
111
208
|
const float* y,
|
|
112
209
|
size_t d,
|
|
113
210
|
size_t ny) {
|
|
114
|
-
|
|
211
|
+
for (size_t i = 0; i < ny; i++) {
|
|
212
|
+
dis[i] = fvec_L2sqr<SIMDLevel::RISCV_RVV>(x, y, d);
|
|
213
|
+
y += d;
|
|
214
|
+
}
|
|
215
|
+
}
|
|
216
|
+
|
|
217
|
+
// Index of the first element equal to the minimum of values[0..n), or n when
|
|
218
|
+
// there is none (e.g. n == 0). Shared by the *_nearest and madd_and_argmin
|
|
219
|
+
// kernels so the vfmin/vfredmin/vfirst sequence lives in one place.
|
|
220
|
+
static size_t rvv_argmin(const float* values, size_t n) {
|
|
221
|
+
size_t vlmax = __riscv_vsetvlmax_e32m8();
|
|
222
|
+
vfloat32m8_t vmin = __riscv_vfmv_v_f_f32m8(__builtin_inff(), vlmax);
|
|
223
|
+
size_t i = 0;
|
|
224
|
+
while (i < n) {
|
|
225
|
+
size_t vl = __riscv_vsetvl_e32m8(n - i);
|
|
226
|
+
vfloat32m8_t vd = __riscv_vle32_v_f32m8(values + i, vl);
|
|
227
|
+
vmin = __riscv_vfmin_vv_f32m8_tu(vmin, vmin, vd, vl);
|
|
228
|
+
i += vl;
|
|
229
|
+
}
|
|
230
|
+
vfloat32m1_t rmin = __riscv_vfmv_s_f_f32m1(__builtin_inff(), 1);
|
|
231
|
+
rmin = __riscv_vfredmin_vs_f32m8_f32m1(vmin, rmin, vlmax);
|
|
232
|
+
float min_val = __riscv_vfmv_f_s_f32m1_f32(rmin);
|
|
233
|
+
i = 0;
|
|
234
|
+
while (i < n) {
|
|
235
|
+
size_t vl = __riscv_vsetvl_e32m8(n - i);
|
|
236
|
+
vfloat32m8_t vd = __riscv_vle32_v_f32m8(values + i, vl);
|
|
237
|
+
long j = __riscv_vfirst_m_b4(
|
|
238
|
+
__riscv_vmfeq_vf_f32m8_b4(vd, min_val, vl), vl);
|
|
239
|
+
if (j >= 0)
|
|
240
|
+
return i + static_cast<size_t>(j);
|
|
241
|
+
i += vl;
|
|
242
|
+
}
|
|
243
|
+
return n;
|
|
115
244
|
}
|
|
116
245
|
|
|
117
246
|
template <>
|
|
@@ -121,8 +250,9 @@ size_t fvec_L2sqr_ny_nearest<SIMDLevel::RISCV_RVV>(
|
|
|
121
250
|
const float* y,
|
|
122
251
|
size_t d,
|
|
123
252
|
size_t ny) {
|
|
124
|
-
|
|
125
|
-
|
|
253
|
+
fvec_L2sqr_ny<SIMDLevel::RISCV_RVV>(distances_tmp_buffer, x, y, d, ny);
|
|
254
|
+
const size_t j = rvv_argmin(distances_tmp_buffer, ny);
|
|
255
|
+
return j < ny ? j : 0;
|
|
126
256
|
}
|
|
127
257
|
|
|
128
258
|
template <>
|
|
@@ -134,8 +264,10 @@ size_t fvec_L2sqr_ny_nearest_y_transposed<SIMDLevel::RISCV_RVV>(
|
|
|
134
264
|
size_t d,
|
|
135
265
|
size_t d_offset,
|
|
136
266
|
size_t ny) {
|
|
137
|
-
|
|
267
|
+
fvec_L2sqr_ny_transposed<SIMDLevel::RISCV_RVV>(
|
|
138
268
|
distances_tmp_buffer, x, y, y_sqlen, d, d_offset, ny);
|
|
269
|
+
const size_t j = rvv_argmin(distances_tmp_buffer, ny);
|
|
270
|
+
return j < ny ? j : 0;
|
|
139
271
|
}
|
|
140
272
|
|
|
141
273
|
template <>
|
|
@@ -145,7 +277,15 @@ void fvec_madd<SIMDLevel::RISCV_RVV>(
|
|
|
145
277
|
float bf,
|
|
146
278
|
const float* b,
|
|
147
279
|
float* c) {
|
|
148
|
-
|
|
280
|
+
size_t i = 0;
|
|
281
|
+
while (i < n) {
|
|
282
|
+
size_t vl = __riscv_vsetvl_e32m8(n - i);
|
|
283
|
+
vfloat32m8_t va = __riscv_vle32_v_f32m8(a + i, vl);
|
|
284
|
+
vfloat32m8_t vb = __riscv_vle32_v_f32m8(b + i, vl);
|
|
285
|
+
va = __riscv_vfmacc_vf_f32m8(va, bf, vb, vl);
|
|
286
|
+
__riscv_vse32_v_f32m8(c + i, va, vl);
|
|
287
|
+
i += vl;
|
|
288
|
+
}
|
|
149
289
|
}
|
|
150
290
|
|
|
151
291
|
template <>
|
|
@@ -155,7 +295,9 @@ int fvec_madd_and_argmin<SIMDLevel::RISCV_RVV>(
|
|
|
155
295
|
float bf,
|
|
156
296
|
const float* b,
|
|
157
297
|
float* c) {
|
|
158
|
-
|
|
298
|
+
fvec_madd<SIMDLevel::RISCV_RVV>(n, a, bf, b, c);
|
|
299
|
+
const size_t j = rvv_argmin(c, n);
|
|
300
|
+
return j < n ? static_cast<int>(j) : -1;
|
|
159
301
|
}
|
|
160
302
|
|
|
161
303
|
#define DEFINE_VECTOR_DISTANCE_RVV_FALLBACK(metric) \
|
|
@@ -592,39 +592,12 @@ simd16uint16 accu8to16(simd32uint8 a8) {
|
|
|
592
592
|
return hadd(a8_0, a8_1);
|
|
593
593
|
}
|
|
594
594
|
|
|
595
|
-
|
|
596
|
-
|
|
597
|
-
|
|
598
|
-
|
|
599
|
-
0,
|
|
600
|
-
4,
|
|
601
|
-
64,
|
|
602
|
-
0,
|
|
603
|
-
0,
|
|
604
|
-
0,
|
|
605
|
-
0,
|
|
606
|
-
1,
|
|
607
|
-
16,
|
|
608
|
-
0,
|
|
609
|
-
0,
|
|
610
|
-
4,
|
|
611
|
-
64,
|
|
612
|
-
1,
|
|
613
|
-
16,
|
|
614
|
-
0,
|
|
615
|
-
0,
|
|
616
|
-
4,
|
|
617
|
-
64,
|
|
618
|
-
0,
|
|
619
|
-
0,
|
|
620
|
-
0,
|
|
621
|
-
0,
|
|
622
|
-
1,
|
|
623
|
-
16,
|
|
624
|
-
0,
|
|
625
|
-
0,
|
|
626
|
-
4,
|
|
627
|
-
64>();
|
|
595
|
+
// Lookup table held as a plain byte array in .rodata. Storing it as a
|
|
596
|
+
// `simd32uint8` global would emit an AVX2 initializer into `.init_array` that
|
|
597
|
+
// runs at dlopen, before runtime SIMD dispatch, and SIGILLs on non-AVX2 CPUs
|
|
598
|
+
alignas(32) static const uint8_t shifts[32] = {
|
|
599
|
+
1, 16, 0, 0, 4, 64, 0, 0, 0, 0, 1, 16, 0, 0, 4, 64,
|
|
600
|
+
1, 16, 0, 0, 4, 64, 0, 0, 0, 0, 1, 16, 0, 0, 4, 64};
|
|
628
601
|
|
|
629
602
|
// 2-bit accumulator: we can add only up to 3 elements
|
|
630
603
|
// on output we return 2*4-bit results
|
|
@@ -644,7 +617,8 @@ void compute_accu2(
|
|
|
644
617
|
v = pp(v);
|
|
645
618
|
// 0x800 -> force second half of table
|
|
646
619
|
simd16uint16 idx = v | (v << 8) | simd16uint16(0x800);
|
|
647
|
-
a2 += simd16uint16(
|
|
620
|
+
a2 += simd16uint16(
|
|
621
|
+
simd32uint8(shifts).lookup_2_lanes(simd32uint8(idx)));
|
|
648
622
|
}
|
|
649
623
|
a4lo += a2 & mask2;
|
|
650
624
|
a4hi += (a2 >> 2) & mask2;
|
|
@@ -694,39 +668,11 @@ simd16uint16 histogram_8(const uint16_t* data, Preproc pp, size_t n_in) {
|
|
|
694
668
|
* 16 bins
|
|
695
669
|
************************************************************/
|
|
696
670
|
|
|
697
|
-
|
|
698
|
-
|
|
699
|
-
|
|
700
|
-
4,
|
|
701
|
-
8,
|
|
702
|
-
16,
|
|
703
|
-
32,
|
|
704
|
-
64,
|
|
705
|
-
128,
|
|
706
|
-
1,
|
|
707
|
-
2,
|
|
708
|
-
4,
|
|
709
|
-
8,
|
|
710
|
-
16,
|
|
711
|
-
32,
|
|
712
|
-
64,
|
|
713
|
-
128,
|
|
714
|
-
1,
|
|
715
|
-
2,
|
|
716
|
-
4,
|
|
717
|
-
8,
|
|
718
|
-
16,
|
|
719
|
-
32,
|
|
720
|
-
64,
|
|
721
|
-
128,
|
|
722
|
-
1,
|
|
723
|
-
2,
|
|
724
|
-
4,
|
|
725
|
-
8,
|
|
726
|
-
16,
|
|
727
|
-
32,
|
|
728
|
-
64,
|
|
729
|
-
128>();
|
|
671
|
+
// See the note on `shifts` above: kept as a .rodata byte array so its
|
|
672
|
+
// initializer does not emit AVX2 into `.init_array`
|
|
673
|
+
alignas(32) static const uint8_t shifts2[32] = {
|
|
674
|
+
1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128,
|
|
675
|
+
1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128};
|
|
730
676
|
|
|
731
677
|
simd32uint8 shiftr_16(simd32uint8 x, int n) {
|
|
732
678
|
return simd32uint8(simd16uint16(x) >> n);
|
|
@@ -754,7 +700,7 @@ void compute_accu2_16(
|
|
|
754
700
|
v = pp(v);
|
|
755
701
|
|
|
756
702
|
simd16uint16 idx = v | (v << 8);
|
|
757
|
-
simd32uint8 a1 = shifts2.lookup_2_lanes(simd32uint8(idx));
|
|
703
|
+
simd32uint8 a1 = simd32uint8(shifts2).lookup_2_lanes(simd32uint8(idx));
|
|
758
704
|
// contains 0s for out-of-bounds elements
|
|
759
705
|
|
|
760
706
|
simd16uint16 lt8 = (v >> 3) == simd16uint16(0);
|
|
@@ -9,6 +9,7 @@
|
|
|
9
9
|
|
|
10
10
|
#include <faiss/utils/rabitq_simd.h>
|
|
11
11
|
#include <immintrin.h>
|
|
12
|
+
#include <limits>
|
|
12
13
|
|
|
13
14
|
namespace faiss::rabitq {
|
|
14
15
|
|
|
@@ -82,8 +83,161 @@ inline uint64_t reduce_add_128(__m128i v) {
|
|
|
82
83
|
return lanes[0] + lanes[1];
|
|
83
84
|
}
|
|
84
85
|
|
|
86
|
+
inline float reduce_min_256(__m256 v) {
|
|
87
|
+
__m128 x =
|
|
88
|
+
_mm_min_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
|
|
89
|
+
x = _mm_min_ps(x, _mm_movehl_ps(x, x));
|
|
90
|
+
x = _mm_min_ss(x, _mm_shuffle_ps(x, x, 1));
|
|
91
|
+
return _mm_cvtss_f32(x);
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
inline float reduce_max_256(__m256 v) {
|
|
95
|
+
__m128 x =
|
|
96
|
+
_mm_max_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
|
|
97
|
+
x = _mm_max_ps(x, _mm_movehl_ps(x, x));
|
|
98
|
+
x = _mm_max_ss(x, _mm_shuffle_ps(x, x, 1));
|
|
99
|
+
return _mm_cvtss_f32(x);
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
inline __m256i round_nonnegative_ps_to_i32(__m256 x) {
|
|
103
|
+
return _mm256_cvttps_epi32(_mm256_add_ps(x, _mm256_set1_ps(0.5f)));
|
|
104
|
+
}
|
|
105
|
+
|
|
106
|
+
inline void store_i32_as_u8_8(__m256i values, uint8_t* out) {
|
|
107
|
+
const __m128i packed16 = _mm_packus_epi32(
|
|
108
|
+
_mm256_castsi256_si128(values),
|
|
109
|
+
_mm256_extracti128_si256(values, 1));
|
|
110
|
+
const __m128i packed8 = _mm_packus_epi16(packed16, _mm_setzero_si128());
|
|
111
|
+
_mm_storel_epi64(reinterpret_cast<__m128i*>(out), packed8);
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
inline void accumulate_i32_as_i64(
|
|
115
|
+
__m256i values,
|
|
116
|
+
__m256i& low_acc,
|
|
117
|
+
__m256i& high_acc) {
|
|
118
|
+
low_acc = _mm256_add_epi64(
|
|
119
|
+
low_acc, _mm256_cvtepi32_epi64(_mm256_castsi256_si128(values)));
|
|
120
|
+
high_acc = _mm256_add_epi64(
|
|
121
|
+
high_acc,
|
|
122
|
+
_mm256_cvtepi32_epi64(_mm256_extracti128_si256(values, 1)));
|
|
123
|
+
}
|
|
124
|
+
|
|
85
125
|
} // namespace
|
|
86
126
|
|
|
127
|
+
template <>
|
|
128
|
+
void lut_minmax_16<SIMDLevel::AVX2>(const float* tab, float& mn, float& mx) {
|
|
129
|
+
const __m256 lo = _mm256_loadu_ps(tab);
|
|
130
|
+
const __m256 hi = _mm256_loadu_ps(tab + 8);
|
|
131
|
+
const __m256 min_vec = _mm256_min_ps(lo, hi);
|
|
132
|
+
const __m256 max_vec = _mm256_max_ps(lo, hi);
|
|
133
|
+
mn = reduce_min_256(min_vec);
|
|
134
|
+
mx = reduce_max_256(max_vec);
|
|
135
|
+
}
|
|
136
|
+
|
|
137
|
+
template <>
|
|
138
|
+
void minmax_values<SIMDLevel::AVX2>(
|
|
139
|
+
const float* values,
|
|
140
|
+
size_t n,
|
|
141
|
+
float& mn,
|
|
142
|
+
float& mx) {
|
|
143
|
+
if (n == 0) {
|
|
144
|
+
return;
|
|
145
|
+
}
|
|
146
|
+
|
|
147
|
+
size_t i = 0;
|
|
148
|
+
__m256 min_vec = _mm256_set1_ps(std::numeric_limits<float>::max());
|
|
149
|
+
__m256 max_vec = _mm256_set1_ps(std::numeric_limits<float>::lowest());
|
|
150
|
+
for (; i + 8 <= n; i += 8) {
|
|
151
|
+
const __m256 values_vec = _mm256_loadu_ps(values + i);
|
|
152
|
+
min_vec = _mm256_min_ps(min_vec, values_vec);
|
|
153
|
+
max_vec = _mm256_max_ps(max_vec, values_vec);
|
|
154
|
+
}
|
|
155
|
+
|
|
156
|
+
mn = reduce_min_256(min_vec);
|
|
157
|
+
mx = reduce_max_256(max_vec);
|
|
158
|
+
for (; i < n; i++) {
|
|
159
|
+
mn = std::min(mn, values[i]);
|
|
160
|
+
mx = std::max(mx, values[i]);
|
|
161
|
+
}
|
|
162
|
+
}
|
|
163
|
+
|
|
164
|
+
template <>
|
|
165
|
+
void lut_quantize_16_to_uint8<SIMDLevel::AVX2>(
|
|
166
|
+
const float* tab,
|
|
167
|
+
float mn,
|
|
168
|
+
float a,
|
|
169
|
+
uint8_t* out) {
|
|
170
|
+
const __m256 a_vec = _mm256_set1_ps(a);
|
|
171
|
+
const __m256 mn_times_a_vec = _mm256_set1_ps(mn * a);
|
|
172
|
+
const __m256i zero = _mm256_setzero_si256();
|
|
173
|
+
for (size_t i = 0; i < 16; i += 8) {
|
|
174
|
+
const __m256 values = _mm256_loadu_ps(tab + i);
|
|
175
|
+
const __m256 scaled = _mm256_fmsub_ps(values, a_vec, mn_times_a_vec);
|
|
176
|
+
const __m256i rounded =
|
|
177
|
+
_mm256_max_epi32(round_nonnegative_ps_to_i32(scaled), zero);
|
|
178
|
+
store_i32_as_u8_8(rounded, out + i);
|
|
179
|
+
}
|
|
180
|
+
}
|
|
181
|
+
|
|
182
|
+
template <>
|
|
183
|
+
void quantize_query_values<SIMDLevel::AVX2>(
|
|
184
|
+
const float* rq,
|
|
185
|
+
size_t d,
|
|
186
|
+
float v_min,
|
|
187
|
+
float inv_delta,
|
|
188
|
+
uint8_t max_code,
|
|
189
|
+
bool centered,
|
|
190
|
+
uint8_t* rqq,
|
|
191
|
+
size_t& sum_qq,
|
|
192
|
+
int64_t& sum2_signed_odd_int) {
|
|
193
|
+
const __m256 inv_delta_vec = _mm256_set1_ps(inv_delta);
|
|
194
|
+
const __m256 v_min_times_inv_delta_vec = _mm256_set1_ps(v_min * inv_delta);
|
|
195
|
+
const __m256 zero = _mm256_setzero_ps();
|
|
196
|
+
const __m256 max_code_ps = _mm256_set1_ps(max_code);
|
|
197
|
+
const __m256i max_code_i32 = _mm256_set1_epi32(max_code);
|
|
198
|
+
const __m256i two = _mm256_set1_epi32(2);
|
|
199
|
+
__m256i sum_acc_lo = _mm256_setzero_si256();
|
|
200
|
+
__m256i sum_acc_hi = _mm256_setzero_si256();
|
|
201
|
+
__m256i sq_acc_lo = _mm256_setzero_si256();
|
|
202
|
+
__m256i sq_acc_hi = _mm256_setzero_si256();
|
|
203
|
+
|
|
204
|
+
size_t i = 0;
|
|
205
|
+
for (; i + 8 <= d; i += 8) {
|
|
206
|
+
const __m256 values = _mm256_loadu_ps(rq + i);
|
|
207
|
+
__m256 scaled = _mm256_fmsub_ps(
|
|
208
|
+
values, inv_delta_vec, v_min_times_inv_delta_vec);
|
|
209
|
+
scaled = _mm256_min_ps(_mm256_max_ps(scaled, zero), max_code_ps);
|
|
210
|
+
const __m256i rounded = round_nonnegative_ps_to_i32(scaled);
|
|
211
|
+
accumulate_i32_as_i64(rounded, sum_acc_lo, sum_acc_hi);
|
|
212
|
+
|
|
213
|
+
if (centered) {
|
|
214
|
+
const __m256i signed_odd = _mm256_sub_epi32(
|
|
215
|
+
_mm256_mullo_epi32(rounded, two), max_code_i32);
|
|
216
|
+
const __m256i signed_odd_sqr =
|
|
217
|
+
_mm256_mullo_epi32(signed_odd, signed_odd);
|
|
218
|
+
accumulate_i32_as_i64(signed_odd_sqr, sq_acc_lo, sq_acc_hi);
|
|
219
|
+
}
|
|
220
|
+
store_i32_as_u8_8(rounded, rqq + i);
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
sum_qq += reduce_add_256(sum_acc_lo) + reduce_add_256(sum_acc_hi);
|
|
224
|
+
if (centered) {
|
|
225
|
+
sum2_signed_odd_int +=
|
|
226
|
+
reduce_add_256(sq_acc_lo) + reduce_add_256(sq_acc_hi);
|
|
227
|
+
}
|
|
228
|
+
|
|
229
|
+
for (; i < d; i++) {
|
|
230
|
+
const uint8_t v_qq = round_clamped_byte_scalar(
|
|
231
|
+
(rq[i] - v_min) * inv_delta, max_code);
|
|
232
|
+
rqq[i] = v_qq;
|
|
233
|
+
sum_qq += v_qq;
|
|
234
|
+
if (centered) {
|
|
235
|
+
const int64_t signed_odd_int = int64_t(v_qq) * 2 - max_code;
|
|
236
|
+
sum2_signed_odd_int += signed_odd_int * signed_odd_int;
|
|
237
|
+
}
|
|
238
|
+
}
|
|
239
|
+
}
|
|
240
|
+
|
|
87
241
|
template <>
|
|
88
242
|
uint64_t bitwise_and_dot_product<SIMDLevel::AVX2>(
|
|
89
243
|
const uint8_t* query,
|
|
@@ -137,6 +291,69 @@ uint64_t bitwise_and_dot_product<SIMDLevel::AVX2>(
|
|
|
137
291
|
return sum;
|
|
138
292
|
}
|
|
139
293
|
|
|
294
|
+
template <>
|
|
295
|
+
BitwiseAndDotProductResult bitwise_and_dot_product_with_popcount<
|
|
296
|
+
SIMDLevel::AVX2>(
|
|
297
|
+
const uint8_t* query,
|
|
298
|
+
const uint8_t* data,
|
|
299
|
+
size_t size,
|
|
300
|
+
size_t qb) {
|
|
301
|
+
uint64_t dot_product = 0;
|
|
302
|
+
uint64_t popcount_sum = 0;
|
|
303
|
+
size_t offset = 0;
|
|
304
|
+
if (size_t step = 256 / 8; offset + step <= size) {
|
|
305
|
+
__m256i dot_256 = _mm256_setzero_si256();
|
|
306
|
+
__m256i pop_256 = _mm256_setzero_si256();
|
|
307
|
+
for (; offset + step <= size; offset += step) {
|
|
308
|
+
__m256i v_x = _mm256_loadu_si256((const __m256i*)(data + offset));
|
|
309
|
+
pop_256 = _mm256_add_epi64(pop_256, popcount_256(v_x));
|
|
310
|
+
for (int j = 0; j < qb; j++) {
|
|
311
|
+
__m256i v_q = _mm256_loadu_si256(
|
|
312
|
+
(const __m256i*)(query + j * size + offset));
|
|
313
|
+
__m256i v_and = _mm256_and_si256(v_q, v_x);
|
|
314
|
+
__m256i v_popcnt = popcount_256(v_and);
|
|
315
|
+
__m256i v_shifted = _mm256_slli_epi64(v_popcnt, j);
|
|
316
|
+
dot_256 = _mm256_add_epi64(dot_256, v_shifted);
|
|
317
|
+
}
|
|
318
|
+
}
|
|
319
|
+
dot_product += reduce_add_256(dot_256);
|
|
320
|
+
popcount_sum += reduce_add_256(pop_256);
|
|
321
|
+
}
|
|
322
|
+
__m128i dot_128 = _mm_setzero_si128();
|
|
323
|
+
__m128i pop_128 = _mm_setzero_si128();
|
|
324
|
+
for (size_t step = 128 / 8; offset + step <= size; offset += step) {
|
|
325
|
+
__m128i v_x = _mm_loadu_si128((const __m128i*)(data + offset));
|
|
326
|
+
pop_128 = _mm_add_epi64(pop_128, popcount_128(v_x));
|
|
327
|
+
for (int j = 0; j < qb; j++) {
|
|
328
|
+
__m128i v_q = _mm_loadu_si128(
|
|
329
|
+
(const __m128i*)(query + j * size + offset));
|
|
330
|
+
__m128i v_and = _mm_and_si128(v_q, v_x);
|
|
331
|
+
__m128i v_popcnt = popcount_128(v_and);
|
|
332
|
+
__m128i v_shifted = _mm_slli_epi64(v_popcnt, j);
|
|
333
|
+
dot_128 = _mm_add_epi64(dot_128, v_shifted);
|
|
334
|
+
}
|
|
335
|
+
}
|
|
336
|
+
dot_product += reduce_add_128(dot_128);
|
|
337
|
+
popcount_sum += reduce_add_128(pop_128);
|
|
338
|
+
for (size_t step = 64 / 8; offset + step <= size; offset += step) {
|
|
339
|
+
const uint64_t yv = *(const uint64_t*)(data + offset);
|
|
340
|
+
popcount_sum += popcount64(yv);
|
|
341
|
+
for (int j = 0; j < qb; j++) {
|
|
342
|
+
const uint64_t qv = *(const uint64_t*)(query + j * size + offset);
|
|
343
|
+
dot_product += popcount64(qv & yv) << j;
|
|
344
|
+
}
|
|
345
|
+
}
|
|
346
|
+
for (; offset < size; ++offset) {
|
|
347
|
+
const uint8_t yv = *(data + offset);
|
|
348
|
+
popcount_sum += popcount32(yv);
|
|
349
|
+
for (int j = 0; j < qb; j++) {
|
|
350
|
+
const uint8_t qv = *(query + j * size + offset);
|
|
351
|
+
dot_product += popcount32(qv & yv) << j;
|
|
352
|
+
}
|
|
353
|
+
}
|
|
354
|
+
return {dot_product, popcount_sum};
|
|
355
|
+
}
|
|
356
|
+
|
|
140
357
|
template <>
|
|
141
358
|
uint64_t bitwise_xor_dot_product<SIMDLevel::AVX2>(
|
|
142
359
|
const uint8_t* query,
|
|
@@ -220,6 +437,34 @@ uint64_t popcount<SIMDLevel::AVX2>(const uint8_t* data, size_t size) {
|
|
|
220
437
|
return sum;
|
|
221
438
|
}
|
|
222
439
|
|
|
440
|
+
template <>
|
|
441
|
+
void rearrange_bit_planes<SIMDLevel::AVX2>(
|
|
442
|
+
const uint8_t* rotated_qq,
|
|
443
|
+
size_t d,
|
|
444
|
+
size_t qb,
|
|
445
|
+
uint8_t* out) {
|
|
446
|
+
const size_t offset = (d + 7) / 8;
|
|
447
|
+
memset(out, 0, offset * qb);
|
|
448
|
+
const size_t nchunks = d / 32;
|
|
449
|
+
for (size_t chunk = 0; chunk < nchunks; chunk++) {
|
|
450
|
+
__m256i vals =
|
|
451
|
+
_mm256_loadu_si256((const __m256i*)(rotated_qq + chunk * 32));
|
|
452
|
+
for (size_t iv = 0; iv < qb; iv++) {
|
|
453
|
+
__m256i mask = _mm256_set1_epi8(static_cast<char>(1 << iv));
|
|
454
|
+
__m256i bits =
|
|
455
|
+
_mm256_cmpeq_epi8(_mm256_and_si256(vals, mask), mask);
|
|
456
|
+
uint32_t packed = static_cast<uint32_t>(_mm256_movemask_epi8(bits));
|
|
457
|
+
memcpy(&out[iv * offset + chunk * 4], &packed, 4);
|
|
458
|
+
}
|
|
459
|
+
}
|
|
460
|
+
for (size_t idim = nchunks * 32; idim < d; idim++) {
|
|
461
|
+
for (size_t iv = 0; iv < qb; iv++) {
|
|
462
|
+
const bool bit = ((rotated_qq[idim] & (1 << iv)) != 0);
|
|
463
|
+
out[iv * offset + idim / 8] |= bit ? (1 << (idim % 8)) : 0;
|
|
464
|
+
}
|
|
465
|
+
}
|
|
466
|
+
}
|
|
467
|
+
|
|
223
468
|
} // namespace faiss::rabitq
|
|
224
469
|
|
|
225
470
|
namespace faiss::rabitq::multibit {
|