faiss 0.6.1 → 0.6.3
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- checksums.yaml +4 -4
- data/CHANGELOG.md +8 -0
- data/lib/faiss/version.rb +1 -1
- data/vendor/faiss/faiss/AutoTune.cpp +3 -1
- data/vendor/faiss/faiss/Clustering.cpp +9 -1
- data/vendor/faiss/faiss/IVFlib.cpp +14 -3
- data/vendor/faiss/faiss/Index.h +2 -2
- data/vendor/faiss/faiss/IndexAdditiveQuantizer.cpp +9 -10
- data/vendor/faiss/faiss/IndexAdditiveQuantizerFastScan.cpp +2 -3
- data/vendor/faiss/faiss/IndexBinaryFromFloat.cpp +1 -2
- data/vendor/faiss/faiss/IndexBinaryHNSW.cpp +10 -12
- data/vendor/faiss/faiss/IndexBinaryHash.cpp +5 -9
- data/vendor/faiss/faiss/IndexBinaryIVF.cpp +5 -7
- data/vendor/faiss/faiss/IndexEDEN.cpp +273 -0
- data/vendor/faiss/faiss/IndexEDEN.h +57 -0
- data/vendor/faiss/faiss/IndexFastScan.cpp +15 -4
- data/vendor/faiss/faiss/IndexFlat.cpp +13 -50
- data/vendor/faiss/faiss/IndexHNSW.cpp +177 -148
- data/vendor/faiss/faiss/IndexIDMap.cpp +16 -3
- data/vendor/faiss/faiss/IndexIDMap.h +2 -0
- data/vendor/faiss/faiss/IndexIVF.cpp +19 -8
- data/vendor/faiss/faiss/IndexIVFAdditiveQuantizer.cpp +3 -3
- data/vendor/faiss/faiss/IndexIVFAdditiveQuantizerFastScan.cpp +3 -4
- data/vendor/faiss/faiss/IndexIVFEDEN.cpp +302 -0
- data/vendor/faiss/faiss/IndexIVFEDEN.h +70 -0
- data/vendor/faiss/faiss/IndexIVFFastScan.cpp +5 -6
- data/vendor/faiss/faiss/IndexIVFFlat.cpp +6 -5
- data/vendor/faiss/faiss/IndexIVFFlatPanorama.cpp +3 -3
- data/vendor/faiss/faiss/IndexIVFIndependentQuantizer.cpp +1 -1
- data/vendor/faiss/faiss/IndexIVFPQ.cpp +42 -25
- data/vendor/faiss/faiss/IndexIVFPQFastScan.cpp +0 -1
- data/vendor/faiss/faiss/IndexIVFPQR.cpp +2 -3
- data/vendor/faiss/faiss/IndexIVFRaBitQ.cpp +23 -62
- data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.cpp +180 -76
- data/vendor/faiss/faiss/IndexIVFRaBitQFastScan.h +5 -4
- data/vendor/faiss/faiss/IndexIVFSpectralHash.cpp +8 -6
- data/vendor/faiss/faiss/IndexLSH.cpp +2 -3
- data/vendor/faiss/faiss/IndexLattice.cpp +5 -0
- data/vendor/faiss/faiss/IndexNNDescent.cpp +10 -3
- data/vendor/faiss/faiss/IndexNSG.cpp +8 -4
- data/vendor/faiss/faiss/IndexPQ.cpp +6 -8
- data/vendor/faiss/faiss/IndexPreTransform.cpp +15 -0
- data/vendor/faiss/faiss/IndexRaBitQ.cpp +2 -2
- data/vendor/faiss/faiss/IndexRaBitQFastScan.cpp +1 -2
- data/vendor/faiss/faiss/IndexRaBitQFastScan.h +5 -1
- data/vendor/faiss/faiss/IndexRefine.cpp +30 -1
- data/vendor/faiss/faiss/IndexReplicas.cpp +1 -2
- data/vendor/faiss/faiss/IndexScalarQuantizer.cpp +68 -6
- data/vendor/faiss/faiss/IndexScalarQuantizer.h +10 -0
- data/vendor/faiss/faiss/IndexShards.cpp +2 -2
- data/vendor/faiss/faiss/IndexShardsIVF.cpp +2 -2
- data/vendor/faiss/faiss/MetaIndexes.cpp +2 -4
- data/vendor/faiss/faiss/SuperKMeans.cpp +256 -240
- data/vendor/faiss/faiss/SuperKMeans.h +30 -0
- data/vendor/faiss/faiss/VectorTransform.cpp +33 -2
- data/vendor/faiss/faiss/clone_index.cpp +5 -0
- data/vendor/faiss/faiss/cppcontrib/SaDecodeKernels.h +1 -1
- data/vendor/faiss/faiss/cppcontrib/sa_decode/Level2-neon-inl.h +902 -12
- data/vendor/faiss/faiss/cppcontrib/sa_decode/PQ-neon-inl.h +702 -10
- data/vendor/faiss/faiss/factory_tools.cpp +51 -4
- data/vendor/faiss/faiss/gpu/GpuCloner.cpp +11 -11
- data/vendor/faiss/faiss/gpu/GpuIndex.h +34 -11
- data/vendor/faiss/faiss/gpu/GpuIndexCagra.h +47 -0
- data/vendor/faiss/faiss/gpu/GpuIndexIVF.h +17 -0
- data/vendor/faiss/faiss/gpu/GpuIndexIVFScalarQuantizer.h +16 -0
- data/vendor/faiss/faiss/gpu/GpuResources.h +3 -2
- data/vendor/faiss/faiss/gpu/StandardGpuResources.cpp +11 -12
- data/vendor/faiss/faiss/gpu/StandardGpuResources.h +3 -3
- data/vendor/faiss/faiss/gpu/perf/PerfClustering.cpp +1 -1
- data/vendor/faiss/faiss/gpu/perf/PerfIVFPQAdd.cpp +2 -2
- data/vendor/faiss/faiss/gpu/test/TestGpuIndexIVFScalarQuantizer.cpp +180 -0
- data/vendor/faiss/faiss/gpu_metal/MetalDistance.h +87 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndex.h +7 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFFlat.h +177 -0
- data/vendor/faiss/faiss/gpu_metal/MetalIndexIVFPQ.h +88 -0
- data/vendor/faiss/faiss/gpu_metal/MetalKernels.h +48 -3
- data/vendor/faiss/faiss/gpu_metal/MetalPythonBridge.h +45 -0
- data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFFlat.h +193 -0
- data/vendor/faiss/faiss/gpu_metal/impl/MetalIVFPQ.h +134 -0
- data/vendor/faiss/faiss/impl/ClusteringInitialization.cpp +2 -2
- data/vendor/faiss/faiss/impl/DistanceComputer.h +34 -0
- data/vendor/faiss/faiss/impl/EDENQuantizer.h +119 -0
- data/vendor/faiss/faiss/impl/HNSW.cpp +658 -344
- data/vendor/faiss/faiss/impl/HNSW.h +51 -13
- data/vendor/faiss/faiss/impl/LocalSearchQuantizer.cpp +2 -2
- data/vendor/faiss/faiss/impl/NSG.cpp +18 -12
- data/vendor/faiss/faiss/impl/Panorama.h +20 -7
- data/vendor/faiss/faiss/impl/PolysemousTraining.cpp +152 -84
- data/vendor/faiss/faiss/impl/ProductQuantizer.cpp +59 -24
- data/vendor/faiss/faiss/impl/RaBitQUtils.cpp +45 -37
- data/vendor/faiss/faiss/impl/RaBitQUtils.h +35 -0
- data/vendor/faiss/faiss/impl/RaBitQuantizer.cpp +175 -68
- data/vendor/faiss/faiss/impl/RaBitQuantizer.h +19 -0
- data/vendor/faiss/faiss/impl/RaBitQuantizerMultiBit.cpp +2 -11
- data/vendor/faiss/faiss/impl/ResultHandler.h +26 -31
- data/vendor/faiss/faiss/impl/ScalarQuantizer.cpp +522 -58
- data/vendor/faiss/faiss/impl/ScalarQuantizer.h +70 -0
- data/vendor/faiss/faiss/impl/ThreadedIndex-inl.h +2 -2
- data/vendor/faiss/faiss/impl/VisitedTable.cpp +33 -13
- data/vendor/faiss/faiss/impl/VisitedTable.h +88 -33
- data/vendor/faiss/faiss/impl/binary_hamming/IndexBinaryIVF_impl.h +1 -1
- data/vendor/faiss/faiss/impl/binary_hamming/avx2.cpp +4 -4
- data/vendor/faiss/faiss/impl/fast_scan/dispatching.h +38 -3
- data/vendor/faiss/faiss/impl/hnsw/LockVector.cpp +1 -1
- data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.cpp +35 -43
- data/vendor/faiss/faiss/impl/hnsw/MinimaxHeap.h +64 -15
- data/vendor/faiss/faiss/impl/hnsw/avx2.cpp +86 -40
- data/vendor/faiss/faiss/impl/hnsw/avx512.cpp +81 -50
- data/vendor/faiss/faiss/impl/index_read.cpp +476 -75
- data/vendor/faiss/faiss/impl/index_write.cpp +56 -4
- data/vendor/faiss/faiss/impl/io_macros.h +25 -0
- data/vendor/faiss/faiss/impl/lattice_Zn.cpp +8 -9
- data/vendor/faiss/faiss/impl/platform_macros.h +15 -9
- data/vendor/faiss/faiss/impl/polysemous_training/avx512.cpp +284 -0
- data/vendor/faiss/faiss/impl/polysemous_training/dispatch.h +115 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/IVFPQ_QueryTables.cpp +0 -1
- data/vendor/faiss/faiss/impl/pq_code_distance/PQDistanceComputer_impl.h +26 -15
- data/vendor/faiss/faiss/impl/pq_code_distance/avx2.cpp +6 -4
- data/vendor/faiss/faiss/impl/pq_code_distance/avx512.cpp +2 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/neon.cpp +2 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-generic.cpp +20 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-inl.h +36 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_code_distance-sve.cpp +5 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/pq_scan_impl.h +105 -0
- data/vendor/faiss/faiss/impl/pq_code_distance/rvv.cpp +2 -0
- data/vendor/faiss/faiss/impl/result_handler/ResultHandler.cpp +195 -0
- data/vendor/faiss/faiss/impl/result_handler/avx2.cpp +133 -0
- data/vendor/faiss/faiss/impl/result_handler/avx512.cpp +281 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx2.cpp +72 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer-avx512.cpp +228 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/EDENQuantizer.cpp +882 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/distance_computers.h +6 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/quantizers.h +336 -26
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx2.cpp +331 -32
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-impl.h +553 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512-spr.cpp +558 -0
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-avx512.cpp +284 -45
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-dispatch.h +502 -3
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-neon.cpp +157 -32
- data/vendor/faiss/faiss/impl/scalar_quantizer/sq-rvv.cpp +26 -0
- data/vendor/faiss/faiss/impl/simd_dispatch.h +86 -8
- data/vendor/faiss/faiss/index_factory.cpp +37 -7
- data/vendor/faiss/faiss/index_io.h +16 -0
- data/vendor/faiss/faiss/invlists/DirectMap.cpp +5 -2
- data/vendor/faiss/faiss/invlists/InvertedLists.cpp +15 -15
- data/vendor/faiss/faiss/invlists/InvertedLists.h +2 -2
- data/vendor/faiss/faiss/invlists/OnDiskInvertedLists.cpp +19 -4
- data/vendor/faiss/faiss/python/python_callbacks.cpp +3 -1
- data/vendor/faiss/faiss/svs/IndexSVSFaissUtils.h +60 -0
- data/vendor/faiss/faiss/svs/IndexSVSFlat.cpp +26 -1
- data/vendor/faiss/faiss/svs/IndexSVSFlat.h +13 -0
- data/vendor/faiss/faiss/svs/IndexSVSIVF.cpp +1 -1
- data/vendor/faiss/faiss/svs/IndexSVSIVFLeanVec.cpp +1 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamana.cpp +150 -23
- data/vendor/faiss/faiss/svs/IndexSVSVamana.h +30 -7
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.cpp +3 -2
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLVQ.h +2 -1
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.cpp +65 -25
- data/vendor/faiss/faiss/svs/IndexSVSVamanaLeanVec.h +3 -2
- data/vendor/faiss/faiss/utils/approx_topk_hamming/approx_topk_hamming.h +1 -1
- data/vendor/faiss/faiss/utils/bf16.h +34 -0
- data/vendor/faiss/faiss/utils/distances.cpp +14 -2
- data/vendor/faiss/faiss/utils/distances_simd.cpp +4 -4
- data/vendor/faiss/faiss/utils/extra_distances.cpp +4 -14
- data/vendor/faiss/faiss/utils/extra_distances.h +1 -2
- data/vendor/faiss/faiss/utils/hamming.cpp +9 -9
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx2.cpp +2 -1
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_avx512_spr.cpp +15 -0
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512.h +6 -30
- data/vendor/faiss/faiss/utils/hamming_distance/hamming_computer-avx512_spr.h +171 -0
- data/vendor/faiss/faiss/utils/partitioning.cpp +0 -2
- data/vendor/faiss/faiss/utils/quantize_lut.cpp +29 -8
- data/vendor/faiss/faiss/utils/rabitq_simd.h +202 -0
- data/vendor/faiss/faiss/utils/simd_impl/distances_avx2.cpp +0 -1
- data/vendor/faiss/faiss/utils/simd_impl/distances_avx512.cpp +263 -15
- data/vendor/faiss/faiss/utils/simd_impl/distances_rvv.cpp +160 -18
- data/vendor/faiss/faiss/utils/simd_impl/partitioning_simdlib256.h +14 -68
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx2.cpp +245 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512.cpp +273 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_avx512_spr.cpp +435 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_neon.cpp +11 -0
- data/vendor/faiss/faiss/utils/simd_impl/rabitq_rvv.cpp +143 -6
- data/vendor/faiss/faiss/utils/simd_levels.cpp +56 -2
- data/vendor/faiss/faiss/utils/simd_levels.h +14 -0
- data/vendor/faiss/faiss/utils/utils.cpp +9 -27
- metadata +27 -2
|
@@ -16,89 +16,135 @@
|
|
|
16
16
|
|
|
17
17
|
namespace faiss {
|
|
18
18
|
|
|
19
|
-
|
|
20
|
-
|
|
21
|
-
|
|
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
|
|
27
|
-
float
|
|
41
|
+
int32_t best_idx = -1;
|
|
42
|
+
float best_dis = worst_v;
|
|
28
43
|
|
|
29
44
|
size_t iii = 0;
|
|
30
45
|
|
|
31
|
-
__m256i
|
|
32
|
-
__m256
|
|
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
|
-
//
|
|
38
|
-
|
|
39
|
-
|
|
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
|
-
//
|
|
59
|
+
// Mask out -1 indices (invalid entries).
|
|
48
60
|
__m256i m1mask = _mm256_cmpgt_epi32(_mm256_setzero_si256(), indices);
|
|
49
61
|
|
|
50
|
-
|
|
51
|
-
|
|
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
|
|
76
|
+
const __m256i best_indices_new = _mm256_castps_si256(_mm256_blendv_ps(
|
|
55
77
|
_mm256_castsi256_ps(current_indices),
|
|
56
|
-
_mm256_castsi256_ps(
|
|
78
|
+
_mm256_castsi256_ps(best_indices),
|
|
57
79
|
finalmask));
|
|
58
80
|
|
|
59
|
-
const __m256
|
|
60
|
-
_mm256_blendv_ps(distances,
|
|
81
|
+
const __m256 best_distances_new =
|
|
82
|
+
_mm256_blendv_ps(distances, best_distances, finalmask);
|
|
61
83
|
|
|
62
|
-
|
|
63
|
-
|
|
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
|
|
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,
|
|
72
|
-
_mm256_storeu_si256((__m256i*)vidx8,
|
|
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
|
-
|
|
76
|
-
|
|
77
|
-
|
|
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
|
-
//
|
|
82
|
-
for (; iii < static_cast<size_t>(k); iii++) {
|
|
83
|
-
if (ids[iii]
|
|
84
|
-
|
|
85
|
-
|
|
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 (
|
|
118
|
+
if (best_idx == -1) {
|
|
90
119
|
return -1;
|
|
91
120
|
}
|
|
92
121
|
|
|
93
122
|
if (vmin_out) {
|
|
94
|
-
*vmin_out =
|
|
123
|
+
*vmin_out = best_dis;
|
|
95
124
|
}
|
|
96
|
-
int ret = ids[
|
|
97
|
-
ids[
|
|
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
|
-
|
|
20
|
-
|
|
21
|
-
|
|
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
|
-
|
|
27
|
-
|
|
31
|
+
constexpr float worst_v = HC::is_max
|
|
32
|
+
? std::numeric_limits<float>::infinity()
|
|
33
|
+
: -std::numeric_limits<float>::infinity();
|
|
28
34
|
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
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
|
-
|
|
37
|
-
|
|
38
|
-
|
|
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
|
|
53
|
-
finalmask, current_indices,
|
|
54
|
-
const __m512
|
|
55
|
-
_mm512_mask_blend_ps(finalmask, 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
|
-
|
|
58
|
-
|
|
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
|
-
//
|
|
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
|
|
80
|
-
finalmask, current_indices,
|
|
81
|
-
const __m512
|
|
82
|
-
_mm512_mask_blend_ps(finalmask, 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
|
-
|
|
85
|
-
|
|
94
|
+
best_indices = best_indices_new;
|
|
95
|
+
best_distances = best_distances_new;
|
|
86
96
|
}
|
|
87
97
|
|
|
88
|
-
//
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
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 (
|
|
110
|
+
if (best_idx == -1) {
|
|
97
111
|
return -1;
|
|
98
112
|
}
|
|
99
113
|
|
|
100
114
|
if (vmin_out) {
|
|
101
|
-
*vmin_out =
|
|
115
|
+
*vmin_out = best_dis;
|
|
102
116
|
}
|
|
103
|
-
int ret = ids[
|
|
104
|
-
ids[
|
|
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
|