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.
Files changed (178) hide show
  1. checksums.yaml +4 -4
  2. data/CHANGELOG.md +8 -0
  3. data/lib/faiss/version.rb +1 -1
  4. data/vendor/faiss/faiss/AutoTune.cpp +3 -1
  5. data/vendor/faiss/faiss/Clustering.cpp +9 -1
  6. data/vendor/faiss/faiss/Clustering.h +8 -0
  7. data/vendor/faiss/faiss/IVFlib.cpp +14 -3
  8. data/vendor/faiss/faiss/Index.h +2 -2
  9. data/vendor/faiss/faiss/IndexAdditiveQuantizer.cpp +9 -10
  10. data/vendor/faiss/faiss/IndexAdditiveQuantizerFastScan.cpp +2 -3
  11. data/vendor/faiss/faiss/IndexBinaryFromFloat.cpp +24 -4
  12. data/vendor/faiss/faiss/IndexBinaryHNSW.cpp +16 -145
  13. data/vendor/faiss/faiss/IndexBinaryHNSW.h +0 -6
  14. data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
  15. data/vendor/faiss/faiss/IndexBinaryIVF.cpp +8 -18
  16. data/vendor/faiss/faiss/IndexBinaryIVF.h +8 -1
  17. data/vendor/faiss/faiss/IndexEDEN.cpp +273 -0
  18. data/vendor/faiss/faiss/IndexEDEN.h +57 -0
  19. data/vendor/faiss/faiss/IndexFastScan.cpp +15 -4
  20. data/vendor/faiss/faiss/IndexFlat.cpp +21 -54
  21. data/vendor/faiss/faiss/IndexFlat.h +2 -2
  22. data/vendor/faiss/faiss/IndexHNSW.cpp +311 -102
  23. data/vendor/faiss/faiss/IndexHNSW.h +31 -7
  24. data/vendor/faiss/faiss/IndexIDMap.cpp +26 -8
  25. data/vendor/faiss/faiss/IndexIDMap.h +2 -0
  26. data/vendor/faiss/faiss/IndexIVF.cpp +36 -10
  27. data/vendor/faiss/faiss/IndexIVFAdditiveQuantizer.cpp +1 -1
  28. data/vendor/faiss/faiss/IndexIVFAdditiveQuantizerFastScan.cpp +3 -4
  29. data/vendor/faiss/faiss/IndexIVFEDEN.cpp +302 -0
  30. data/vendor/faiss/faiss/IndexIVFEDEN.h +70 -0
  31. data/vendor/faiss/faiss/IndexIVFFastScan.cpp +5 -6
  32. data/vendor/faiss/faiss/IndexIVFFlat.cpp +3 -4
  33. data/vendor/faiss/faiss/IndexIVFIndependentQuantizer.cpp +1 -1
  34. data/vendor/faiss/faiss/IndexIVFPQ.cpp +49 -23
  35. data/vendor/faiss/faiss/IndexIVFPQ.h +11 -0
  36. data/vendor/faiss/faiss/IndexIVFPQFastScan.cpp +0 -1
  37. data/vendor/faiss/faiss/IndexIVFRaBitQ.cpp +19 -49
  38. data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.cpp +180 -76
  39. data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.h +5 -4
  40. data/vendor/faiss/faiss/IndexIVFSpectralHash.cpp +8 -6
  41. data/vendor/faiss/faiss/IndexLSH.cpp +2 -3
  42. data/vendor/faiss/faiss/IndexLattice.cpp +5 -0
  43. data/vendor/faiss/faiss/IndexNNDescent.cpp +9 -2
  44. data/vendor/faiss/faiss/IndexNSG.cpp +7 -2
  45. data/vendor/faiss/faiss/IndexPQ.cpp +6 -8
  46. data/vendor/faiss/faiss/IndexPreTransform.cpp +15 -0
  47. data/vendor/faiss/faiss/IndexRaBitQ.cpp +2 -2
  48. data/vendor/faiss/faiss/IndexRaBitQFastScan.cpp +1 -2
  49. data/vendor/faiss/faiss/IndexRaBitQFastScan.h +5 -1
  50. data/vendor/faiss/faiss/IndexRefine.cpp +30 -1
  51. data/vendor/faiss/faiss/IndexReplicas.cpp +1 -2
  52. data/vendor/faiss/faiss/IndexShards.cpp +5 -5
  53. data/vendor/faiss/faiss/IndexShardsIVF.cpp +6 -5
  54. data/vendor/faiss/faiss/MetaIndexes.cpp +2 -4
  55. data/vendor/faiss/faiss/SuperKMeans.cpp +286 -247
  56. data/vendor/faiss/faiss/SuperKMeans.h +33 -2
  57. data/vendor/faiss/faiss/VectorTransform.cpp +71 -2
  58. data/vendor/faiss/faiss/VectorTransform.h +3 -0
  59. data/vendor/faiss/faiss/clone_index.cpp +8 -0
  60. data/vendor/faiss/faiss/factory_tools.cpp +47 -4
  61. data/vendor/faiss/faiss/gpu/GpuCloner.cpp +11 -11
  62. data/vendor/faiss/faiss/gpu/GpuClonerOptions.h +1 -5
  63. data/vendor/faiss/faiss/gpu/GpuDistance.h +2 -5
  64. data/vendor/faiss/faiss/gpu/GpuIndex.h +38 -16
  65. data/vendor/faiss/faiss/gpu/GpuIndexCagra.h +71 -1
  66. data/vendor/faiss/faiss/gpu/GpuIndexIVF.h +17 -0
  67. data/vendor/faiss/faiss/gpu/GpuIndexIVFScalarQuantizer.h +16 -0
  68. data/vendor/faiss/faiss/gpu/perf/PerfClustering.cpp +1 -1
  69. data/vendor/faiss/faiss/gpu/perf/PerfIVFPQAdd.cpp +2 -2
  70. data/vendor/faiss/faiss/gpu/test/TestGpuIndexIVFScalarQuantizer.cpp +180 -0
  71. data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFFlat.h +1 -5
  72. data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFPQ.h +88 -0
  73. data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFPQ.h +134 -0
  74. data/vendor/faiss/faiss/impl/AdditiveQuantizer.cpp +1 -1
  75. data/vendor/faiss/faiss/impl/ClusteringInitialization.cpp +7 -4
  76. data/vendor/faiss/faiss/impl/DistanceComputer.h +34 -0
  77. data/vendor/faiss/faiss/impl/EDENQuantizer.h +119 -0
  78. data/vendor/faiss/faiss/impl/HNSW.cpp +528 -267
  79. data/vendor/faiss/faiss/impl/HNSW.h +46 -7
  80. data/vendor/faiss/faiss/impl/IDSelector.h +44 -0
  81. data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
  82. data/vendor/faiss/faiss/impl/NNDescent.cpp +10 -3
  83. data/vendor/faiss/faiss/impl/NSG.cpp +3 -1
  84. data/vendor/faiss/faiss/impl/Panorama.h +20 -9
  85. data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
  86. data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +38 -26
  87. data/vendor/faiss/faiss/impl/RaBitQUtils.cpp +45 -37
  88. data/vendor/faiss/faiss/impl/RaBitQUtils.h +35 -0
  89. data/vendor/faiss/faiss/impl/RaBitQuantizer.cpp +239 -72
  90. data/vendor/faiss/faiss/impl/RaBitQuantizer.h +66 -4
  91. data/vendor/faiss/faiss/impl/RaBitQuantizerMultiBit.cpp +4 -13
  92. data/vendor/faiss/faiss/impl/ResultHandler.h +34 -34
  93. data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +287 -84
  94. data/vendor/faiss/faiss/impl/ScalarQuantizer.h +26 -10
  95. data/vendor/faiss/faiss/impl/ThreadedIndex-inl.h +2 -2
  96. data/vendor/faiss/faiss/impl/VisitedTable.cpp +22 -2
  97. data/vendor/faiss/faiss/impl/VisitedTable.h +20 -0
  98. data/vendor/faiss/faiss/impl/binary_hamming/IndexBinaryIVF_impl.h +90 -14
  99. data/vendor/faiss/faiss/impl/binary_hamming/avx2.cpp +4 -4
  100. data/vendor/faiss/faiss/impl/expanded_scanners.h +5 -1
  101. data/vendor/faiss/faiss/impl/fast_scan/decompose_qbs.h +1 -0
  102. data/vendor/faiss/faiss/impl/fast_scan/dispatching.h +35 -2
  103. data/vendor/faiss/faiss/impl/hnsw/LockVector.cpp +1 -1
  104. data/vendor/faiss/faiss/impl/index_read.cpp +491 -50
  105. data/vendor/faiss/faiss/impl/index_write.cpp +86 -30
  106. data/vendor/faiss/faiss/impl/lattice_Zn.cpp +8 -9
  107. data/vendor/faiss/faiss/impl/platform_macros.h +3 -1
  108. data/vendor/faiss/faiss/impl/polysemous_training/avx512.cpp +284 -0
  109. data/vendor/faiss/faiss/impl/polysemous_training/dispatch.h +115 -0
  110. data/vendor/faiss/faiss/impl/pq_code_distance/IVFPQScanner_impl.h +73 -39
  111. data/vendor/faiss/faiss/impl/pq_code_distance/IVFPQ_QueryTables.cpp +0 -1
  112. data/vendor/faiss/faiss/impl/pq_code_distance/PQDistanceComputer_impl.h +26 -15
  113. data/vendor/faiss/faiss/impl/pq_code_distance/avx2.cpp +4 -4
  114. data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-generic.cpp +4 -4
  115. data/vendor/faiss/faiss/impl/result_handler/ResultHandler.cpp +195 -0
  116. data/vendor/faiss/faiss/impl/result_handler/avx2.cpp +133 -0
  117. data/vendor/faiss/faiss/impl/result_handler/avx512.cpp +281 -0
  118. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx2.cpp +72 -0
  119. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx512.cpp +228 -0
  120. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer.cpp +887 -0
  121. data/vendor/faiss/faiss/impl/scalar_quantizer/distance_computers.h +2 -2
  122. data/vendor/faiss/faiss/impl/scalar_quantizer/quantizers.h +9 -8
  123. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +90 -24
  124. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-impl.h +30 -30
  125. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-spr.cpp +4 -5
  126. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512.cpp +101 -34
  127. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-dispatch.h +169 -13
  128. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +125 -26
  129. data/vendor/faiss/faiss/impl/simd_dispatch.h +70 -31
  130. data/vendor/faiss/faiss/index_factory.cpp +40 -7
  131. data/vendor/faiss/faiss/invlists/DirectMap.cpp +1 -1
  132. data/vendor/faiss/faiss/invlists/InvertedLists.cpp +9 -6
  133. data/vendor/faiss/faiss/invlists/OnDiskInvertedLists.cpp +29 -8
  134. data/vendor/faiss/faiss/python/python_callbacks.cpp +3 -1
  135. data/vendor/faiss/faiss/svs/IndexSVSFaissUtils.h +60 -0
  136. data/vendor/faiss/faiss/svs/IndexSVSFlat.cpp +26 -1
  137. data/vendor/faiss/faiss/svs/IndexSVSFlat.h +13 -0
  138. data/vendor/faiss/faiss/svs/IndexSVSIVF.cpp +1 -1
  139. data/vendor/faiss/faiss/svs/IndexSVSIVFLeanVec.cpp +1 -1
  140. data/vendor/faiss/faiss/svs/IndexSVSVamana.cpp +47 -5
  141. data/vendor/faiss/faiss/svs/IndexSVSVamana.h +23 -3
  142. data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.cpp +4 -2
  143. data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.h +2 -1
  144. data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +10 -4
  145. data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.h +2 -1
  146. data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
  147. data/vendor/faiss/faiss/utils/distances.cpp +30 -11
  148. data/vendor/faiss/faiss/utils/distances_dispatch.h +30 -24
  149. data/vendor/faiss/faiss/utils/distances_fused/distances_fused.cpp +1 -1
  150. data/vendor/faiss/faiss/utils/distances_simd.cpp +4 -3
  151. data/vendor/faiss/faiss/utils/extra_distances.cpp +4 -14
  152. data/vendor/faiss/faiss/utils/extra_distances.h +1 -2
  153. data/vendor/faiss/faiss/utils/hamming.cpp +16 -10
  154. data/vendor/faiss/faiss/utils/hamming.h +10 -1
  155. data/vendor/faiss/faiss/utils/hamming_distance/common.h +14 -3
  156. data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_vpopcnt.cpp +24 -0
  157. data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512.h +1 -1
  158. data/vendor/faiss/faiss/utils/hamming_distance/{hamming_computer-avx512_spr.h → hamming_computer-avx512_vpopcnt.h} +85 -24
  159. data/vendor/faiss/faiss/utils/hamming_distance/hamming_impl.h +141 -0
  160. data/vendor/faiss/faiss/utils/quantize_lut.cpp +29 -8
  161. data/vendor/faiss/faiss/utils/rabitq_simd.h +202 -0
  162. data/vendor/faiss/faiss/utils/simd_impl/distances_arm_sve.cpp +194 -30
  163. data/vendor/faiss/faiss/utils/simd_impl/distances_avx2.cpp +0 -1
  164. data/vendor/faiss/faiss/utils/simd_impl/distances_avx512.cpp +263 -15
  165. data/vendor/faiss/faiss/utils/simd_impl/distances_rvv.cpp +198 -18
  166. data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx2.cpp +245 -0
  167. data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512.cpp +330 -40
  168. data/vendor/faiss/faiss/utils/simd_impl/{rabitq_avx512_spr.cpp → rabitq_avx512_vpopcnt.cpp} +112 -23
  169. data/vendor/faiss/faiss/utils/simd_impl/rabitq_neon.cpp +11 -0
  170. data/vendor/faiss/faiss/utils/simd_impl/rabitq_rvv.cpp +143 -6
  171. data/vendor/faiss/faiss/utils/simd_impl/super_kmeans_dispatch.h +2 -7
  172. data/vendor/faiss/faiss/utils/simd_impl/super_kmeans_kernels.h +6 -1
  173. data/vendor/faiss/faiss/utils/simd_impl/super_kmeans_kernels_sve.cpp +34 -0
  174. data/vendor/faiss/faiss/utils/simd_levels.cpp +196 -47
  175. data/vendor/faiss/faiss/utils/simd_levels.h +33 -8
  176. data/vendor/faiss/faiss/utils/utils.cpp +9 -27
  177. metadata +21 -5
  178. 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
