faiss 0.6.2 → 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 +4 -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 +4 -5
- data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
- data/vendor/faiss/faiss/IndexBinaryIVF.cpp +2 -4
- 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 +10 -11
- data/vendor/faiss/faiss/IndexIDMap.cpp +16 -3
- data/vendor/faiss/faiss/IndexIDMap.h +2 -0
- data/vendor/faiss/faiss/IndexIVF.cpp +17 -6
- data/vendor/faiss/faiss/IndexIVFAdditiveQuantizer.cpp +1 -1
- 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 +3 -4
- data/vendor/faiss/faiss/IndexIVFIndependentQuantizer.cpp +1 -1
- data/vendor/faiss/faiss/IndexIVFPQ.cpp +40 -22
- data/vendor/faiss/faiss/IndexIVFPQFastScan.cpp +0 -1
- data/vendor/faiss/faiss/IndexIVFRaBitQ.cpp +19 -49
- 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 +9 -2
- data/vendor/faiss/faiss/IndexNSG.cpp +7 -2
- 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/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/factory_tools.cpp +47 -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/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/MetalIndexIVFFlat.h +1 -5
- data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFPQ.h +88 -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 +109 -152
- data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
- data/vendor/faiss/faiss/impl/NSG.cpp +3 -1
- data/vendor/faiss/faiss/impl/Panorama.h +9 -7
- data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
- data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +34 -22
- data/vendor/faiss/faiss/impl/RaBitQUtils.cpp +44 -36
- data/vendor/faiss/faiss/impl/RaBitQUtils.h +35 -0
- data/vendor/faiss/faiss/impl/RaBitQuantizer.cpp +168 -67
- 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 +25 -31
- data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +258 -57
- data/vendor/faiss/faiss/impl/ScalarQuantizer.h +20 -0
- data/vendor/faiss/faiss/impl/ThreadedIndex-inl.h +2 -2
- data/vendor/faiss/faiss/impl/VisitedTable.cpp +22 -2
- data/vendor/faiss/faiss/impl/VisitedTable.h +20 -0
- 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 +35 -2
- data/vendor/faiss/faiss/impl/hnsw/LockVector.cpp +1 -1
- data/vendor/faiss/faiss/impl/index_read.cpp +376 -36
- data/vendor/faiss/faiss/impl/index_write.cpp +55 -4
- data/vendor/faiss/faiss/impl/lattice_Zn.cpp +8 -9
- data/vendor/faiss/faiss/impl/platform_macros.h +3 -1
- 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 +4 -4
- 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/quantizers.h +9 -8
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +85 -23
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-impl.h +30 -30
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-spr.cpp +4 -5
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512.cpp +101 -34
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-dispatch.h +136 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +16 -16
- data/vendor/faiss/faiss/impl/simd_dispatch.h +30 -9
- data/vendor/faiss/faiss/index_factory.cpp +32 -6
- data/vendor/faiss/faiss/invlists/DirectMap.cpp +1 -1
- data/vendor/faiss/faiss/invlists/InvertedLists.cpp +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 +31 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamana.h +15 -2
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +1 -2
- data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
- data/vendor/faiss/faiss/utils/distances.cpp +14 -2
- data/vendor/faiss/faiss/utils/distances_simd.cpp +4 -3
- 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 +1 -1
- 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/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 +92 -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 +44 -0
- data/vendor/faiss/faiss/utils/simd_levels.h +14 -0
- data/vendor/faiss/faiss/utils/utils.cpp +9 -27
- metadata +16 -1
|
@@ -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 {
|
|
@@ -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
|
|
|
@@ -167,8 +168,164 @@ inline uint64_t reduce_add_128(__m128i v) {
|
|
|
167
168
|
return lanes[0] + lanes[1];
|
|
168
169
|
}
|
|
169
170
|
|
|
171
|
+
inline __m512i round_nonnegative_ps_to_i32(__m512 x) {
|
|
172
|
+
return _mm512_cvttps_epi32(_mm512_add_ps(x, _mm512_set1_ps(0.5f)));
|
|
173
|
+
}
|
|
174
|
+
|
|
175
|
+
inline void store_i32_as_u8_16(__m512i values, uint8_t* out) {
|
|
176
|
+
const __m128i packed = _mm512_cvtusepi32_epi8(values);
|
|
177
|
+
_mm_storeu_si128(reinterpret_cast<__m128i*>(out), packed);
|
|
178
|
+
}
|
|
179
|
+
|
|
170
180
|
} // namespace
|
|
171
181
|
|
|
182
|
+
template <>
|
|
183
|
+
void lut_minmax_16<SIMDLevel::AVX512>(const float* tab, float& mn, float& mx) {
|
|
184
|
+
const __m512 v = _mm512_loadu_ps(tab);
|
|
185
|
+
mn = _mm512_reduce_min_ps(v);
|
|
186
|
+
mx = _mm512_reduce_max_ps(v);
|
|
187
|
+
}
|
|
188
|
+
|
|
189
|
+
template <>
|
|
190
|
+
void minmax_values<SIMDLevel::AVX512>(
|
|
191
|
+
const float* values,
|
|
192
|
+
size_t n,
|
|
193
|
+
float& mn,
|
|
194
|
+
float& mx) {
|
|
195
|
+
if (n == 0) {
|
|
196
|
+
return;
|
|
197
|
+
}
|
|
198
|
+
|
|
199
|
+
size_t i = 0;
|
|
200
|
+
__m512 mn_vec = _mm512_set1_ps(std::numeric_limits<float>::max());
|
|
201
|
+
__m512 mx_vec = _mm512_set1_ps(std::numeric_limits<float>::lowest());
|
|
202
|
+
for (; i + 16 <= n; i += 16) {
|
|
203
|
+
const __m512 v = _mm512_loadu_ps(values + i);
|
|
204
|
+
mn_vec = _mm512_min_ps(mn_vec, v);
|
|
205
|
+
mx_vec = _mm512_max_ps(mx_vec, v);
|
|
206
|
+
}
|
|
207
|
+
|
|
208
|
+
if (i < n) {
|
|
209
|
+
const __mmask16 mask =
|
|
210
|
+
static_cast<__mmask16>((uint32_t(1) << (n - i)) - 1);
|
|
211
|
+
const __m512 v = _mm512_maskz_loadu_ps(mask, values + i);
|
|
212
|
+
mn_vec = _mm512_mask_min_ps(mn_vec, mask, mn_vec, v);
|
|
213
|
+
mx_vec = _mm512_mask_max_ps(mx_vec, mask, mx_vec, v);
|
|
214
|
+
}
|
|
215
|
+
|
|
216
|
+
mn = _mm512_reduce_min_ps(mn_vec);
|
|
217
|
+
mx = _mm512_reduce_max_ps(mx_vec);
|
|
218
|
+
}
|
|
219
|
+
|
|
220
|
+
template <>
|
|
221
|
+
void lut_quantize_16_to_uint8<SIMDLevel::AVX512>(
|
|
222
|
+
const float* tab,
|
|
223
|
+
float mn,
|
|
224
|
+
float a,
|
|
225
|
+
uint8_t* out) {
|
|
226
|
+
const __m512 values = _mm512_loadu_ps(tab);
|
|
227
|
+
const __m512 a_vec = _mm512_set1_ps(a);
|
|
228
|
+
const __m512 scaled =
|
|
229
|
+
_mm512_fmsub_ps(values, a_vec, _mm512_set1_ps(mn * a));
|
|
230
|
+
__m512i rounded = round_nonnegative_ps_to_i32(scaled);
|
|
231
|
+
rounded = _mm512_max_epi32(rounded, _mm512_setzero_si512());
|
|
232
|
+
store_i32_as_u8_16(rounded, out);
|
|
233
|
+
}
|
|
234
|
+
|
|
235
|
+
template <>
|
|
236
|
+
void quantize_query_values<SIMDLevel::AVX512>(
|
|
237
|
+
const float* rq,
|
|
238
|
+
size_t d,
|
|
239
|
+
float v_min,
|
|
240
|
+
float inv_delta,
|
|
241
|
+
uint8_t max_code,
|
|
242
|
+
bool centered,
|
|
243
|
+
uint8_t* rqq,
|
|
244
|
+
size_t& sum_qq,
|
|
245
|
+
int64_t& sum2_signed_odd_int) {
|
|
246
|
+
const __m512 inv_delta_vec = _mm512_set1_ps(inv_delta);
|
|
247
|
+
const __m512 v_min_times_inv_delta_vec = _mm512_set1_ps(v_min * inv_delta);
|
|
248
|
+
const __m512 zero = _mm512_setzero_ps();
|
|
249
|
+
const __m512 max_code_ps = _mm512_set1_ps(static_cast<float>(max_code));
|
|
250
|
+
const __m512i max_code_i32 = _mm512_set1_epi32(max_code);
|
|
251
|
+
const __m512i two = _mm512_set1_epi32(2);
|
|
252
|
+
|
|
253
|
+
size_t i = 0;
|
|
254
|
+
if (centered) {
|
|
255
|
+
__m512i sum_acc_lo = _mm512_setzero_si512();
|
|
256
|
+
__m512i sum_acc_hi = _mm512_setzero_si512();
|
|
257
|
+
__m512i sq_acc_lo = _mm512_setzero_si512();
|
|
258
|
+
__m512i sq_acc_hi = _mm512_setzero_si512();
|
|
259
|
+
for (; i + 16 <= d; i += 16) {
|
|
260
|
+
const __m512 values = _mm512_loadu_ps(rq + i);
|
|
261
|
+
__m512 scaled = _mm512_fmsub_ps(
|
|
262
|
+
values, inv_delta_vec, v_min_times_inv_delta_vec);
|
|
263
|
+
scaled = _mm512_min_ps(_mm512_max_ps(scaled, zero), max_code_ps);
|
|
264
|
+
__m512i rounded = round_nonnegative_ps_to_i32(scaled);
|
|
265
|
+
|
|
266
|
+
sum_acc_lo = _mm512_add_epi64(
|
|
267
|
+
sum_acc_lo,
|
|
268
|
+
_mm512_cvtepi32_epi64(_mm512_castsi512_si256(rounded)));
|
|
269
|
+
sum_acc_hi = _mm512_add_epi64(
|
|
270
|
+
sum_acc_hi,
|
|
271
|
+
_mm512_cvtepi32_epi64(
|
|
272
|
+
_mm512_extracti64x4_epi64(rounded, 1)));
|
|
273
|
+
const __m512i signed_odd = _mm512_sub_epi32(
|
|
274
|
+
_mm512_mullo_epi32(rounded, two), max_code_i32);
|
|
275
|
+
const __m512i signed_odd_sqr =
|
|
276
|
+
_mm512_mullo_epi32(signed_odd, signed_odd);
|
|
277
|
+
sq_acc_lo = _mm512_add_epi64(
|
|
278
|
+
sq_acc_lo,
|
|
279
|
+
_mm512_cvtepi32_epi64(
|
|
280
|
+
_mm512_castsi512_si256(signed_odd_sqr)));
|
|
281
|
+
sq_acc_hi = _mm512_add_epi64(
|
|
282
|
+
sq_acc_hi,
|
|
283
|
+
_mm512_cvtepi32_epi64(
|
|
284
|
+
_mm512_extracti64x4_epi64(signed_odd_sqr, 1)));
|
|
285
|
+
store_i32_as_u8_16(rounded, rqq + i);
|
|
286
|
+
}
|
|
287
|
+
sum_qq += static_cast<uint64_t>(
|
|
288
|
+
_mm512_reduce_add_epi64(sum_acc_lo) +
|
|
289
|
+
_mm512_reduce_add_epi64(sum_acc_hi));
|
|
290
|
+
sum2_signed_odd_int += _mm512_reduce_add_epi64(sq_acc_lo);
|
|
291
|
+
sum2_signed_odd_int += _mm512_reduce_add_epi64(sq_acc_hi);
|
|
292
|
+
} else {
|
|
293
|
+
__m512i sum_acc_lo = _mm512_setzero_si512();
|
|
294
|
+
__m512i sum_acc_hi = _mm512_setzero_si512();
|
|
295
|
+
for (; i + 16 <= d; i += 16) {
|
|
296
|
+
const __m512 values = _mm512_loadu_ps(rq + i);
|
|
297
|
+
__m512 scaled = _mm512_fmsub_ps(
|
|
298
|
+
values, inv_delta_vec, v_min_times_inv_delta_vec);
|
|
299
|
+
scaled = _mm512_min_ps(_mm512_max_ps(scaled, zero), max_code_ps);
|
|
300
|
+
__m512i rounded = round_nonnegative_ps_to_i32(scaled);
|
|
301
|
+
|
|
302
|
+
sum_acc_lo = _mm512_add_epi64(
|
|
303
|
+
sum_acc_lo,
|
|
304
|
+
_mm512_cvtepi32_epi64(_mm512_castsi512_si256(rounded)));
|
|
305
|
+
sum_acc_hi = _mm512_add_epi64(
|
|
306
|
+
sum_acc_hi,
|
|
307
|
+
_mm512_cvtepi32_epi64(
|
|
308
|
+
_mm512_extracti64x4_epi64(rounded, 1)));
|
|
309
|
+
store_i32_as_u8_16(rounded, rqq + i);
|
|
310
|
+
}
|
|
311
|
+
sum_qq += static_cast<uint64_t>(
|
|
312
|
+
_mm512_reduce_add_epi64(sum_acc_lo) +
|
|
313
|
+
_mm512_reduce_add_epi64(sum_acc_hi));
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
for (; i < d; i++) {
|
|
317
|
+
const uint8_t v_qq = round_clamped_byte_scalar(
|
|
318
|
+
(rq[i] - v_min) * inv_delta, max_code);
|
|
319
|
+
rqq[i] = v_qq;
|
|
320
|
+
sum_qq += v_qq;
|
|
321
|
+
|
|
322
|
+
if (centered) {
|
|
323
|
+
const int64_t signed_odd_int = int64_t(v_qq) * 2 - max_code;
|
|
324
|
+
sum2_signed_odd_int += signed_odd_int * signed_odd_int;
|
|
325
|
+
}
|
|
326
|
+
}
|
|
327
|
+
}
|
|
328
|
+
|
|
172
329
|
template <>
|
|
173
330
|
uint64_t bitwise_and_dot_product<SIMDLevel::AVX512>(
|
|
174
331
|
const uint8_t* query,
|
|
@@ -237,6 +394,87 @@ uint64_t bitwise_and_dot_product<SIMDLevel::AVX512>(
|
|
|
237
394
|
return sum;
|
|
238
395
|
}
|
|
239
396
|
|
|
397
|
+
template <>
|
|
398
|
+
BitwiseAndDotProductResult bitwise_and_dot_product_with_popcount<
|
|
399
|
+
SIMDLevel::AVX512>(
|
|
400
|
+
const uint8_t* query,
|
|
401
|
+
const uint8_t* data,
|
|
402
|
+
size_t size,
|
|
403
|
+
size_t qb) {
|
|
404
|
+
uint64_t dot_product = 0;
|
|
405
|
+
uint64_t popcount_sum = 0;
|
|
406
|
+
size_t offset = 0;
|
|
407
|
+
if (size_t step = 512 / 8; offset + step <= size) {
|
|
408
|
+
__m512i dot_512 = _mm512_setzero_si512();
|
|
409
|
+
__m512i pop_512 = _mm512_setzero_si512();
|
|
410
|
+
for (; offset + step <= size; offset += step) {
|
|
411
|
+
__m512i v_x = _mm512_loadu_si512((const __m512i*)(data + offset));
|
|
412
|
+
pop_512 = _mm512_add_epi64(pop_512, popcount_512(v_x));
|
|
413
|
+
for (int j = 0; j < qb; j++) {
|
|
414
|
+
__m512i v_q = _mm512_loadu_si512(
|
|
415
|
+
(const __m512i*)(query + j * size + offset));
|
|
416
|
+
__m512i v_and = _mm512_and_si512(v_q, v_x);
|
|
417
|
+
__m512i v_popcnt = popcount_512(v_and);
|
|
418
|
+
__m512i v_shifted = _mm512_slli_epi64(v_popcnt, j);
|
|
419
|
+
dot_512 = _mm512_add_epi64(dot_512, v_shifted);
|
|
420
|
+
}
|
|
421
|
+
}
|
|
422
|
+
dot_product += _mm512_reduce_add_epi64(dot_512);
|
|
423
|
+
popcount_sum += _mm512_reduce_add_epi64(pop_512);
|
|
424
|
+
}
|
|
425
|
+
if (size_t step = 256 / 8; offset + step <= size) {
|
|
426
|
+
__m256i dot_256 = _mm256_setzero_si256();
|
|
427
|
+
__m256i pop_256 = _mm256_setzero_si256();
|
|
428
|
+
for (; offset + step <= size; offset += step) {
|
|
429
|
+
__m256i v_x = _mm256_loadu_si256((const __m256i*)(data + offset));
|
|
430
|
+
pop_256 = _mm256_add_epi64(pop_256, popcount_256(v_x));
|
|
431
|
+
for (int j = 0; j < qb; j++) {
|
|
432
|
+
__m256i v_q = _mm256_loadu_si256(
|
|
433
|
+
(const __m256i*)(query + j * size + offset));
|
|
434
|
+
__m256i v_and = _mm256_and_si256(v_q, v_x);
|
|
435
|
+
__m256i v_popcnt = popcount_256(v_and);
|
|
436
|
+
__m256i v_shifted = _mm256_slli_epi64(v_popcnt, j);
|
|
437
|
+
dot_256 = _mm256_add_epi64(dot_256, v_shifted);
|
|
438
|
+
}
|
|
439
|
+
}
|
|
440
|
+
dot_product += reduce_add_256(dot_256);
|
|
441
|
+
popcount_sum += reduce_add_256(pop_256);
|
|
442
|
+
}
|
|
443
|
+
__m128i dot_128 = _mm_setzero_si128();
|
|
444
|
+
__m128i pop_128 = _mm_setzero_si128();
|
|
445
|
+
for (size_t step = 128 / 8; offset + step <= size; offset += step) {
|
|
446
|
+
__m128i v_x = _mm_loadu_si128((const __m128i*)(data + offset));
|
|
447
|
+
pop_128 = _mm_add_epi64(pop_128, popcount_128(v_x));
|
|
448
|
+
for (int j = 0; j < qb; j++) {
|
|
449
|
+
__m128i v_q = _mm_loadu_si128(
|
|
450
|
+
(const __m128i*)(query + j * size + offset));
|
|
451
|
+
__m128i v_and = _mm_and_si128(v_q, v_x);
|
|
452
|
+
__m128i v_popcnt = popcount_128(v_and);
|
|
453
|
+
__m128i v_shifted = _mm_slli_epi64(v_popcnt, j);
|
|
454
|
+
dot_128 = _mm_add_epi64(dot_128, v_shifted);
|
|
455
|
+
}
|
|
456
|
+
}
|
|
457
|
+
dot_product += reduce_add_128(dot_128);
|
|
458
|
+
popcount_sum += reduce_add_128(pop_128);
|
|
459
|
+
for (size_t step = 64 / 8; offset + step <= size; offset += step) {
|
|
460
|
+
const auto yv = *(const uint64_t*)(data + offset);
|
|
461
|
+
popcount_sum += popcount64(yv);
|
|
462
|
+
for (int j = 0; j < qb; j++) {
|
|
463
|
+
const auto qv = *(const uint64_t*)(query + j * size + offset);
|
|
464
|
+
dot_product += popcount64(qv & yv) << j;
|
|
465
|
+
}
|
|
466
|
+
}
|
|
467
|
+
for (; offset < size; ++offset) {
|
|
468
|
+
const auto yv = *(data + offset);
|
|
469
|
+
popcount_sum += popcount32(yv);
|
|
470
|
+
for (int j = 0; j < qb; j++) {
|
|
471
|
+
const auto qv = *(query + j * size + offset);
|
|
472
|
+
dot_product += popcount32(qv & yv) << j;
|
|
473
|
+
}
|
|
474
|
+
}
|
|
475
|
+
return {dot_product, popcount_sum};
|
|
476
|
+
}
|
|
477
|
+
|
|
240
478
|
template <>
|
|
241
479
|
uint64_t bitwise_xor_dot_product<SIMDLevel::AVX512>(
|
|
242
480
|
const uint8_t* query,
|
|
@@ -344,6 +582,41 @@ uint64_t popcount<SIMDLevel::AVX512>(const uint8_t* data, size_t size) {
|
|
|
344
582
|
return sum;
|
|
345
583
|
}
|
|
346
584
|
|
|
585
|
+
template <>
|
|
586
|
+
void rearrange_bit_planes<SIMDLevel::AVX512>(
|
|
587
|
+
const uint8_t* rotated_qq,
|
|
588
|
+
size_t d,
|
|
589
|
+
size_t qb,
|
|
590
|
+
uint8_t* out) {
|
|
591
|
+
const size_t offset = (d + 7) / 8;
|
|
592
|
+
memset(out, 0, offset * qb);
|
|
593
|
+
size_t idim = 0;
|
|
594
|
+
for (; idim + 64 <= d; idim += 64) {
|
|
595
|
+
__m512i vals = _mm512_loadu_si512((const __m512i*)(rotated_qq + idim));
|
|
596
|
+
for (size_t iv = 0; iv < qb; iv++) {
|
|
597
|
+
__m512i mask = _mm512_set1_epi8(static_cast<char>(1 << iv));
|
|
598
|
+
__mmask64 bits = _mm512_test_epi8_mask(vals, mask);
|
|
599
|
+
memcpy(&out[iv * offset + idim / 8], &bits, 8);
|
|
600
|
+
}
|
|
601
|
+
}
|
|
602
|
+
for (; idim + 32 <= d; idim += 32) {
|
|
603
|
+
__m256i vals = _mm256_loadu_si256((const __m256i*)(rotated_qq + idim));
|
|
604
|
+
for (size_t iv = 0; iv < qb; iv++) {
|
|
605
|
+
__m256i mask = _mm256_set1_epi8(static_cast<char>(1 << iv));
|
|
606
|
+
__m256i bits =
|
|
607
|
+
_mm256_cmpeq_epi8(_mm256_and_si256(vals, mask), mask);
|
|
608
|
+
uint32_t packed = static_cast<uint32_t>(_mm256_movemask_epi8(bits));
|
|
609
|
+
memcpy(&out[iv * offset + idim / 8], &packed, 4);
|
|
610
|
+
}
|
|
611
|
+
}
|
|
612
|
+
for (; idim < d; idim++) {
|
|
613
|
+
for (size_t iv = 0; iv < qb; iv++) {
|
|
614
|
+
const bool bit = ((rotated_qq[idim] & (1 << iv)) != 0);
|
|
615
|
+
out[iv * offset + idim / 8] |= bit ? (1 << (idim % 8)) : 0;
|
|
616
|
+
}
|
|
617
|
+
}
|
|
618
|
+
}
|
|
619
|
+
|
|
347
620
|
} // namespace faiss::rabitq
|
|
348
621
|
|
|
349
622
|
namespace faiss::rabitq::multibit {
|
|
@@ -185,6 +185,98 @@ uint64_t bitwise_and_dot_product<SIMDLevel::AVX512_SPR>(
|
|
|
185
185
|
return sum;
|
|
186
186
|
}
|
|
187
187
|
|
|
188
|
+
template <>
|
|
189
|
+
BitwiseAndDotProductResult bitwise_and_dot_product_with_popcount<
|
|
190
|
+
SIMDLevel::AVX512_SPR>(
|
|
191
|
+
const uint8_t* query,
|
|
192
|
+
const uint8_t* data,
|
|
193
|
+
size_t size,
|
|
194
|
+
size_t qb) {
|
|
195
|
+
uint64_t dot_product = 0;
|
|
196
|
+
uint64_t popcount_sum = 0;
|
|
197
|
+
size_t offset = 0;
|
|
198
|
+
|
|
199
|
+
if (size_t step = 512 / 8; offset + step <= size) {
|
|
200
|
+
__m512i dot_512 = _mm512_setzero_si512();
|
|
201
|
+
__m512i pop_512 = _mm512_setzero_si512();
|
|
202
|
+
for (; offset + step <= size; offset += step) {
|
|
203
|
+
__m512i v_x = _mm512_loadu_si512(
|
|
204
|
+
reinterpret_cast<const __m512i*>(data + offset));
|
|
205
|
+
pop_512 = _mm512_add_epi64(pop_512, popcount_512_vpopcntdq(v_x));
|
|
206
|
+
for (size_t j = 0; j < qb; j++) {
|
|
207
|
+
__m512i v_q = _mm512_loadu_si512(
|
|
208
|
+
reinterpret_cast<const __m512i*>(
|
|
209
|
+
query + j * size + offset));
|
|
210
|
+
__m512i v_and = _mm512_and_si512(v_q, v_x);
|
|
211
|
+
__m512i v_popcnt = popcount_512_vpopcntdq(v_and);
|
|
212
|
+
__m512i v_shifted = _mm512_slli_epi64(v_popcnt, j);
|
|
213
|
+
dot_512 = _mm512_add_epi64(dot_512, v_shifted);
|
|
214
|
+
}
|
|
215
|
+
}
|
|
216
|
+
dot_product += _mm512_reduce_add_epi64(dot_512);
|
|
217
|
+
popcount_sum += _mm512_reduce_add_epi64(pop_512);
|
|
218
|
+
}
|
|
219
|
+
|
|
220
|
+
if (size_t step = 256 / 8; offset + step <= size) {
|
|
221
|
+
__m256i dot_256 = _mm256_setzero_si256();
|
|
222
|
+
__m256i pop_256 = _mm256_setzero_si256();
|
|
223
|
+
for (; offset + step <= size; offset += step) {
|
|
224
|
+
__m256i v_x = _mm256_loadu_si256(
|
|
225
|
+
reinterpret_cast<const __m256i*>(data + offset));
|
|
226
|
+
pop_256 = _mm256_add_epi64(pop_256, popcount_256_vpopcntdq(v_x));
|
|
227
|
+
for (size_t j = 0; j < qb; j++) {
|
|
228
|
+
__m256i v_q = _mm256_loadu_si256(
|
|
229
|
+
reinterpret_cast<const __m256i*>(
|
|
230
|
+
query + j * size + offset));
|
|
231
|
+
__m256i v_and = _mm256_and_si256(v_q, v_x);
|
|
232
|
+
__m256i v_popcnt = popcount_256_vpopcntdq(v_and);
|
|
233
|
+
__m256i v_shifted = _mm256_slli_epi64(v_popcnt, j);
|
|
234
|
+
dot_256 = _mm256_add_epi64(dot_256, v_shifted);
|
|
235
|
+
}
|
|
236
|
+
}
|
|
237
|
+
dot_product += reduce_add_256(dot_256);
|
|
238
|
+
popcount_sum += reduce_add_256(pop_256);
|
|
239
|
+
}
|
|
240
|
+
|
|
241
|
+
__m128i dot_128 = _mm_setzero_si128();
|
|
242
|
+
__m128i pop_128 = _mm_setzero_si128();
|
|
243
|
+
for (size_t step = 128 / 8; offset + step <= size; offset += step) {
|
|
244
|
+
__m128i v_x = _mm_loadu_si128(
|
|
245
|
+
reinterpret_cast<const __m128i*>(data + offset));
|
|
246
|
+
pop_128 = _mm_add_epi64(pop_128, popcount_128_vpopcntdq(v_x));
|
|
247
|
+
for (size_t j = 0; j < qb; j++) {
|
|
248
|
+
__m128i v_q = _mm_loadu_si128(
|
|
249
|
+
reinterpret_cast<const __m128i*>(
|
|
250
|
+
query + j * size + offset));
|
|
251
|
+
__m128i v_and = _mm_and_si128(v_q, v_x);
|
|
252
|
+
__m128i v_popcnt = popcount_128_vpopcntdq(v_and);
|
|
253
|
+
__m128i v_shifted = _mm_slli_epi64(v_popcnt, j);
|
|
254
|
+
dot_128 = _mm_add_epi64(dot_128, v_shifted);
|
|
255
|
+
}
|
|
256
|
+
}
|
|
257
|
+
dot_product += reduce_add_128(dot_128);
|
|
258
|
+
popcount_sum += reduce_add_128(pop_128);
|
|
259
|
+
|
|
260
|
+
for (size_t step = 64 / 8; offset + step <= size; offset += step) {
|
|
261
|
+
const auto yv = *reinterpret_cast<const uint64_t*>(data + offset);
|
|
262
|
+
popcount_sum += popcount64(yv);
|
|
263
|
+
for (size_t j = 0; j < qb; j++) {
|
|
264
|
+
const auto qv = *reinterpret_cast<const uint64_t*>(
|
|
265
|
+
query + j * size + offset);
|
|
266
|
+
dot_product += static_cast<uint64_t>(popcount64(qv & yv)) << j;
|
|
267
|
+
}
|
|
268
|
+
}
|
|
269
|
+
for (; offset < size; ++offset) {
|
|
270
|
+
const auto yv = *(data + offset);
|
|
271
|
+
popcount_sum += popcount32(yv);
|
|
272
|
+
for (size_t j = 0; j < qb; j++) {
|
|
273
|
+
const auto qv = *(query + j * size + offset);
|
|
274
|
+
dot_product += static_cast<uint64_t>(popcount32(qv & yv)) << j;
|
|
275
|
+
}
|
|
276
|
+
}
|
|
277
|
+
return {dot_product, popcount_sum};
|
|
278
|
+
}
|
|
279
|
+
|
|
188
280
|
template <>
|
|
189
281
|
uint64_t bitwise_xor_dot_product<SIMDLevel::AVX512_SPR>(
|
|
190
282
|
const uint8_t* query,
|
|
@@ -20,6 +20,17 @@ uint64_t bitwise_and_dot_product<SIMDLevel::ARM_NEON>(
|
|
|
20
20
|
return bitwise_and_dot_product<SIMDLevel::NONE>(query, data, size, qb);
|
|
21
21
|
}
|
|
22
22
|
|
|
23
|
+
template <>
|
|
24
|
+
BitwiseAndDotProductResult bitwise_and_dot_product_with_popcount<
|
|
25
|
+
SIMDLevel::ARM_NEON>(
|
|
26
|
+
const uint8_t* query,
|
|
27
|
+
const uint8_t* data,
|
|
28
|
+
size_t size,
|
|
29
|
+
size_t qb) {
|
|
30
|
+
return bitwise_and_dot_product_with_popcount<SIMDLevel::NONE>(
|
|
31
|
+
query, data, size, qb);
|
|
32
|
+
}
|
|
33
|
+
|
|
23
34
|
template <>
|
|
24
35
|
uint64_t bitwise_xor_dot_product<SIMDLevel::ARM_NEON>(
|
|
25
36
|
const uint8_t* query,
|