faiss 0.6.2 → 0.6.4
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/Clustering.h +8 -0
- 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 +24 -4
- data/vendor/faiss/faiss/IndexBinaryHNSW.cpp +16 -145
- data/vendor/faiss/faiss/IndexBinaryHNSW.h +0 -6
- data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
- data/vendor/faiss/faiss/IndexBinaryIVF.cpp +8 -18
- data/vendor/faiss/faiss/IndexBinaryIVF.h +8 -1
- 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 +21 -54
- data/vendor/faiss/faiss/IndexFlat.h +2 -2
- data/vendor/faiss/faiss/IndexHNSW.cpp +311 -102
- data/vendor/faiss/faiss/IndexHNSW.h +31 -7
- data/vendor/faiss/faiss/IndexIDMap.cpp +26 -8
- data/vendor/faiss/faiss/IndexIDMap.h +2 -0
- data/vendor/faiss/faiss/IndexIVF.cpp +36 -10
- 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 +49 -23
- data/vendor/faiss/faiss/IndexIVFPQ.h +11 -0
- 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 +5 -5
- data/vendor/faiss/faiss/IndexShardsIVF.cpp +6 -5
- data/vendor/faiss/faiss/MetaIndexes.cpp +2 -4
- data/vendor/faiss/faiss/SuperKMeans.cpp +286 -247
- data/vendor/faiss/faiss/SuperKMeans.h +33 -2
- data/vendor/faiss/faiss/VectorTransform.cpp +71 -2
- data/vendor/faiss/faiss/VectorTransform.h +3 -0
- data/vendor/faiss/faiss/clone_index.cpp +8 -0
- data/vendor/faiss/faiss/factory_tools.cpp +47 -4
- data/vendor/faiss/faiss/gpu/GpuCloner.cpp +11 -11
- data/vendor/faiss/faiss/gpu/GpuClonerOptions.h +1 -5
- data/vendor/faiss/faiss/gpu/GpuDistance.h +2 -5
- data/vendor/faiss/faiss/gpu/GpuIndex.h +38 -16
- data/vendor/faiss/faiss/gpu/GpuIndexCagra.h +71 -1
- 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/AdditiveQuantizer.cpp +1 -1
- data/vendor/faiss/faiss/impl/ClusteringInitialization.cpp +7 -4
- 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 +528 -267
- data/vendor/faiss/faiss/impl/HNSW.h +46 -7
- data/vendor/faiss/faiss/impl/IDSelector.h +44 -0
- data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
- data/vendor/faiss/faiss/impl/NNDescent.cpp +10 -3
- data/vendor/faiss/faiss/impl/NSG.cpp +3 -1
- data/vendor/faiss/faiss/impl/Panorama.h +20 -9
- data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
- data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +38 -26
- 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 +239 -72
- data/vendor/faiss/faiss/impl/RaBitQuantizer.h +66 -4
- data/vendor/faiss/faiss/impl/RaBitQuantizerMultiBit.cpp +4 -13
- data/vendor/faiss/faiss/impl/ResultHandler.h +34 -34
- data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +287 -84
- data/vendor/faiss/faiss/impl/ScalarQuantizer.h +26 -10
- 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 +90 -14
- data/vendor/faiss/faiss/impl/binary_hamming/avx2.cpp +4 -4
- data/vendor/faiss/faiss/impl/expanded_scanners.h +5 -1
- data/vendor/faiss/faiss/impl/fast_scan/decompose_qbs.h +1 -0
- 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 +491 -50
- data/vendor/faiss/faiss/impl/index_write.cpp +86 -30
- 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/IVFPQScanner_impl.h +73 -39
- 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/pq_code_distance/pq_code_distance-generic.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 +887 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/distance_computers.h +2 -2
- data/vendor/faiss/faiss/impl/scalar_quantizer/quantizers.h +9 -8
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +90 -24
- 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 +169 -13
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +125 -26
- data/vendor/faiss/faiss/impl/simd_dispatch.h +70 -31
- data/vendor/faiss/faiss/index_factory.cpp +40 -7
- data/vendor/faiss/faiss/invlists/DirectMap.cpp +1 -1
- data/vendor/faiss/faiss/invlists/InvertedLists.cpp +9 -6
- data/vendor/faiss/faiss/invlists/OnDiskInvertedLists.cpp +29 -8
- 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 +47 -5
- data/vendor/faiss/faiss/svs/IndexSVSVamana.h +23 -3
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.cpp +4 -2
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.h +2 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +10 -4
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.h +2 -1
- data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
- data/vendor/faiss/faiss/utils/distances.cpp +30 -11
- data/vendor/faiss/faiss/utils/distances_dispatch.h +30 -24
- data/vendor/faiss/faiss/utils/distances_fused/distances_fused.cpp +1 -1
- 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 +16 -10
- data/vendor/faiss/faiss/utils/hamming.h +10 -1
- data/vendor/faiss/faiss/utils/hamming_distance/common.h +14 -3
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_vpopcnt.cpp +24 -0
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512.h +1 -1
- data/vendor/faiss/faiss/utils/hamming_distance/{hamming_computer-avx512_spr.h → hamming_computer-avx512_vpopcnt.h} +85 -24
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_impl.h +141 -0
- 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_arm_sve.cpp +194 -30
- 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 +198 -18
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx2.cpp +245 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512.cpp +330 -40
- data/vendor/faiss/faiss/utils/simd_impl/{rabitq_avx512_spr.cpp → rabitq_avx512_vpopcnt.cpp} +112 -23
- 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_impl/super_kmeans_dispatch.h +2 -7
- data/vendor/faiss/faiss/utils/simd_impl/super_kmeans_kernels.h +6 -1
- data/vendor/faiss/faiss/utils/simd_impl/super_kmeans_kernels_sve.cpp +34 -0
- data/vendor/faiss/faiss/utils/simd_levels.cpp +196 -47
- data/vendor/faiss/faiss/utils/simd_levels.h +33 -8
- data/vendor/faiss/faiss/utils/utils.cpp +9 -27
- metadata +21 -5
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_spr.cpp +0 -15
|
@@ -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,22 +582,47 @@ 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 {
|
|
350
623
|
|
|
351
624
|
namespace {
|
|
352
625
|
|
|
353
|
-
inline float hsum_avx2(__m256 v) {
|
|
354
|
-
__m128 hi = _mm256_extractf128_ps(v, 1);
|
|
355
|
-
__m128 lo = _mm256_castps256_ps128(v);
|
|
356
|
-
lo = _mm_add_ps(lo, hi);
|
|
357
|
-
__m128 shuf = _mm_movehdup_ps(lo);
|
|
358
|
-
lo = _mm_add_ps(lo, shuf);
|
|
359
|
-
shuf = _mm_movehl_ps(shuf, lo);
|
|
360
|
-
return _mm_cvtss_f32(_mm_add_ss(lo, shuf));
|
|
361
|
-
}
|
|
362
|
-
|
|
363
626
|
inline float ip_1exbit_avx512(
|
|
364
627
|
const uint8_t* __restrict sign_bits,
|
|
365
628
|
const uint8_t* __restrict ex_code,
|
|
@@ -391,60 +654,86 @@ inline float ip_1exbit_avx512(
|
|
|
391
654
|
return result;
|
|
392
655
|
}
|
|
393
656
|
|
|
394
|
-
//
|
|
395
|
-
//
|
|
396
|
-
// VIA Eden X4 has AVX2 without BMI2.
|
|
657
|
+
// Needs BMI2 for _pext_u64. Some AVX2 CPUs lack it, and FAISS_BMI2_FLAGS can
|
|
658
|
+
// be empty, so the dispatcher falls back to the scalar path without it.
|
|
397
659
|
#ifdef __BMI2__
|
|
398
|
-
|
|
660
|
+
// Bitplane kernel for ex_bits >= 2, 16 dims per iteration. A bitplane is
|
|
661
|
+
// already a bitmask, so it goes into a mask register and one masked add
|
|
662
|
+
// applies its weight. Reads of ex_code run a few bytes past the ex-code
|
|
663
|
+
// section into the record's own trailing factors, so they stay in bounds.
|
|
664
|
+
inline float ip_bitplane_avx512(
|
|
399
665
|
const uint8_t* __restrict sign_bits,
|
|
400
666
|
const uint8_t* __restrict ex_code,
|
|
401
667
|
const float* __restrict rotated_q,
|
|
402
668
|
size_t d,
|
|
403
669
|
size_t ex_bits,
|
|
404
670
|
float cb) {
|
|
405
|
-
|
|
406
|
-
const
|
|
407
|
-
const __m256i bit_pos = _mm256_setr_epi32(1, 2, 4, 8, 16, 32, 64, 128);
|
|
408
|
-
const __m256i zero = _mm256_setzero_si256();
|
|
409
|
-
const __m256 v_cb = _mm256_set1_ps(cb);
|
|
671
|
+
__m512 acc = _mm512_setzero_ps();
|
|
672
|
+
const __m512 v_cb = _mm512_set1_ps(cb);
|
|
410
673
|
|
|
411
674
|
uint64_t pext_masks[7];
|
|
412
|
-
|
|
675
|
+
__m512 v_weights[8];
|
|
413
676
|
for (size_t b = 0; b < ex_bits; b++) {
|
|
414
677
|
uint64_t m = 0;
|
|
415
678
|
for (int j = 0; j < 8; j++) {
|
|
416
679
|
m |= (1ULL << (b + j * ex_bits));
|
|
417
680
|
}
|
|
418
681
|
pext_masks[b] = m;
|
|
419
|
-
v_weights[b] =
|
|
682
|
+
v_weights[b] = _mm512_set1_ps(static_cast<float>(1u << b));
|
|
420
683
|
}
|
|
421
|
-
v_weights[ex_bits] =
|
|
684
|
+
v_weights[ex_bits] = _mm512_set1_ps(static_cast<float>(1u << ex_bits));
|
|
422
685
|
|
|
423
686
|
size_t i = 0;
|
|
424
|
-
for (; i +
|
|
425
|
-
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
429
|
-
_mm256_and_ps(_mm256_castsi256_ps(sb_cmp), v_one),
|
|
430
|
-
v_weights[ex_bits]);
|
|
687
|
+
for (; i + 16 <= d; i += 16) {
|
|
688
|
+
uint16_t sb = 0;
|
|
689
|
+
memcpy(&sb, sign_bits + (i / 8), sizeof(uint16_t));
|
|
690
|
+
__m512 recon = _mm512_maskz_mov_ps(
|
|
691
|
+
static_cast<__mmask16>(sb), v_weights[ex_bits]);
|
|
431
692
|
|
|
432
|
-
uint64_t
|
|
433
|
-
|
|
693
|
+
uint64_t lo64 = 0;
|
|
694
|
+
uint64_t hi64 = 0;
|
|
695
|
+
memcpy(&lo64, ex_code + (i / 8) * ex_bits, sizeof(uint64_t));
|
|
696
|
+
memcpy(&hi64, ex_code + ((i / 8) + 1) * ex_bits, sizeof(uint64_t));
|
|
434
697
|
|
|
435
698
|
for (size_t b = 0; b < ex_bits; b++) {
|
|
436
|
-
|
|
437
|
-
|
|
438
|
-
|
|
439
|
-
|
|
440
|
-
recon =
|
|
699
|
+
const uint32_t plane =
|
|
700
|
+
static_cast<uint32_t>(_pext_u64(lo64, pext_masks[b])) |
|
|
701
|
+
(static_cast<uint32_t>(_pext_u64(hi64, pext_masks[b]))
|
|
702
|
+
<< 8);
|
|
703
|
+
recon = _mm512_mask_add_ps(
|
|
704
|
+
recon, static_cast<__mmask16>(plane), recon, v_weights[b]);
|
|
441
705
|
}
|
|
442
706
|
|
|
443
|
-
|
|
444
|
-
acc =
|
|
707
|
+
__m512 rq = _mm512_loadu_ps(rotated_q + i);
|
|
708
|
+
acc = _mm512_fmadd_ps(rq, _mm512_add_ps(recon, v_cb), acc);
|
|
445
709
|
}
|
|
446
710
|
|
|
447
|
-
|
|
711
|
+
// Half-width step: keeps the scalar tail under 8 dims when d is a multiple
|
|
712
|
+
// of 8 but not of 16 (e.g. 200, 1000). The upper 8 lanes are masked off
|
|
713
|
+
// throughout, and rotated_q is loaded masked so nothing is read past the
|
|
714
|
+
// end.
|
|
715
|
+
if (i + 8 <= d) {
|
|
716
|
+
const __mmask16 low8 = static_cast<__mmask16>(0x00ff);
|
|
717
|
+
__m512 recon = _mm512_maskz_mov_ps(
|
|
718
|
+
static_cast<__mmask16>(sign_bits[i / 8]), v_weights[ex_bits]);
|
|
719
|
+
|
|
720
|
+
uint64_t lo64 = 0;
|
|
721
|
+
memcpy(&lo64, ex_code + (i / 8) * ex_bits, sizeof(uint64_t));
|
|
722
|
+
|
|
723
|
+
for (size_t b = 0; b < ex_bits; b++) {
|
|
724
|
+
const uint32_t plane =
|
|
725
|
+
static_cast<uint32_t>(_pext_u64(lo64, pext_masks[b]));
|
|
726
|
+
recon = _mm512_mask_add_ps(
|
|
727
|
+
recon, static_cast<__mmask16>(plane), recon, v_weights[b]);
|
|
728
|
+
}
|
|
729
|
+
|
|
730
|
+
__m512 rq = _mm512_maskz_loadu_ps(low8, rotated_q + i);
|
|
731
|
+
acc = _mm512_fmadd_ps(
|
|
732
|
+
rq, _mm512_mask_add_ps(recon, low8, recon, v_cb), acc);
|
|
733
|
+
i += 8;
|
|
734
|
+
}
|
|
735
|
+
|
|
736
|
+
float result = _mm512_reduce_add_ps(acc);
|
|
448
737
|
result += ip_scalar(sign_bits, ex_code, rotated_q, i, d, ex_bits, cb);
|
|
449
738
|
return result;
|
|
450
739
|
}
|
|
@@ -466,7 +755,8 @@ float compute_inner_product<SIMDLevel::AVX512>(
|
|
|
466
755
|
|
|
467
756
|
#ifdef __BMI2__
|
|
468
757
|
if (ex_bits <= 7) {
|
|
469
|
-
return
|
|
758
|
+
return ip_bitplane_avx512(
|
|
759
|
+
sign_bits, ex_code, rotated_q, d, ex_bits, cb);
|
|
470
760
|
}
|
|
471
761
|
#endif
|
|
472
762
|
return ip_scalar(sign_bits, ex_code, rotated_q, 0, d, ex_bits, cb);
|
|
@@ -6,34 +6,31 @@
|
|
|
6
6
|
*/
|
|
7
7
|
|
|
8
8
|
/**
|
|
9
|
-
* @file
|
|
9
|
+
* @file rabitq_avx512_vpopcnt.cpp
|
|
10
10
|
*
|
|
11
|
-
* RaBitQ SIMD kernels specialized for SIMDLevel::
|
|
11
|
+
* RaBitQ SIMD kernels specialized for SIMDLevel::AVX512_VPOPCNT.
|
|
12
12
|
*
|
|
13
|
-
*
|
|
14
|
-
*
|
|
15
|
-
*
|
|
16
|
-
* multi-step shuffle
|
|
17
|
-
* specialization in rabitq_avx512.cpp.
|
|
18
|
-
* (bitwise_and_dot_product, bitwise_xor_dot_product, popcount) become
|
|
19
|
-
* substantially shorter and faster on SPR+ as a result.
|
|
13
|
+
* AVX-512 VPOPCNTDQ performs a per-lane 64-bit popcount in a single
|
|
14
|
+
* instruction. It is available on CPUs including Ice Lake, Zen 4, and
|
|
15
|
+
* Sapphire Rapids, independently of the other SPR-only extensions. This
|
|
16
|
+
* replaces the multi-step shuffle-based popcount used by the generic
|
|
17
|
+
* AVX-512 specialization in rabitq_avx512.cpp.
|
|
20
18
|
*
|
|
21
19
|
* Build / dispatch behavior:
|
|
22
20
|
* - faiss_avx512 (AVX-512 only, no SPR features): NOT compiled.
|
|
23
21
|
* The existing AVX512 specialization in rabitq_avx512.cpp is used.
|
|
24
|
-
* - faiss_avx512_spr
|
|
25
|
-
*
|
|
26
|
-
* selected by static dispatch.
|
|
22
|
+
* - faiss_avx512_spr: compiled alongside the full SPR specialization and
|
|
23
|
+
* selected through the SPR -> VPOPCNT fallback.
|
|
27
24
|
* - faiss with FAISS_OPT_LEVEL=dd (dynamic dispatch): compiled with
|
|
28
25
|
* -mavx512vpopcntdq as a per-file flag. Selected at runtime when
|
|
29
|
-
*
|
|
26
|
+
* the CPU exposes AVX512_VPOPCNTDQ.
|
|
30
27
|
*
|
|
31
28
|
* The floating-point multi-bit inner-product kernel does not benefit
|
|
32
|
-
* from VPOPCNTDQ, so this TU forwards compute_inner_product<
|
|
29
|
+
* from VPOPCNTDQ, so this TU forwards compute_inner_product<VPOPCNT> to
|
|
33
30
|
* the AVX512 implementation to avoid duplicating that code path.
|
|
34
31
|
*/
|
|
35
32
|
|
|
36
|
-
#ifdef
|
|
33
|
+
#ifdef COMPILE_SIMD_AVX512_VPOPCNT
|
|
37
34
|
|
|
38
35
|
#include <faiss/utils/popcount.h>
|
|
39
36
|
#include <faiss/utils/rabitq_simd.h>
|
|
@@ -47,7 +44,7 @@
|
|
|
47
44
|
namespace faiss::rabitq {
|
|
48
45
|
|
|
49
46
|
// Forward declarations for the AVX512 specializations defined in
|
|
50
|
-
// rabitq_avx512.cpp. They live in the same TU group
|
|
47
|
+
// rabitq_avx512.cpp. They live in the same TU group in supported builds, so
|
|
51
48
|
// we can reuse them as a tail handler / fallback. Declaring rather
|
|
52
49
|
// than redefining avoids ODR risk and keeps a single source of truth
|
|
53
50
|
// for the floating-point kernel.
|
|
@@ -75,8 +72,8 @@ inline __m512i popcount_512_vpopcntdq(__m512i v) {
|
|
|
75
72
|
}
|
|
76
73
|
|
|
77
74
|
// 256-bit popcount using AVX-512VL VPOPCNTDQ.
|
|
78
|
-
//
|
|
79
|
-
//
|
|
75
|
+
// Baseline AVX-512 includes AVX512VL, so VPOPCNTDQ is also available on
|
|
76
|
+
// 256-bit registers via _mm256_popcnt_epi64.
|
|
80
77
|
inline __m256i popcount_256_vpopcntdq(__m256i v) {
|
|
81
78
|
return _mm256_popcnt_epi64(v);
|
|
82
79
|
}
|
|
@@ -101,7 +98,7 @@ inline uint64_t reduce_add_128(__m128i v) {
|
|
|
101
98
|
} // namespace
|
|
102
99
|
|
|
103
100
|
template <>
|
|
104
|
-
uint64_t bitwise_and_dot_product<SIMDLevel::
|
|
101
|
+
uint64_t bitwise_and_dot_product<SIMDLevel::AVX512_VPOPCNT>(
|
|
105
102
|
const uint8_t* query,
|
|
106
103
|
const uint8_t* data,
|
|
107
104
|
size_t size,
|
|
@@ -186,7 +183,99 @@ uint64_t bitwise_and_dot_product<SIMDLevel::AVX512_SPR>(
|
|
|
186
183
|
}
|
|
187
184
|
|
|
188
185
|
template <>
|
|
189
|
-
|
|
186
|
+
BitwiseAndDotProductResult bitwise_and_dot_product_with_popcount<
|
|
187
|
+
SIMDLevel::AVX512_VPOPCNT>(
|
|
188
|
+
const uint8_t* query,
|
|
189
|
+
const uint8_t* data,
|
|
190
|
+
size_t size,
|
|
191
|
+
size_t qb) {
|
|
192
|
+
uint64_t dot_product = 0;
|
|
193
|
+
uint64_t popcount_sum = 0;
|
|
194
|
+
size_t offset = 0;
|
|
195
|
+
|
|
196
|
+
if (size_t step = 512 / 8; offset + step <= size) {
|
|
197
|
+
__m512i dot_512 = _mm512_setzero_si512();
|
|
198
|
+
__m512i pop_512 = _mm512_setzero_si512();
|
|
199
|
+
for (; offset + step <= size; offset += step) {
|
|
200
|
+
__m512i v_x = _mm512_loadu_si512(
|
|
201
|
+
reinterpret_cast<const __m512i*>(data + offset));
|
|
202
|
+
pop_512 = _mm512_add_epi64(pop_512, popcount_512_vpopcntdq(v_x));
|
|
203
|
+
for (size_t j = 0; j < qb; j++) {
|
|
204
|
+
__m512i v_q = _mm512_loadu_si512(
|
|
205
|
+
reinterpret_cast<const __m512i*>(
|
|
206
|
+
query + j * size + offset));
|
|
207
|
+
__m512i v_and = _mm512_and_si512(v_q, v_x);
|
|
208
|
+
__m512i v_popcnt = popcount_512_vpopcntdq(v_and);
|
|
209
|
+
__m512i v_shifted = _mm512_slli_epi64(v_popcnt, j);
|
|
210
|
+
dot_512 = _mm512_add_epi64(dot_512, v_shifted);
|
|
211
|
+
}
|
|
212
|
+
}
|
|
213
|
+
dot_product += _mm512_reduce_add_epi64(dot_512);
|
|
214
|
+
popcount_sum += _mm512_reduce_add_epi64(pop_512);
|
|
215
|
+
}
|
|
216
|
+
|
|
217
|
+
if (size_t step = 256 / 8; offset + step <= size) {
|
|
218
|
+
__m256i dot_256 = _mm256_setzero_si256();
|
|
219
|
+
__m256i pop_256 = _mm256_setzero_si256();
|
|
220
|
+
for (; offset + step <= size; offset += step) {
|
|
221
|
+
__m256i v_x = _mm256_loadu_si256(
|
|
222
|
+
reinterpret_cast<const __m256i*>(data + offset));
|
|
223
|
+
pop_256 = _mm256_add_epi64(pop_256, popcount_256_vpopcntdq(v_x));
|
|
224
|
+
for (size_t j = 0; j < qb; j++) {
|
|
225
|
+
__m256i v_q = _mm256_loadu_si256(
|
|
226
|
+
reinterpret_cast<const __m256i*>(
|
|
227
|
+
query + j * size + offset));
|
|
228
|
+
__m256i v_and = _mm256_and_si256(v_q, v_x);
|
|
229
|
+
__m256i v_popcnt = popcount_256_vpopcntdq(v_and);
|
|
230
|
+
__m256i v_shifted = _mm256_slli_epi64(v_popcnt, j);
|
|
231
|
+
dot_256 = _mm256_add_epi64(dot_256, v_shifted);
|
|
232
|
+
}
|
|
233
|
+
}
|
|
234
|
+
dot_product += reduce_add_256(dot_256);
|
|
235
|
+
popcount_sum += reduce_add_256(pop_256);
|
|
236
|
+
}
|
|
237
|
+
|
|
238
|
+
__m128i dot_128 = _mm_setzero_si128();
|
|
239
|
+
__m128i pop_128 = _mm_setzero_si128();
|
|
240
|
+
for (size_t step = 128 / 8; offset + step <= size; offset += step) {
|
|
241
|
+
__m128i v_x = _mm_loadu_si128(
|
|
242
|
+
reinterpret_cast<const __m128i*>(data + offset));
|
|
243
|
+
pop_128 = _mm_add_epi64(pop_128, popcount_128_vpopcntdq(v_x));
|
|
244
|
+
for (size_t j = 0; j < qb; j++) {
|
|
245
|
+
__m128i v_q = _mm_loadu_si128(
|
|
246
|
+
reinterpret_cast<const __m128i*>(
|
|
247
|
+
query + j * size + offset));
|
|
248
|
+
__m128i v_and = _mm_and_si128(v_q, v_x);
|
|
249
|
+
__m128i v_popcnt = popcount_128_vpopcntdq(v_and);
|
|
250
|
+
__m128i v_shifted = _mm_slli_epi64(v_popcnt, j);
|
|
251
|
+
dot_128 = _mm_add_epi64(dot_128, v_shifted);
|
|
252
|
+
}
|
|
253
|
+
}
|
|
254
|
+
dot_product += reduce_add_128(dot_128);
|
|
255
|
+
popcount_sum += reduce_add_128(pop_128);
|
|
256
|
+
|
|
257
|
+
for (size_t step = 64 / 8; offset + step <= size; offset += step) {
|
|
258
|
+
const auto yv = *reinterpret_cast<const uint64_t*>(data + offset);
|
|
259
|
+
popcount_sum += popcount64(yv);
|
|
260
|
+
for (size_t j = 0; j < qb; j++) {
|
|
261
|
+
const auto qv = *reinterpret_cast<const uint64_t*>(
|
|
262
|
+
query + j * size + offset);
|
|
263
|
+
dot_product += static_cast<uint64_t>(popcount64(qv & yv)) << j;
|
|
264
|
+
}
|
|
265
|
+
}
|
|
266
|
+
for (; offset < size; ++offset) {
|
|
267
|
+
const auto yv = *(data + offset);
|
|
268
|
+
popcount_sum += popcount32(yv);
|
|
269
|
+
for (size_t j = 0; j < qb; j++) {
|
|
270
|
+
const auto qv = *(query + j * size + offset);
|
|
271
|
+
dot_product += static_cast<uint64_t>(popcount32(qv & yv)) << j;
|
|
272
|
+
}
|
|
273
|
+
}
|
|
274
|
+
return {dot_product, popcount_sum};
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
template <>
|
|
278
|
+
uint64_t bitwise_xor_dot_product<SIMDLevel::AVX512_VPOPCNT>(
|
|
190
279
|
const uint8_t* query,
|
|
191
280
|
const uint8_t* data,
|
|
192
281
|
size_t size,
|
|
@@ -265,7 +354,7 @@ uint64_t bitwise_xor_dot_product<SIMDLevel::AVX512_SPR>(
|
|
|
265
354
|
}
|
|
266
355
|
|
|
267
356
|
template <>
|
|
268
|
-
uint64_t popcount<SIMDLevel::
|
|
357
|
+
uint64_t popcount<SIMDLevel::AVX512_VPOPCNT>(const uint8_t* data, size_t size) {
|
|
269
358
|
uint64_t sum = 0;
|
|
270
359
|
size_t offset = 0;
|
|
271
360
|
|
|
@@ -327,7 +416,7 @@ float compute_inner_product<SIMDLevel::AVX512>(
|
|
|
327
416
|
float cb);
|
|
328
417
|
|
|
329
418
|
template <>
|
|
330
|
-
float compute_inner_product<SIMDLevel::
|
|
419
|
+
float compute_inner_product<SIMDLevel::AVX512_VPOPCNT>(
|
|
331
420
|
const uint8_t* __restrict sign_bits,
|
|
332
421
|
const uint8_t* __restrict ex_code,
|
|
333
422
|
const float* __restrict rotated_q,
|
|
@@ -340,4 +429,4 @@ float compute_inner_product<SIMDLevel::AVX512_SPR>(
|
|
|
340
429
|
|
|
341
430
|
} // namespace faiss::rabitq::multibit
|
|
342
431
|
|
|
343
|
-
#endif //
|
|
432
|
+
#endif // COMPILE_SIMD_AVX512_VPOPCNT
|
|
@@ -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,
|