- // AVX2+BMI2 bitplane kernel used as fallback for ex_bits >= 2.
395
- // AVX512 TU has AVX2 available. BMI2 guarded separately since
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
- inline float ip_bitplane_avx2(
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
- __m256 acc = _mm256_setzero_ps();
406
- const __m256 v_one = _mm256_set1_ps(1.0f);
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
- __m256 v_weights[8];
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] = _mm256_set1_ps(static_cast<float>(1u << b));
682
+ v_weights[b] = _mm512_set1_ps(static_cast<float>(1u << b));
420
683
  }
421
- v_weights[ex_bits] = _mm256_set1_ps(static_cast<float>(1u << 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 + 8 <= d; i += 8) {
425
- __m256i sb_cmp = _mm256_cmpgt_epi32(
426
- _mm256_and_si256(_mm256_set1_epi32(sign_bits[i / 8]), bit_pos),
427
- zero);
428
- __m256 recon = _mm256_mul_ps(
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 ex64 = 0;
433
- memcpy(&ex64, ex_code + (i / 8) * ex_bits, sizeof(uint64_t));
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
- auto plane = static_cast<uint8_t>(_pext_u64(ex64, pext_masks[b]));
437
- __m256i p_cmp = _mm256_cmpgt_epi32(
438
- _mm256_and_si256(_mm256_set1_epi32(plane), bit_pos), zero);
439
- __m256 p_f = _mm256_and_ps(_mm256_castsi256_ps(p_cmp), v_one);
440
- recon = _mm256_fmadd_ps(p_f, v_weights[b], 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
- __m256 rq = _mm256_loadu_ps(rotated_q + i);
444
- acc = _mm256_fmadd_ps(rq, _mm256_add_ps(recon, v_cb), 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
- float result = hsum_avx2(acc);
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 ip_bitplane_avx2(sign_bits, ex_code, rotated_q, d, ex_bits, cb);
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 rabitq_avx512_spr.cpp
9
+ * @file rabitq_avx512_vpopcnt.cpp
10
10
  *
11
- * RaBitQ SIMD kernels specialized for SIMDLevel::AVX512_SPR.
11
+ * RaBitQ SIMD kernels specialized for SIMDLevel::AVX512_VPOPCNT.
12
12
  *
13
- * Sapphire Rapids (SPR) and later Intel microarchitectures expose
14
- * AVX-512 VPOPCNTDQ (vpopcntq), which performs a per-lane 64-bit
15
- * popcount in a single instruction. This is used here to replace the
16
- * multi-step shuffle/pshufb-based popcount used by the generic AVX-512
17
- * specialization in rabitq_avx512.cpp. The popcount-heavy kernels
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 (statically built for SPR+): compiled. The
25
- * SINGLE_SIMD_LEVEL is AVX512_SPR, so this specialization is
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
- * SIMDConfig::level == SIMDLevel::AVX512_SPR.
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<SPR> to
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 COMPILE_SIMD_AVX512_SPR
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 on SPR builds, so
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
- // AVX512VL is part of the SPR feature set, so vpopcntq is available
79
- // on 256-bit registers via _mm256_popcnt_epi64.
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::AVX512_SPR>(
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
- uint64_t bitwise_xor_dot_product<SIMDLevel::AVX512_SPR>(
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::AVX512_SPR>(const uint8_t* data, size_t size) {
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::AVX512_SPR>(
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 // COMPILE_SIMD_AVX512_SPR
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,