faiss 0.6.1 → 0.6.3

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (186) 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/IVFlib.cpp +14 -3
  7. data/vendor/faiss/faiss/Index.h +2 -2
  8. data/vendor/faiss/faiss/IndexAdditiveQuantizer.cpp +9 -10
  9. data/vendor/faiss/faiss/IndexAdditiveQuantizerFastScan.cpp +2 -3
  10. data/vendor/faiss/faiss/IndexBinaryFromFloat.cpp +1 -2
  11. data/vendor/faiss/faiss/IndexBinaryHNSW.cpp +10 -12
  12. data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
  13. data/vendor/faiss/faiss/IndexBinaryIVF.cpp +5 -7
  14. data/vendor/faiss/faiss/IndexEDEN.cpp +273 -0
  15. data/vendor/faiss/faiss/IndexEDEN.h +57 -0
  16. data/vendor/faiss/faiss/IndexFastScan.cpp +15 -4
  17. data/vendor/faiss/faiss/IndexFlat.cpp +13 -50
  18. data/vendor/faiss/faiss/IndexHNSW.cpp +177 -148
  19. data/vendor/faiss/faiss/IndexIDMap.cpp +16 -3
  20. data/vendor/faiss/faiss/IndexIDMap.h +2 -0
  21. data/vendor/faiss/faiss/IndexIVF.cpp +19 -8
  22. data/vendor/faiss/faiss/IndexIVFAdditiveQuantizer.cpp +3 -3
  23. data/vendor/faiss/faiss/IndexIVFAdditiveQuantizerFastScan.cpp +3 -4
  24. data/vendor/faiss/faiss/IndexIVFEDEN.cpp +302 -0
  25. data/vendor/faiss/faiss/IndexIVFEDEN.h +70 -0
  26. data/vendor/faiss/faiss/IndexIVFFastScan.cpp +5 -6
  27. data/vendor/faiss/faiss/IndexIVFFlat.cpp +6 -5
  28. data/vendor/faiss/faiss/IndexIVFFlatPanorama.cpp +3 -3
  29. data/vendor/faiss/faiss/IndexIVFIndependentQuantizer.cpp +1 -1
  30. data/vendor/faiss/faiss/IndexIVFPQ.cpp +42 -25
  31. data/vendor/faiss/faiss/IndexIVFPQFastScan.cpp +0 -1
  32. data/vendor/faiss/faiss/IndexIVFPQR.cpp +2 -3
  33. data/vendor/faiss/faiss/IndexIVFRaBitQ.cpp +23 -62
  34. data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.cpp +180 -76
  35. data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.h +5 -4
  36. data/vendor/faiss/faiss/IndexIVFSpectralHash.cpp +8 -6
  37. data/vendor/faiss/faiss/IndexLSH.cpp +2 -3
  38. data/vendor/faiss/faiss/IndexLattice.cpp +5 -0
  39. data/vendor/faiss/faiss/IndexNNDescent.cpp +10 -3
  40. data/vendor/faiss/faiss/IndexNSG.cpp +8 -4
  41. data/vendor/faiss/faiss/IndexPQ.cpp +6 -8
  42. data/vendor/faiss/faiss/IndexPreTransform.cpp +15 -0
  43. data/vendor/faiss/faiss/IndexRaBitQ.cpp +2 -2
  44. data/vendor/faiss/faiss/IndexRaBitQFastScan.cpp +1 -2
  45. data/vendor/faiss/faiss/IndexRaBitQFastScan.h +5 -1
  46. data/vendor/faiss/faiss/IndexRefine.cpp +30 -1
  47. data/vendor/faiss/faiss/IndexReplicas.cpp +1 -2
  48. data/vendor/faiss/faiss/IndexScalarQuantizer.cpp +68 -6
  49. data/vendor/faiss/faiss/IndexScalarQuantizer.h +10 -0
  50. data/vendor/faiss/faiss/IndexShards.cpp +2 -2
  51. data/vendor/faiss/faiss/IndexShardsIVF.cpp +2 -2
  52. data/vendor/faiss/faiss/MetaIndexes.cpp +2 -4
  53. data/vendor/faiss/faiss/SuperKMeans.cpp +256 -240
  54. data/vendor/faiss/faiss/SuperKMeans.h +30 -0
  55. data/vendor/faiss/faiss/VectorTransform.cpp +33 -2
  56. data/vendor/faiss/faiss/clone_index.cpp +5 -0
  57. data/vendor/faiss/faiss/cppcontrib/SaDecodeKernels.h +1 -1
  58. data/vendor/faiss/faiss/cppcontrib/sa_decode/Level2-neon-inl.h +902 -12
  59. data/vendor/faiss/faiss/cppcontrib/sa_decode/PQ-neon-inl.h +702 -10
  60. data/vendor/faiss/faiss/factory_tools.cpp +51 -4
  61. data/vendor/faiss/faiss/gpu/GpuCloner.cpp +11 -11
  62. data/vendor/faiss/faiss/gpu/GpuIndex.h +34 -11
  63. data/vendor/faiss/faiss/gpu/GpuIndexCagra.h +47 -0
  64. data/vendor/faiss/faiss/gpu/GpuIndexIVF.h +17 -0
  65. data/vendor/faiss/faiss/gpu/GpuIndexIVFScalarQuantizer.h +16 -0
  66. data/vendor/faiss/faiss/gpu/GpuResources.h +3 -2
  67. data/vendor/faiss/faiss/gpu/StandardGpuResources.cpp +11 -12
  68. data/vendor/faiss/faiss/gpu/StandardGpuResources.h +3 -3
  69. data/vendor/faiss/faiss/gpu/perf/PerfClustering.cpp +1 -1
  70. data/vendor/faiss/faiss/gpu/perf/PerfIVFPQAdd.cpp +2 -2
  71. data/vendor/faiss/faiss/gpu/test/TestGpuIndexIVFScalarQuantizer.cpp +180 -0
  72. data/vendor/faiss/faiss/gpu_metal/MetalDistance.h +87 -0
  73. data/vendor/faiss/faiss/gpu_metal/MetalIndex.h +7 -0
  74. data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFFlat.h +177 -0
  75. data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFPQ.h +88 -0
  76. data/vendor/faiss/faiss/gpu_metal/MetalKernels.h +48 -3
  77. data/vendor/faiss/faiss/gpu_metal/MetalPythonBridge.h +45 -0
  78. data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFFlat.h +193 -0
  79. data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFPQ.h +134 -0
  80. data/vendor/faiss/faiss/impl/ClusteringInitialization.cpp +2 -2
  81. data/vendor/faiss/faiss/impl/DistanceComputer.h +34 -0
  82. data/vendor/faiss/faiss/impl/EDENQuantizer.h +119 -0
  83. data/vendor/faiss/faiss/impl/HNSW.cpp +658 -344
  84. data/vendor/faiss/faiss/impl/HNSW.h +51 -13
  85. data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
  86. data/vendor/faiss/faiss/impl/NSG.cpp +18 -12
  87. data/vendor/faiss/faiss/impl/Panorama.h +20 -7
  88. data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
  89. data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +59 -24
  90. data/vendor/faiss/faiss/impl/RaBitQUtils.cpp +45 -37
  91. data/vendor/faiss/faiss/impl/RaBitQUtils.h +35 -0
  92. data/vendor/faiss/faiss/impl/RaBitQuantizer.cpp +175 -68
  93. data/vendor/faiss/faiss/impl/RaBitQuantizer.h +19 -0
  94. data/vendor/faiss/faiss/impl/RaBitQuantizerMultiBit.cpp +2 -11
  95. data/vendor/faiss/faiss/impl/ResultHandler.h +26 -31
  96. data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +522 -58
  97. data/vendor/faiss/faiss/impl/ScalarQuantizer.h +70 -0
  98. data/vendor/faiss/faiss/impl/ThreadedIndex-inl.h +2 -2
  99. data/vendor/faiss/faiss/impl/VisitedTable.cpp +33 -13
  100. data/vendor/faiss/faiss/impl/VisitedTable.h +88 -33
  101. data/vendor/faiss/faiss/impl/binary_hamming/IndexBinaryIVF_impl.h +1 -1
  102. data/vendor/faiss/faiss/impl/binary_hamming/avx2.cpp +4 -4
  103. data/vendor/faiss/faiss/impl/fast_scan/dispatching.h +38 -3
  104. data/vendor/faiss/faiss/impl/hnsw/LockVector.cpp +1 -1
  105. data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.cpp +35 -43
  106. data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.h +64 -15
  107. data/vendor/faiss/faiss/impl/hnsw/avx2.cpp +86 -40
  108. data/vendor/faiss/faiss/impl/hnsw/avx512.cpp +81 -50
  109. data/vendor/faiss/faiss/impl/index_read.cpp +476 -75
  110. data/vendor/faiss/faiss/impl/index_write.cpp +56 -4
  111. data/vendor/faiss/faiss/impl/io_macros.h +25 -0
  112. data/vendor/faiss/faiss/impl/lattice_Zn.cpp +8 -9
  113. data/vendor/faiss/faiss/impl/platform_macros.h +15 -9
  114. data/vendor/faiss/faiss/impl/polysemous_training/avx512.cpp +284 -0
  115. data/vendor/faiss/faiss/impl/polysemous_training/dispatch.h +115 -0
  116. data/vendor/faiss/faiss/impl/pq_code_distance/IVFPQ_QueryTables.cpp +0 -1
  117. data/vendor/faiss/faiss/impl/pq_code_distance/PQDistanceComputer_impl.h +26 -15
  118. data/vendor/faiss/faiss/impl/pq_code_distance/avx2.cpp +6 -4
  119. data/vendor/faiss/faiss/impl/pq_code_distance/avx512.cpp +2 -0
  120. data/vendor/faiss/faiss/impl/pq_code_distance/neon.cpp +2 -0
  121. data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-generic.cpp +20 -0
  122. data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-inl.h +36 -0
  123. data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-sve.cpp +5 -0
  124. data/vendor/faiss/faiss/impl/pq_code_distance/pq_scan_impl.h +105 -0
  125. data/vendor/faiss/faiss/impl/pq_code_distance/rvv.cpp +2 -0
  126. data/vendor/faiss/faiss/impl/result_handler/ResultHandler.cpp +195 -0
  127. data/vendor/faiss/faiss/impl/result_handler/avx2.cpp +133 -0
  128. data/vendor/faiss/faiss/impl/result_handler/avx512.cpp +281 -0
  129. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx2.cpp +72 -0
  130. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx512.cpp +228 -0
  131. data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer.cpp +882 -0
  132. data/vendor/faiss/faiss/impl/scalar_quantizer/distance_computers.h +6 -0
  133. data/vendor/faiss/faiss/impl/scalar_quantizer/quantizers.h +336 -26
  134. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +331 -32
  135. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-impl.h +553 -0
  136. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-spr.cpp +558 -0
  137. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512.cpp +284 -45
  138. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-dispatch.h +502 -3
  139. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +157 -32
  140. data/vendor/faiss/faiss/impl/scalar_quantizer/sq-rvv.cpp +26 -0
  141. data/vendor/faiss/faiss/impl/simd_dispatch.h +86 -8
  142. data/vendor/faiss/faiss/index_factory.cpp +37 -7
  143. data/vendor/faiss/faiss/index_io.h +16 -0
  144. data/vendor/faiss/faiss/invlists/DirectMap.cpp +5 -2
  145. data/vendor/faiss/faiss/invlists/InvertedLists.cpp +15 -15
  146. data/vendor/faiss/faiss/invlists/InvertedLists.h +2 -2
  147. data/vendor/faiss/faiss/invlists/OnDiskInvertedLists.cpp +19 -4
  148. data/vendor/faiss/faiss/python/python_callbacks.cpp +3 -1
  149. data/vendor/faiss/faiss/svs/IndexSVSFaissUtils.h +60 -0
  150. data/vendor/faiss/faiss/svs/IndexSVSFlat.cpp +26 -1
  151. data/vendor/faiss/faiss/svs/IndexSVSFlat.h +13 -0
  152. data/vendor/faiss/faiss/svs/IndexSVSIVF.cpp +1 -1
  153. data/vendor/faiss/faiss/svs/IndexSVSIVFLeanVec.cpp +1 -1
  154. data/vendor/faiss/faiss/svs/IndexSVSVamana.cpp +150 -23
  155. data/vendor/faiss/faiss/svs/IndexSVSVamana.h +30 -7
  156. data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.cpp +3 -2
  157. data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.h +2 -1
  158. data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +65 -25
  159. data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.h +3 -2
  160. data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
  161. data/vendor/faiss/faiss/utils/bf16.h +34 -0
  162. data/vendor/faiss/faiss/utils/distances.cpp +14 -2
  163. data/vendor/faiss/faiss/utils/distances_simd.cpp +4 -4
  164. data/vendor/faiss/faiss/utils/extra_distances.cpp +4 -14
  165. data/vendor/faiss/faiss/utils/extra_distances.h +1 -2
  166. data/vendor/faiss/faiss/utils/hamming.cpp +9 -9
  167. data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx2.cpp +2 -1
  168. data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_spr.cpp +15 -0
  169. data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512.h +6 -30
  170. data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512_spr.h +171 -0
  171. data/vendor/faiss/faiss/utils/partitioning.cpp +0 -2
  172. data/vendor/faiss/faiss/utils/quantize_lut.cpp +29 -8
  173. data/vendor/faiss/faiss/utils/rabitq_simd.h +202 -0
  174. data/vendor/faiss/faiss/utils/simd_impl/distances_avx2.cpp +0 -1
  175. data/vendor/faiss/faiss/utils/simd_impl/distances_avx512.cpp +263 -15
  176. data/vendor/faiss/faiss/utils/simd_impl/distances_rvv.cpp +160 -18
  177. data/vendor/faiss/faiss/utils/simd_impl/partitioning_simdlib256.h +14 -68
  178. data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx2.cpp +245 -0
  179. data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512.cpp +273 -0
  180. data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512_spr.cpp +435 -0
  181. data/vendor/faiss/faiss/utils/simd_impl/rabitq_neon.cpp +11 -0
  182. data/vendor/faiss/faiss/utils/simd_impl/rabitq_rvv.cpp +143 -6
  183. data/vendor/faiss/faiss/utils/simd_levels.cpp +56 -2
  184. data/vendor/faiss/faiss/utils/simd_levels.h +14 -0
  185. data/vendor/faiss/faiss/utils/utils.cpp +9 -27
  186. metadata +27 -2
