Skip to content

Commit ed7e1f2

Browse files
algoriddlemeta-codesync[bot]
authored andcommitted
Hoist SIMD dispatch outside loops in 4 call sites (#5074)
Summary: Pull Request resolved: #5074 Move `with_simd_level` / `with_simd_level_256bit` calls outside the enclosing loops so the SIMD level is resolved once rather than on every iteration. Sites fixed: - distances.cpp: knn_inner_products_by_idx, knn_L2sqr_by_idx - NeuralNet.cpp: ZnLUTCodec::encode - ClusteringInitialization.cpp: init_kmpp_plus_plus Reviewed By: mdouze Differential Revision: D100144174 fbshipit-source-id: bd2369ed4fd9c3b5b54e435c7ee66a03f0e152df
1 parent 6e64c5d commit ed7e1f2

3 files changed

Lines changed: 50 additions & 50 deletions

File tree

faiss/impl/ClusteringInitialization.cpp

Lines changed: 14 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -221,27 +221,27 @@ void ClusteringInitialization::init_kmeans_plus_plus(
221221
std::vector<double> cumsum(n);
222222

223223
// Select remaining centroids using D² sampling
224-
for (size_t c = result.first_new_centroid_idx; c < k; c++) {
225-
// Compute cumulative sum
226-
cumsum[0] = min_distances[0];
227-
for (size_t i = 1; i < n; i++) {
228-
cumsum[i] = cumsum[i - 1] + min_distances[i];
229-
}
224+
with_simd_level([&]<SIMDLevel SL>() {
225+
for (size_t c = result.first_new_centroid_idx; c < k; c++) {
226+
// Compute cumulative sum
227+
cumsum[0] = min_distances[0];
228+
for (size_t i = 1; i < n; i++) {
229+
cumsum[i] = cumsum[i - 1] + min_distances[i];
230+
}
230231

231-
// Sample using precomputed cumsum
232-
size_t next_idx = sample_from_cumsum(cumsum, rng);
232+
// Sample using precomputed cumsum
233+
size_t next_idx = sample_from_cumsum(cumsum, rng);
233234

234-
float* new_centroid = centroids + c * d;
235-
std::memcpy(new_centroid, x + next_idx * d, d * sizeof(float));
235+
float* new_centroid = centroids + c * d;
236+
std::memcpy(new_centroid, x + next_idx * d, d * sizeof(float));
236237

237-
// Update min distances incrementally
238-
with_simd_level([&]<SIMDLevel SL>() {
238+
// Update min distances incrementally
239239
for (size_t i = 0; i < n; i++) {
240240
double dist = fvec_L2sqr<SL>(x + i * d, new_centroid, d);
241241
min_distances[i] = std::min(min_distances[i], dist);
242242
}
243-
});
244-
}
243+
}
244+
});
245245
}
246246

247247
void ClusteringInitialization::init_afkmc2(

faiss/utils/NeuralNet.cpp

Lines changed: 15 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -268,12 +268,12 @@ nn::Int32Tensor2D QINCoStep::encode(
268268
res = residuals->data();
269269
}
270270

271-
for (size_t i = 0; i < n; i++) {
272-
const float* q = x.data() + i * d;
273-
const float* db = zqs_r.data() + i * K * d;
274-
float dis_min = HUGE_VALF;
275-
int64_t idx = -1;
276-
with_simd_level([&]<SIMDLevel SL>() {
271+
with_simd_level([&]<SIMDLevel SL>() {
272+
for (size_t i = 0; i < n; i++) {
273+
const float* q = x.data() + i * d;
274+
const float* db = zqs_r.data() + i * K * d;
275+
float dis_min = HUGE_VALF;
276+
int64_t idx = -1;
277277
for (size_t j = 0; j < static_cast<size_t>(K); j++) {
278278
float dis = fvec_L2sqr<SL>(q, db, d);
279279
if (dis < dis_min) {
@@ -282,17 +282,17 @@ nn::Int32Tensor2D QINCoStep::encode(
282282
}
283283
db += d;
284284
}
285-
});
286-
codes.v[i] = idx;
287-
if (res) {
288-
const float* xhat_row = xhat.data() + i * d;
289-
const float* xhat_next_row = zqs_r.data() + (i * K + idx) * d;
290-
for (size_t j = 0; j < static_cast<size_t>(d); j++) {
291-
res[j] = xhat_next_row[j] - xhat_row[j];
285+
codes.v[i] = idx;
286+
if (res) {
287+
const float* xhat_row = xhat.data() + i * d;
288+
const float* xhat_next_row = zqs_r.data() + (i * K + idx) * d;
289+
for (size_t j = 0; j < static_cast<size_t>(d); j++) {
290+
res[j] = xhat_next_row[j] - xhat_row[j];
291+
}
292+
res += d;
292293
}
293-
res += d;
294294
}
295-
}
295+
});
296296
return codes;
297297
}
298298

faiss/utils/distances.cpp

Lines changed: 21 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -838,16 +838,16 @@ void knn_inner_products_by_idx(
838838
ld_ids = ny;
839839
}
840840

841+
with_simd_level([&]<SIMDLevel SL>() {
841842
#pragma omp parallel for if (nx > 100)
842-
for (int64_t i = 0; i < static_cast<int64_t>(nx); i++) {
843-
const float* x_ = x + i * d;
844-
const int64_t* idsi = ids + i * ld_ids;
845-
size_t j;
846-
float* __restrict simi = res_vals + i * k;
847-
int64_t* __restrict idxi = res_ids + i * k;
848-
minheap_heapify(k, simi, idxi);
843+
for (int64_t i = 0; i < static_cast<int64_t>(nx); i++) {
844+
const float* x_ = x + i * d;
845+
const int64_t* idsi = ids + i * ld_ids;
846+
size_t j;
847+
float* __restrict simi = res_vals + i * k;
848+
int64_t* __restrict idxi = res_ids + i * k;
849+
minheap_heapify(k, simi, idxi);
849850

850-
with_simd_level([&]<SIMDLevel SL>() {
851851
for (j = 0; j < nsubset; j++) {
852852
if (idsi[j] < 0 || static_cast<size_t>(idsi[j]) >= ny) {
853853
break;
@@ -858,9 +858,9 @@ void knn_inner_products_by_idx(
858858
minheap_replace_top(k, simi, idxi, ip, idsi[j]);
859859
}
860860
}
861-
});
862-
minheap_reorder(k, simi, idxi);
863-
}
861+
minheap_reorder(k, simi, idxi);
862+
}
863+
});
864864
}
865865

866866
void knn_L2sqr_by_idx(
@@ -878,14 +878,14 @@ void knn_L2sqr_by_idx(
878878
if (ld_ids < 0) {
879879
ld_ids = ny;
880880
}
881+
with_simd_level([&]<SIMDLevel SL>() {
881882
#pragma omp parallel for if (nx > 100)
882-
for (int64_t i = 0; i < static_cast<int64_t>(nx); i++) {
883-
const float* x_ = x + i * d;
884-
const int64_t* __restrict idsi = ids + i * ld_ids;
885-
float* __restrict simi = res_vals + i * k;
886-
int64_t* __restrict idxi = res_ids + i * k;
887-
maxheap_heapify(k, simi, idxi);
888-
with_simd_level([&]<SIMDLevel SL>() {
883+
for (int64_t i = 0; i < static_cast<int64_t>(nx); i++) {
884+
const float* x_ = x + i * d;
885+
const int64_t* __restrict idsi = ids + i * ld_ids;
886+
float* __restrict simi = res_vals + i * k;
887+
int64_t* __restrict idxi = res_ids + i * k;
888+
maxheap_heapify(k, simi, idxi);
889889
for (size_t j = 0; j < nsubset; j++) {
890890
if (idsi[j] < 0 || static_cast<size_t>(idsi[j]) >= ny) {
891891
break;
@@ -896,9 +896,9 @@ void knn_L2sqr_by_idx(
896896
maxheap_replace_top(k, simi, idxi, disij, idsi[j]);
897897
}
898898
}
899-
});
900-
maxheap_reorder(k, simi, idxi);
901-
}
899+
maxheap_reorder(k, simi, idxi);
900+
}
901+
});
902902
}
903903

904904
void pairwise_L2sqr(

0 commit comments

Comments
 (0)