@@ -16,89 +16,135 @@
16
16
 
17
17
  namespace faiss {
18
18
 
19
- template <>
20
- int MinimaxHeap::pop_min_tpl<SIMDLevel::AVX2>(float* vmin_out) {
21
- assert(k > 0);
19
+ namespace {
20
+
21
+ /// Templated AVX2 implementation of "pop best" for both CMax (returns
22
+ /// the smallest distance) and CMin (returns the largest similarity).
23
+ /// The only differences between the two flavors are: (1) the initial
24
+ /// "worst possible" value, (2) the running-best update comparison
25
+ /// (`_CMP_LT_OS` vs `_CMP_GT_OS`), and (3) the tiebreaker direction.
26
+ template <class HC>
27
+ int pop_best_avx2(MinimaxHeapT<HC>& heap, float* vmin_out) {
28
+ using storage_idx_t = typename MinimaxHeapT<HC>::storage_idx_t;
22
29
  static_assert(
23
30
  std::is_same<storage_idx_t, int32_t>::value,
24
31
  "This code expects storage_idx_t to be int32_t");
32
+ assert(heap.k > 0);
33
+
34
+ // For CMax (distance) the "best" candidate is the smallest value, so
35
+ // we initialize the running best to +inf. For CMin (similarity) the
36
+ // best is the largest value, so we initialize to -inf.
37
+ constexpr float worst_v = HC::is_max
38
+ ? std::numeric_limits<float>::infinity()
39
+ : -std::numeric_limits<float>::infinity();
25
40
 
26
- int32_t min_idx = -1;
27
- float min_dis = std::numeric_limits<float>::infinity();
41
+ int32_t best_idx = -1;
42
+ float best_dis = worst_v;
28
43
 
29
44
  size_t iii = 0;
30
45
 
31
- __m256i min_indices = _mm256_setr_epi32(-1, -1, -1, -1, -1, -1, -1, -1);
32
- __m256 min_distances =
33
- _mm256_set1_ps(std::numeric_limits<float>::infinity());
46
+ __m256i best_indices = _mm256_setr_epi32(-1, -1, -1, -1, -1, -1, -1, -1);
47
+ __m256 best_distances = _mm256_set1_ps(worst_v);
34
48
  __m256i current_indices = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
35
49
  __m256i offset = _mm256_set1_epi32(8);
36
50
 
37
- // The baseline version is available in the NONE specialization.
38
-
39
- // The following loop tracks the rightmost index with the min distance.
40
- // -1 index values are ignored.
41
- const size_t k8 = (k / 8) * 8;
51
+ // Track the rightmost index whose distance equals the running best.
52
+ // -1 index values are filtered out via m1mask.
53
+ const size_t k8 = (heap.k / 8) * 8;
42
54
  for (; iii < k8; iii += 8) {
43
55
  __m256i indices =
44
- _mm256_loadu_si256((const __m256i*)(ids.data() + iii));
45
- __m256 distances = _mm256_loadu_ps(dis.data() + iii);
56
+ _mm256_loadu_si256((const __m256i*)(heap.ids.data() + iii));
57
+ __m256 distances = _mm256_loadu_ps(heap.dis.data() + iii);
46
58
 
47
- // This mask filters out -1 values among indices.
59
+ // Mask out -1 indices (invalid entries).
48
60
  __m256i m1mask = _mm256_cmpgt_epi32(_mm256_setzero_si256(), indices);
49
61
 
50
- __m256i dmask = _mm256_castps_si256(
51
- _mm256_cmp_ps(min_distances, distances, _CMP_LT_OS));
62
+ // dmask is "true where best is already (strictly) better than the
63
+ // candidate" — entries the candidate should NOT update. For CMax,
64
+ // best < candidate means we keep best (we want the smallest);
65
+ // for CMin we keep best when best > candidate (we want the largest).
66
+ __m256i dmask;
67
+ if constexpr (HC::is_max) {
68
+ dmask = _mm256_castps_si256(
69
+ _mm256_cmp_ps(best_distances, distances, _CMP_LT_OS));
70
+ } else {
71
+ dmask = _mm256_castps_si256(
72
+ _mm256_cmp_ps(best_distances, distances, _CMP_GT_OS));
73
+ }
52
74
  __m256 finalmask = _mm256_castsi256_ps(_mm256_or_si256(m1mask, dmask));
53
75
 
54
- const __m256i min_indices_new = _mm256_castps_si256(_mm256_blendv_ps(
76
+ const __m256i best_indices_new = _mm256_castps_si256(_mm256_blendv_ps(
55
77
  _mm256_castsi256_ps(current_indices),
56
- _mm256_castsi256_ps(min_indices),
78
+ _mm256_castsi256_ps(best_indices),
57
79
  finalmask));
58
80
 
59
- const __m256 min_distances_new =
60
- _mm256_blendv_ps(distances, min_distances, finalmask);
81
+ const __m256 best_distances_new =
82
+ _mm256_blendv_ps(distances, best_distances, finalmask);
61
83
 
62
- min_indices = min_indices_new;
63
- min_distances = min_distances_new;
84
+ best_indices = best_indices_new;
85
+ best_distances = best_distances_new;
64
86
 
65
87
  current_indices = _mm256_add_epi32(current_indices, offset);
66
88
  }
67
89
 
68
- // Vectorizing is doable, but is not practical
90
+ // Vectorizing the horizontal reduction is doable but not practical.
69
91
  int32_t vidx8[8];
70
92
  float vdis8[8];
71
- _mm256_storeu_ps(vdis8, min_distances);
72
- _mm256_storeu_si256((__m256i*)vidx8, min_indices);
93
+ _mm256_storeu_ps(vdis8, best_distances);
94
+ _mm256_storeu_si256((__m256i*)vidx8, best_indices);
73
95
 
74
96
  for (size_t j = 0; j < 8; j++) {
75
- if (min_dis > vdis8[j] || (min_dis == vdis8[j] && min_idx < vidx8[j])) {
76
- min_idx = vidx8[j];
77
- min_dis = vdis8[j];
97
+ const bool strictly_better =
98
+ HC::is_max ? (best_dis > vdis8[j]) : (best_dis < vdis8[j]);
99
+ if (strictly_better || (best_dis == vdis8[j] && best_idx < vidx8[j])) {
100
+ best_idx = vidx8[j];
101
+ best_dis = vdis8[j];
78
102
  }
79
103
  }
80
104
 
81
- // process last values. Vectorizing is doable, but is not practical
82
- for (; iii < static_cast<size_t>(k); iii++) {
83
- if (ids[iii] != -1 && dis[iii] <= min_dis) {
84
- min_dis = dis[iii];
85
- min_idx = iii;
105
+ // Tail (under 8 entries). Vectorizing is doable but not practical.
106
+ for (; iii < static_cast<size_t>(heap.k); iii++) {
107
+ if (heap.ids[iii] == -1) {
108
+ continue;
109
+ }
110
+ const bool weakly_better = HC::is_max ? (best_dis >= heap.dis[iii])
111
+ : (best_dis <= heap.dis[iii]);
112
+ if (weakly_better) {
113
+ best_dis = heap.dis[iii];
114
+ best_idx = iii;
86
115
  }
87
116
  }
88
117
 
89
- if (min_idx == -1) {
118
+ if (best_idx == -1) {
90
119
  return -1;
91
120
  }
92
121
 
93
122
  if (vmin_out) {
94
- *vmin_out = min_dis;
123
+ *vmin_out = best_dis;
95
124
  }
96
- int ret = ids[min_idx];
97
- ids[min_idx] = -1;
98
- --nvalid;
125
+ int ret = heap.ids[best_idx];
126
+ heap.ids[best_idx] = -1;
127
+ --heap.nvalid;
99
128
  return ret;
100
129
  }
101
130
 
131
+ } // namespace
132
+
133
+ // Explicit specializations for AVX2
134
+ template <>
135
+ int pop_min_tpl<CMax<float, int32_t>, SIMDLevel::AVX2>(
136
+ MinimaxHeapT<CMax<float, int32_t>>* heap,
137
+ float* vmin_out) {
138
+ return pop_best_avx2<CMax<float, int32_t>>(*heap, vmin_out);
139
+ }
140
+
141
+ template <>
142
+ int pop_min_tpl<CMin<float, int32_t>, SIMDLevel::AVX2>(
143
+ MinimaxHeapT<CMin<float, int32_t>>* heap,
144
+ float* vmin_out) {
145
+ return pop_best_avx2<CMin<float, int32_t>>(*heap, vmin_out);
146
+ }
147
+
102
148
  } // namespace faiss
103
149
 
104
150
  #endif // COMPILE_SIMD_AVX2
@@ -16,96 +16,127 @@
16
16
 
17
17
  namespace faiss {
18
18
 
19
- template <>
20
- int MinimaxHeap::pop_min_tpl<SIMDLevel::AVX512>(float* vmin_out) {
21
- assert(k > 0);
19
+ namespace {
20
+
21
+ /// Templated AVX512 implementation of "pop best" for both CMax (returns
22
+ /// the smallest distance) and CMin (returns the largest similarity).
23
+ template <class HC>
24
+ int pop_best_avx512(MinimaxHeapT<HC>& heap, float* vmin_out) {
25
+ using storage_idx_t = typename MinimaxHeapT<HC>::storage_idx_t;
22
26
  static_assert(
23
27
  std::is_same<storage_idx_t, int32_t>::value,
24
28
  "This code expects storage_idx_t to be int32_t");
29
+ assert(heap.k > 0);
25
30
 
26
- int32_t min_idx = -1;
27
- float min_dis = std::numeric_limits<float>::infinity();
31
+ constexpr float worst_v = HC::is_max
32
+ ? std::numeric_limits<float>::infinity()
33
+ : -std::numeric_limits<float>::infinity();
28
34
 
29
- __m512i min_indices = _mm512_set1_epi32(-1);
30
- __m512 min_distances =
31
- _mm512_set1_ps(std::numeric_limits<float>::infinity());
35
+ int32_t best_idx = -1;
36
+ float best_dis = worst_v;
37
+
38
+ __m512i best_indices = _mm512_set1_epi32(-1);
39
+ __m512 best_distances = _mm512_set1_ps(worst_v);
32
40
  __m512i current_indices = _mm512_setr_epi32(
33
41
  0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
34
42
  __m512i offset = _mm512_set1_epi32(16);
35
43
 
36
- // The following loop tracks the rightmost index with the min distance.
37
- // -1 index values are ignored.
38
- const size_t k16 = (k / 16) * 16;
44
+ auto best_vs_cand_mask = [](__m512 best_d, __m512 cand_d) -> __mmask16 {
45
+ // Returns the mask of lanes where the current best is already
46
+ // (strictly) better than the candidate.
47
+ if constexpr (HC::is_max) {
48
+ return _mm512_cmp_ps_mask(best_d, cand_d, _CMP_LT_OS);
49
+ } else {
50
+ return _mm512_cmp_ps_mask(best_d, cand_d, _CMP_GT_OS);
51
+ }
52
+ };
53
+
54
+ const size_t k16 = (heap.k / 16) * 16;
39
55
  for (size_t iii = 0; iii < k16; iii += 16) {
40
56
  __m512i indices =
41
- _mm512_loadu_si512((const __m512i*)(ids.data() + iii));
42
- __m512 distances = _mm512_loadu_ps(dis.data() + iii);
57
+ _mm512_loadu_si512((const __m512i*)(heap.ids.data() + iii));
58
+ __m512 distances = _mm512_loadu_ps(heap.dis.data() + iii);
43
59
 
44
- // This mask filters out -1 values among indices.
45
60
  __mmask16 m1mask =
46
61
  _mm512_cmpgt_epi32_mask(_mm512_setzero_si512(), indices);
47
-
48
- __mmask16 dmask =
49
- _mm512_cmp_ps_mask(min_distances, distances, _CMP_LT_OS);
62
+ __mmask16 dmask = best_vs_cand_mask(best_distances, distances);
50
63
  __mmask16 finalmask = m1mask | dmask;
51
64
 
52
- const __m512i min_indices_new = _mm512_mask_blend_epi32(
53
- finalmask, current_indices, min_indices);
54
- const __m512 min_distances_new =
55
- _mm512_mask_blend_ps(finalmask, distances, min_distances);
65
+ const __m512i best_indices_new = _mm512_mask_blend_epi32(
66
+ finalmask, current_indices, best_indices);
67
+ const __m512 best_distances_new =
68
+ _mm512_mask_blend_ps(finalmask, distances, best_distances);
56
69
 
57
- min_indices = min_indices_new;
58
- min_distances = min_distances_new;
70
+ best_indices = best_indices_new;
71
+ best_distances = best_distances_new;
59
72
 
60
73
  current_indices = _mm512_add_epi32(current_indices, offset);
61
74
  }
62
75
 
63
- // leftovers
64
- if (k16 != static_cast<size_t>(k)) {
65
- const __mmask16 kmask = (1 << (k - k16)) - 1;
76
+ // Leftovers.
77
+ if (k16 != static_cast<size_t>(heap.k)) {
78
+ const __mmask16 kmask = (1 << (heap.k - k16)) - 1;
66
79
 
67
80
  __m512i indices = _mm512_mask_loadu_epi32(
68
- _mm512_set1_epi32(-1), kmask, ids.data() + k16);
69
- __m512 distances = _mm512_maskz_loadu_ps(kmask, dis.data() + k16);
81
+ _mm512_set1_epi32(-1), kmask, heap.ids.data() + k16);
82
+ __m512 distances = _mm512_maskz_loadu_ps(kmask, heap.dis.data() + k16);
70
83
 
71
- // This mask filters out -1 values among indices.
72
84
  __mmask16 m1mask =
73
85
  _mm512_cmpgt_epi32_mask(_mm512_setzero_si512(), indices);
74
-
75
- __mmask16 dmask =
76
- _mm512_cmp_ps_mask(min_distances, distances, _CMP_LT_OS);
86
+ __mmask16 dmask = best_vs_cand_mask(best_distances, distances);
77
87
  __mmask16 finalmask = m1mask | dmask;
78
88
 
79
- const __m512i min_indices_new = _mm512_mask_blend_epi32(
80
- finalmask, current_indices, min_indices);
81
- const __m512 min_distances_new =
82
- _mm512_mask_blend_ps(finalmask, distances, min_distances);
89
+ const __m512i best_indices_new = _mm512_mask_blend_epi32(
90
+ finalmask, current_indices, best_indices);
91
+ const __m512 best_distances_new =
92
+ _mm512_mask_blend_ps(finalmask, distances, best_distances);
83
93
 
84
- min_indices = min_indices_new;
85
- min_distances = min_distances_new;
94
+ best_indices = best_indices_new;
95
+ best_distances = best_distances_new;
86
96
  }
87
97
 
88
- // grab min distance
89
- min_dis = _mm512_reduce_min_ps(min_distances);
90
- // blend
91
- __mmask16 mindmask =
92
- _mm512_cmpeq_ps_mask(min_distances, _mm512_set1_ps(min_dis));
93
- // pick the max one
94
- min_idx = _mm512_mask_reduce_max_epi32(mindmask, min_indices);
98
+ // Horizontal best: min for CMax (distance), max for CMin (similarity).
99
+ if constexpr (HC::is_max) {
100
+ best_dis = _mm512_reduce_min_ps(best_distances);
101
+ } else {
102
+ best_dis = _mm512_reduce_max_ps(best_distances);
103
+ }
104
+ // Tiebreak by picking the rightmost (largest) index among lanes
105
+ // matching the best distance, matching the original behavior.
106
+ __mmask16 best_lane_mask =
107
+ _mm512_cmpeq_ps_mask(best_distances, _mm512_set1_ps(best_dis));
108
+ best_idx = _mm512_mask_reduce_max_epi32(best_lane_mask, best_indices);
95
109
 
96
- if (min_idx == -1) {
110
+ if (best_idx == -1) {
97
111
  return -1;
98
112
  }
99
113
 
100
114
  if (vmin_out) {
101
- *vmin_out = min_dis;
115
+ *vmin_out = best_dis;
102
116
  }
103
- int ret = ids[min_idx];
104
- ids[min_idx] = -1;
105
- --nvalid;
117
+ int ret = heap.ids[best_idx];
118
+ heap.ids[best_idx] = -1;
119
+ --heap.nvalid;
106
120
  return ret;
107
121
  }
108
122
 
123
+ } // namespace
124
+
125
+ // Explicit specializations for AVX512
126
+ template <>
127
+ int pop_min_tpl<CMax<float, int32_t>, SIMDLevel::AVX512>(
128
+ MinimaxHeapT<CMax<float, int32_t>>* heap,
129
+ float* vmin_out) {
130
+ return pop_best_avx512<CMax<float, int32_t>>(*heap, vmin_out);
131
+ }
132
+
133
+ template <>
134
+ int pop_min_tpl<CMin<float, int32_t>, SIMDLevel::AVX512>(
135
+ MinimaxHeapT<CMin<float, int32_t>>* heap,
136
+ float* vmin_out) {
137
+ return pop_best_avx512<CMin<float, int32_t>>(*heap, vmin_out);
138
+ }
139
+
109
140
  } // namespace faiss
110
141
 
111
142
  #endif // COMPILE_SIMD_AVX512