52#include <oneapi/dpl/algorithm>
53#include <oneapi/dpl/execution>
54#include <sycl/sycl.hpp>
78 int64_t& tile_queries,
80 int64_t max_tile_queries = 128,
81 int64_t tile_points_alignment = 128) {
82 tile_queries = std::min<int64_t>(num_queries, max_tile_queries);
83 tile_queries = std::max<int64_t>(tile_queries, 1);
84 tile_points = std::max<int64_t>(tile_bytes / (tile_queries * element_size),
86 if (tile_points > tile_points_alignment) {
88 (tile_points / tile_points_alignment) * tile_points_alignment;
90 tile_points = std::min<int64_t>(tile_points, num_points);
91 tile_points = std::max<int64_t>(tile_points, 1);
102template <
typename T,
typename TIndex,
int K>
105 int left = 2 * root + 1, right = 2 * root + 2, largest = root;
106 if (left < K && (d[left] > d[largest] ||
107 (d[left] == d[largest] && idx[left] > idx[largest])))
109 if (right < K && (d[right] > d[largest] || (d[right] == d[largest] &&
110 idx[right] > idx[largest])))
112 if (largest == root)
break;
114 d[root] = d[largest];
116 TIndex ti = idx[root];
117 idx[root] = idx[largest];
124template <
typename T,
typename TIndex,
int K>
126 for (
int end = K - 1; end > 0; --end) {
135 int left = 2 * root + 1, right = 2 * root + 2, largest = root;
137 (d[left] > d[largest] ||
138 (d[left] == d[largest] && idx[left] > idx[largest])))
141 (d[right] > d[largest] ||
142 (d[right] == d[largest] && idx[right] > idx[largest])))
144 if (largest == root)
break;
146 d[root] = d[largest];
148 TIndex ti2 = idx[root];
149 idx[root] = idx[largest];
186template <
typename T,
typename TIndex,
int K>
189 int64_t distance_stride,
190 const T* point_norms_ptr,
195 TIndex* best_idx_ptr,
198 const size_t wg = core::sy::PreferredWorkGroupSize(
queue.get_device());
199 const size_t global_size =
200 ((
static_cast<size_t>(num_queries) + wg - 1) / wg) * wg;
202 sycl::nd_range<1>(sycl::range<1>(global_size), sycl::range<1>(wg)),
203 [=](sycl::nd_item<1> it) [[intel::kernel_args_restrict]] {
204 const int64_t q = it.get_global_id(0);
205 if (q >= num_queries)
return;
206 const T* qrow = neg2qp_ptr + q * distance_stride;
207 T* qd = best_dist_ptr + q * K;
208 TIndex* qi = best_idx_ptr + q * K;
214 for (
int i = 0; i < K; ++i) {
224 for (int64_t p = 0; p < num_points; ++p) {
225 const T
dist = qrow[p] + point_norms_ptr[p];
226 if (use_threshold &&
dist > threshold)
continue;
227 const TIndex gp = point_offset +
static_cast<TIndex
>(p);
229 if (
dist < d[0] || (
dist == d[0] && gp < idx[0])) {
232 HeapifyDown<T, TIndex, K>(d, idx, 0);
236 for (
int i = 0; i < K; ++i) {
253template <
typename T,
typename TIndex,
int K>
256 const T* running_dist_ptr,
257 const TIndex* running_idx_ptr,
261 const T* query_norms_ptr) {
262 const size_t wg = core::sy::PreferredWorkGroupSize(
queue.get_device());
263 const size_t global_size =
264 ((
static_cast<size_t>(num_queries) + wg - 1) / wg) * wg;
266 sycl::nd_range<1>(sycl::range<1>(global_size), sycl::range<1>(wg)),
267 [=](sycl::nd_item<1> it) [[intel::kernel_args_restrict]] {
268 const int64_t q = it.get_global_id(0);
269 if (q >= num_queries)
return;
272 for (
int i = 0; i < K; ++i) {
273 d[i] = running_dist_ptr[q * K + i];
274 idx[i] = running_idx_ptr[q * K + i];
276 HeapSort<T, TIndex, K>(d, idx);
278 const T qnorm = query_norms_ptr ? query_norms_ptr[q] : T(0);
279 T* qout_d = out_dist_ptr + q * actual_k;
280 TIndex* qout_i = out_idx_ptr + q * actual_k;
281 for (int64_t i = 0; i < actual_k; ++i) {
283 if (query_norms_ptr)
dist = sycl::fmax(T(0),
dist + qnorm);
298 if (k <= 1)
return 1;
299 if (k <= 2)
return 2;
300 if (k <= 4)
return 4;
301 if (k <= 8)
return 8;
302 if (k <= 16)
return 16;
303 if (k <= 32)
return 32;
304 if (k <= 64)
return 64;
305 if (k <= 128)
return 128;
306 if (k <= 256)
return 256;
311template <
typename T,
typename TIndex>
314 int64_t distance_stride,
315 const T* point_norms_ptr,
321 TIndex* best_idx_ptr,
324#define CALL_UPDATE(Kval) \
325 UpdateTopKFromTile<T, TIndex, Kval>( \
326 queue, neg2qp_ptr, distance_stride, point_norms_ptr, num_queries, \
327 num_points, point_offset, best_dist_ptr, best_idx_ptr, \
328 use_threshold, threshold)
331 else if (k_bucket <= 2)
333 else if (k_bucket <= 4)
335 else if (k_bucket <= 8)
337 else if (k_bucket <= 16)
339 else if (k_bucket <= 32)
341 else if (k_bucket <= 64)
343 else if (k_bucket <= 128)
345 else if (k_bucket <= 256)
353template <
typename T,
typename TIndex>
356 const T* running_dist_ptr,
357 const TIndex* running_idx_ptr,
362 const T* query_norms_ptr) {
363#define CALL_FINALIZE(Kval) \
364 FinalizeTopK<T, TIndex, Kval>(queue, num_queries, running_dist_ptr, \
365 running_idx_ptr, out_dist_ptr, out_idx_ptr, \
366 actual_k, query_norms_ptr)
369 else if (k_bucket <= 2)
371 else if (k_bucket <= 4)
373 else if (k_bucket <= 8)
375 else if (k_bucket <= 16)
377 else if (k_bucket <= 32)
379 else if (k_bucket <= 64)
381 else if (k_bucket <= 128)
383 else if (k_bucket <= 256)
418template <
typename T,
typename TIndex,
int NDIM,
int K,
int SG>
422template <
typename T,
typename TIndex,
int NDIM,
int K,
int SG>
425 const T* queries_ptr,
431 int64_t subgroups_per_wg,
432 int64_t tile_points) {
433 if (num_points <= 0 || num_queries <= 0)
return;
435 const int64_t wg_size = subgroups_per_wg * SG;
436 const int64_t num_wgs =
437 (num_queries + subgroups_per_wg - 1) / subgroups_per_wg;
438 const int64_t global_size = num_wgs * wg_size;
439 const int64_t tp = std::min<int64_t>(tile_points, num_points);
440 const int64_t num_tiles = (num_points + tp - 1) / tp;
442 queue.submit([&](sycl::handler& h) {
443 sycl::local_accessor<T, 1> slm(sycl::range<1>(2 * tp * NDIM), h);
445 sycl::nd_range<1>(sycl::range<1>(global_size),
446 sycl::range<1>(wg_size)),
447 [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(
448 SG)]] [[intel::kernel_args_restrict]] {
449 const auto sg = it.get_sub_group();
450 const int64_t lane = sg.get_local_id()[0];
451 const int64_t sg_id_in_wg = sg.get_group_id()[0];
452 const int64_t wg_id = it.get_group(0);
453 const int64_t local_lin = it.get_local_linear_id();
454 const int64_t local_range = it.get_local_range(0);
456 const int64_t query_idx =
457 wg_id * subgroups_per_wg + sg_id_in_wg;
458 const bool active_query = query_idx < num_queries;
466 const int64_t qrow = active_query ? query_idx : 0;
467 for (
int d = 0; d < NDIM; ++d) {
468 q[d] = queries_ptr[qrow * NDIM + d];
475 for (
int i = 0; i < K; ++i) {
476 d[i] = std::numeric_limits<T>::max();
480 for (int64_t
t = 0;
t < num_tiles; ++
t) {
481 const int64_t cur =
t & 1;
482 const int64_t cur_start =
t * tp;
483 const int64_t cur_n =
484 std::min<int64_t>(tp, num_points - cur_start);
488 for (int64_t e = local_lin; e < cur_n * NDIM;
490 const int64_t p = e / NDIM, dd = e % NDIM;
491 slm[cur * tp * NDIM + e] =
492 points_ptr[(cur_start + p) * NDIM + dd];
494 sycl::group_barrier(it.get_group());
502 if (
t + 1 < num_tiles) {
503 const int64_t nxt = 1 - cur;
504 const int64_t nxt_start = (
t + 1) * tp;
505 const int64_t nxt_n = std::min<int64_t>(
506 tp, num_points - nxt_start);
507 for (int64_t e = local_lin; e < nxt_n * NDIM;
509 const int64_t p = e / NDIM, dd = e % NDIM;
510 slm[nxt * tp * NDIM + e] =
511 points_ptr[(nxt_start + p) * NDIM + dd];
516 for (int64_t p_local = lane; p_local < cur_n;
520 cur * tp * NDIM + p_local * NDIM;
521 for (
int dd = 0; dd < NDIM; ++dd) {
522 const T diff = q[dd] - slm[base + dd];
525 const TIndex gp =
static_cast<TIndex
>(
526 cur_start + p_local);
527 if (
dist < d[K - 1] ||
528 (
dist == d[K - 1] && gp < idx[K - 1])) {
533 (d[pos - 1] > d[pos] ||
534 (d[pos - 1] == d[pos] &&
535 idx[pos - 1] > idx[pos]))) {
539 TIndex ti = idx[pos - 1];
540 idx[pos - 1] = idx[pos];
553 sycl::group_barrier(it.get_group());
556 if (!active_query)
return;
562 for (
int step = 1; step < SG; step <<= 1) {
563 const int64_t partner = lane ^ step;
566 for (
int i = 0; i < K; ++i) {
567 od[i] = sycl::select_from_group(sg, d[i], partner);
568 oidx[i] = sycl::select_from_group(sg, idx[i],
574 for (
int o = 0; o < K; ++o) {
579 (d[a] == od[b] && idx[a] <= oidx[b])));
590 for (
int o = 0; o < K; ++o) {
597 T* od = out_dist_ptr + query_idx * actual_k;
598 TIndex* oi = out_idx_ptr + query_idx * actual_k;
599 for (int64_t i = 0; i < actual_k; ++i) {
600 od[i] = sycl::fmax(T(0), d[i]);
612template <
typename T,
typename TIndex,
int NDIM,
int SG>
615 const T* queries_ptr,
621 int64_t subgroups_per_wg,
622 int64_t tile_points) {
623 const int64_t k_bucket =
KBucket(actual_k);
624#define CALL_DIRECT(Kval) \
625 KnnDirect<T, TIndex, NDIM, Kval, SG>( \
626 queue, points_ptr, queries_ptr, num_points, num_queries, actual_k, \
627 out_dist_ptr, out_idx_ptr, subgroups_per_wg, tile_points)
630 else if (k_bucket <= 2)
632 else if (k_bucket <= 4)
634 else if (k_bucket <= 8)
636 else if (k_bucket <= 16)
645template <
typename T,
typename TIndex,
int NDIM>
648 const T* queries_ptr,
654 int64_t subgroups_per_wg,
655 int64_t tile_points) {
656 if constexpr (std::is_same_v<T, double>) {
657 const auto sg_sizes =
659 .get_info<sycl::info::device::sub_group_sizes>();
660 const bool supports_subgroup_8 =
661 std::find(sg_sizes.begin(), sg_sizes.end(),
size_t(8)) !=
663 if (supports_subgroup_8) {
664 DispatchKnnDirectKForSG<T, TIndex, NDIM, 8>(
665 queue, points_ptr, queries_ptr, num_points, num_queries,
666 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
669 DispatchKnnDirectKForSG<T, TIndex, NDIM, 16>(
670 queue, points_ptr, queries_ptr, num_points, num_queries,
671 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
675 DispatchKnnDirectKForSG<T, TIndex, NDIM, 16>(
676 queue, points_ptr, queries_ptr, num_points, num_queries,
677 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
687template <
typename T,
typename TIndex>
690 const T* queries_ptr,
707 const size_t local_mem_bytes =
709 .get_info<sycl::info::device::local_mem_size>();
711 const int64_t max_tile_points_by_slm =
static_cast<int64_t
>(
712 (local_mem_bytes * 9 / 10) / (2 * dim *
sizeof(T)));
713 tile_points = std::min(tile_points,
714 std::max<int64_t>(max_tile_points_by_slm, 1));
730 const size_t max_wg_size =
732 .get_info<sycl::info::device::max_work_group_size>();
733 subgroups_per_wg = std::min(subgroups_per_wg,
734 static_cast<int64_t
>(max_wg_size / 16));
735 subgroups_per_wg = std::max<int64_t>(subgroups_per_wg, 1);
737#define CALL_DIM(NDIMVAL) \
738 DispatchKnnDirectK<T, TIndex, NDIMVAL>( \
739 queue, points_ptr, queries_ptr, num_points, num_queries, actual_k, \
740 out_dist_ptr, out_idx_ptr, subgroups_per_wg, tile_points)
767 utility::LogError(
"DispatchKnnDirect only supports dim 1 to {}.",
787template <
typename T,
typename TIndex,
int K>
788inline void HeapifyDownActive(T* local_d,
794 int left = 2 * i + 1, right = 2 * i + 2, largest = i;
795 if (left < active_k && (local_d[left] > local_d[largest] ||
796 (local_d[left] == local_d[largest] &&
797 local_i[left] > local_i[largest])))
799 if (right < active_k && (local_d[right] > local_d[largest] ||
800 (local_d[right] == local_d[largest] &&
801 local_i[right] > local_i[largest])))
803 if (largest == i)
break;
805 local_d[i] = local_d[largest];
806 local_d[largest] = td;
807 TIndex ti = local_i[i];
808 local_i[i] = local_i[largest];
809 local_i[largest] = ti;
815template <
typename T,
typename TIndex,
int K>
816void SelectTopKQueriesHeap(sycl::queue&
queue,
817 const T* distances_ptr,
818 int64_t distance_query_stride,
823 TIndex* out_indices_ptr,
824 T* out_distances_ptr,
825 int64_t out_query_stride,
827 const T* query_norms_ptr,
829 T scalar_threshold) {
830 const T inf = std::numeric_limits<T>::max();
831 const int64_t actual_knn = std::min(
knn, num_points);
833 const size_t wg = core::sy::PreferredWorkGroupSize(
queue.get_device());
834 const size_t global_size =
835 ((
static_cast<size_t>(num_queries) + wg - 1) / wg) * wg;
837 sycl::nd_range<1>(sycl::range<1>(global_size), sycl::range<1>(wg)),
838 [=](sycl::nd_item<1> it) [[intel::kernel_args_restrict]] {
839 const int64_t q = it.get_global_id(0);
840 if (q >= num_queries)
return;
841 const T* qd = distances_ptr + q * distance_query_stride;
842 TIndex* qout_i = out_indices_ptr + q * out_query_stride;
843 T* qout_d = out_distances_ptr + q * out_query_stride;
845 const T thr = (use_threshold && query_norms_ptr)
846 ? (radius_sq - query_norms_ptr[q])
851 for (
int k = 0; k < actual_knn; ++k) {
853 local_i[k] = TIndex(-1);
856 for (TIndex p = 0; p < static_cast<TIndex>(num_points); ++p) {
857 const T
dist = qd[p];
858 if (use_threshold &&
dist > thr)
continue;
859 if (
dist < local_d[0] ||
860 (
dist == local_d[0] &&
861 index_offset + p < index_offset + local_i[0])) {
864 HeapifyDownActive<T, TIndex, K>(
866 static_cast<int>(actual_knn));
870 for (
int i = 1; i < actual_knn; ++i) {
871 T key_d = local_d[i];
872 TIndex key_i = local_i[i];
875 (local_d[j] > key_d ||
876 (local_d[j] == key_d && local_i[j] > key_i))) {
877 local_d[j + 1] = local_d[j];
878 local_i[j + 1] = local_i[j];
881 local_d[j + 1] = key_d;
882 local_i[j + 1] = key_i;
885 for (int64_t k = 0; k <
knn; ++k) {
886 if (k >= actual_knn || local_i[k] == TIndex(-1)) {
887 qout_i[k] = TIndex(-1);
890 qout_i[k] = index_offset + local_i[k];
891 qout_d[k] = local_d[k];
901template <
typename T,
typename TIndex>
903 const T* distances_ptr,
904 int64_t distance_query_stride,
910 TIndex* out_indices_ptr,
911 T* out_distances_ptr,
912 int64_t out_query_stride,
914 const T* query_norms_ptr,
916 T scalar_threshold) {
917#define CALL_SELECT(Kval) \
918 SelectTopKQueriesHeap<T, TIndex, Kval>( \
919 queue, distances_ptr, distance_query_stride, num_queries, \
920 num_points, knn, index_offset, out_indices_ptr, out_distances_ptr, \
921 out_query_stride, use_threshold, query_norms_ptr, radius_sq, \
925 else if (k_bucket <= 2)
927 else if (k_bucket <= 4)
929 else if (k_bucket <= 8)
931 else if (k_bucket <= 16)
933 else if (k_bucket <= 32)
935 else if (k_bucket <= 64)
937 else if (k_bucket <= 128)
939 else if (k_bucket <= 256)
954template <
typename T,
typename TIndex>
956 const T* distances_ptr,
957 int64_t distance_query_stride,
962 TIndex* scratch_indices_ptr,
963 int64_t scratch_query_stride,
964 TIndex* out_indices_ptr,
965 T* out_distances_ptr,
966 int64_t out_query_stride,
967 bool use_threshold =
false,
968 const T* query_norms_ptr =
nullptr,
970 T scalar_threshold = T(0)) {
971 if (num_queries == 0 || num_points == 0 ||
knn <= 0)
return;
973 const T inf = std::numeric_limits<T>::max();
974 const int64_t actual_knn = std::min(
knn, num_points);
979 DispatchSelectTopKQueries<T, TIndex>(
980 queue, distances_ptr, distance_query_stride, num_queries,
981 num_points,
knn, k_bucket, index_offset, out_indices_ptr,
982 out_distances_ptr, out_query_stride, use_threshold,
983 query_norms_ptr, radius_sq, scalar_threshold);
986 auto policy = oneapi::dpl::execution::make_device_policy(
queue);
988 sycl::range<2>(num_queries, num_points),
989 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
990 scratch_indices_ptr[
id[0] * scratch_query_stride +
id[1]] =
991 static_cast<TIndex
>(
id[1]);
993 queue.wait_and_throw();
995 for (int64_t qi = 0; qi < num_queries; ++qi) {
996 TIndex* q_scratch = scratch_indices_ptr + qi * scratch_query_stride;
997 const T* q_dist = distances_ptr + qi * distance_query_stride;
1007 std::partial_sort(policy, q_scratch, q_scratch + actual_knn,
1008 q_scratch + num_points,
1009 [q_dist](TIndex lhs, TIndex rhs) {
1010 const T ld = q_dist[lhs];
1011 const T rd = q_dist[rhs];
1012 if (ld < rd)
return true;
1013 if (rd < ld)
return false;
1019 sycl::range<2>(num_queries,
knn),
1020 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
1021 const int64_t qi =
id[0], k =
id[1];
1022 TIndex* qout_i = out_indices_ptr + qi * out_query_stride;
1023 T* qout_d = out_distances_ptr + qi * out_query_stride;
1024 if (k >= actual_knn) {
1025 qout_i[k] = TIndex(-1);
1030 scratch_indices_ptr[qi * scratch_query_stride + k];
1041 distances_ptr[qi * distance_query_stride + li];
1042 const T thr = (use_threshold && query_norms_ptr)
1043 ? (radius_sq - query_norms_ptr[qi])
1045 if (use_threshold &&
dist > thr) {
1046 qout_i[k] = TIndex(-1);
1050 qout_i[k] = index_offset + li;
1058template <
typename T,
typename TIndex>
1060 const T* curr_dist_ptr,
1061 const TIndex* curr_idx_ptr,
1062 int64_t curr_stride,
1063 const T* cand_dist_ptr,
1064 const TIndex* cand_idx_ptr,
1065 int64_t cand_stride,
1066 int64_t num_queries,
1068 TIndex* scratch_ptr,
1069 int64_t scratch_stride,
1070 TIndex* out_idx_ptr,
1072 int64_t out_stride) {
1073 if (num_queries == 0 ||
knn <= 0)
return;
1074 const T inf = std::numeric_limits<T>::max();
1079 sycl::range<1>(num_queries),
1080 [=](sycl::id<1>
id) [[intel::kernel_args_restrict]] {
1081 const int64_t q =
id[0];
1082 const T* qcd = curr_dist_ptr + q * curr_stride;
1083 const TIndex* qci = curr_idx_ptr + q * curr_stride;
1084 const T* qad = cand_dist_ptr + q * cand_stride;
1085 const TIndex* qai = cand_idx_ptr + q * cand_stride;
1086 TIndex* qout_i = out_idx_ptr + q * out_stride;
1087 T* qout_d = out_dist_ptr + q * out_stride;
1089 int64_t ic = 0, ia = 0;
1090 for (int64_t k = 0; k <
knn; ++k) {
1091 const TIndex ci = (ic <
knn) ? qci[ic] : TIndex(-1);
1092 const TIndex ai = (ia <
knn) ? qai[ia] : TIndex(-1);
1093 if (ci < 0 && ai < 0) {
1094 qout_i[k] = TIndex(-1);
1101 }
else if (ai < 0) {
1104 const T cd = qcd[ic], ad = qad[ia];
1110 take_curr = (ci < ai);
1113 qout_d[k] = qcd[ic];
1117 qout_d[k] = qad[ia];
1125 const int64_t combined = 2 *
knn;
1126 auto policy = oneapi::dpl::execution::make_device_policy(
queue);
1127 queue.parallel_for(sycl::range<2>(num_queries, combined),
1128 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
1129 scratch_ptr[
id[0] * scratch_stride +
id[1]] =
1130 static_cast<TIndex
>(
id[1]);
1132 queue.wait_and_throw();
1134 for (int64_t qi = 0; qi < num_queries; ++qi) {
1135 TIndex* qs = scratch_ptr + qi * scratch_stride;
1136 const T* qcd = curr_dist_ptr + qi * curr_stride;
1137 const TIndex* qci = curr_idx_ptr + qi * curr_stride;
1138 const T* qad = cand_dist_ptr + qi * cand_stride;
1139 const TIndex* qai = cand_idx_ptr + qi * cand_stride;
1141 policy, qs, qs +
knn, qs + combined,
1142 [qcd, qci, qad, qai,
knn](TIndex lhs, TIndex rhs) {
1143 const bool lc = (lhs <
knn), rc = (rhs <
knn);
1144 const TIndex li = lc ? qci[lhs] : qai[lhs -
knn];
1145 const TIndex ri = rc ? qci[rhs] : qai[rhs -
knn];
1146 const T ld = lc ? qcd[lhs] : qad[lhs -
knn];
1147 const T rd = rc ? qcd[rhs] : qad[rhs -
knn];
1148 if ((li >= 0) != (ri >= 0))
return li >= 0;
1149 if (ld < rd)
return true;
1150 if (rd < ld)
return false;
1156 sycl::range<2>(num_queries,
knn),
1157 [=](sycl::id<2> id) [[intel::kernel_args_restrict]] {
1158 const int64_t qi =
id[0], k =
id[1];
1159 const TIndex src = scratch_ptr[qi * scratch_stride + k];
1160 const bool is_curr = (src <
knn);
1161 const int64_t off = is_curr ? src : src -
knn;
1163 is_curr ? curr_idx_ptr[qi * curr_stride + off]
1164 : cand_idx_ptr[qi * cand_stride + off];
1166 is_curr ? curr_dist_ptr[qi * curr_stride + off]
1167 : cand_dist_ptr[qi * cand_stride + off];
1168 TIndex* qout_i = out_idx_ptr + qi * out_stride;
1169 T* qout_d = out_dist_ptr + qi * out_stride;
1171 qout_i[k] = TIndex(-1);
1188template <
typename T,
typename TIndex>
1190 int64_t num_queries,
1192 const TIndex* indices_ptr,
1194 const T* query_norms_ptr) {
1196 queue.parallel_for(sycl::range<2>(num_queries,
knn),
1197 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
1198 const int64_t q =
id[0], k =
id[1];
1199 if (indices_ptr[q *
knn + k] < 0)
return;
1200 distances_ptr[q *
knn + k] =
1201 sycl::fmax(T(0), distances_ptr[q *
knn + k] +
1202 query_norms_ptr[q]);
T dist
Definition FixedRadiusSearchSYCLImpl.h:167
#define CALL_SELECT(Kval)
#define CALL_FINALIZE(Kval)
#define CALL_DIM(NDIMVAL)
#define CALL_DIRECT(Kval)
#define CALL_UPDATE(Kval)
int knn
Definition PointCloudSmoothing.cpp:131
sycl::queue queue
Definition SYCLContext.cpp:88
SYCL device properties and (when built) queue manager.
double t
Definition SurfaceReconstructionPoisson.cpp:175
Named kernel tag for KnnDirect (SYCL kernel naming).
Definition KnnSearchSYCLImpl.h:419
Shared types and SYCL nearest-neighbor search tuning defaults.
constexpr int64_t kKnnDirectMaxDim
Maximum point dimension compiled for DispatchKnnDirect.
Definition KnnSearchSYCLImpl.h:415
void FinalizeTopK(sycl::queue &queue, int64_t num_queries, const T *running_dist_ptr, const TIndex *running_idx_ptr, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t actual_k, const T *query_norms_ptr)
Definition KnnSearchSYCLImpl.h:254
void DispatchUpdateTopKFromTile(sycl::queue &queue, const T *neg2qp_ptr, int64_t distance_stride, const T *point_norms_ptr, int64_t num_queries, int64_t num_points, int64_t k_bucket, TIndex point_offset, T *best_dist_ptr, TIndex *best_idx_ptr, bool use_threshold, T threshold)
Instantiate UpdateTopKFromTile for the given k_bucket.
Definition KnnSearchSYCLImpl.h:312
void DispatchFinalizeTopK(sycl::queue &queue, int64_t num_queries, const T *running_dist_ptr, const TIndex *running_idx_ptr, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t actual_k, int64_t k_bucket, const T *query_norms_ptr)
Instantiate FinalizeTopK for the given k_bucket.
Definition KnnSearchSYCLImpl.h:354
constexpr int64_t kKnnDirectSubgroupSize
Default sub-group width for the direct KNN kernel (float path).
Definition KnnSearchSYCLImpl.h:409
void DispatchKnnDirectKForSG(sycl::queue &queue, const T *points_ptr, const T *queries_ptr, int64_t num_points, int64_t num_queries, int64_t actual_k, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t subgroups_per_wg, int64_t tile_points)
Definition KnnSearchSYCLImpl.h:613
void SelectTopKQueries(const Device &device, const T *distances_ptr, int64_t distance_query_stride, int64_t num_queries, int64_t num_points, int64_t knn, TIndex index_offset, TIndex *scratch_indices_ptr, int64_t scratch_query_stride, TIndex *out_indices_ptr, T *out_distances_ptr, int64_t out_query_stride, bool use_threshold=false, const T *query_norms_ptr=nullptr, T radius_sq=T(0), T scalar_threshold=T(0))
Definition KnnSearchSYCLImpl.h:955
void UpdateTopKFromTile(sycl::queue &queue, const T *neg2qp_ptr, int64_t distance_stride, const T *point_norms_ptr, int64_t num_queries, int64_t num_points, TIndex point_offset, T *best_dist_ptr, TIndex *best_idx_ptr, bool use_threshold, T threshold)
Definition KnnSearchSYCLImpl.h:187
void HeapifyDown(T *d, TIndex *idx, int root)
Definition KnnSearchSYCLImpl.h:103
void DispatchKnnDirectK(sycl::queue &queue, const T *points_ptr, const T *queries_ptr, int64_t num_points, int64_t num_queries, int64_t actual_k, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t subgroups_per_wg, int64_t tile_points)
Definition KnnSearchSYCLImpl.h:646
constexpr int64_t kKnnDirectTilePoints
Default point tile size for SLM staging.
Definition KnnSearchSYCLImpl.h:413
void DispatchSelectTopKQueries(sycl::queue &queue, const T *distances_ptr, int64_t distance_query_stride, int64_t num_queries, int64_t num_points, int64_t knn, int64_t k_bucket, TIndex index_offset, TIndex *out_indices_ptr, T *out_distances_ptr, int64_t out_query_stride, bool use_threshold, const T *query_norms_ptr, T radius_sq, T scalar_threshold)
Definition KnnSearchSYCLImpl.h:902
void MergeTopKQueries(const Device &device, const T *curr_dist_ptr, const TIndex *curr_idx_ptr, int64_t curr_stride, const T *cand_dist_ptr, const TIndex *cand_idx_ptr, int64_t cand_stride, int64_t num_queries, int64_t knn, TIndex *scratch_ptr, int64_t scratch_stride, TIndex *out_idx_ptr, T *out_dist_ptr, int64_t out_stride)
Definition KnnSearchSYCLImpl.h:1059
void HeapSort(T *d, TIndex *idx)
Heap-sort a compile-time max-heap of size K into ascending order.
Definition KnnSearchSYCLImpl.h:125
void ChooseTileSize(int64_t num_queries, int64_t num_points, int64_t element_size, int64_t tile_bytes, int64_t &tile_queries, int64_t &tile_points, int64_t max_tile_queries=128, int64_t tile_points_alignment=128)
Definition KnnSearchSYCLImpl.h:74
void KnnDirect(sycl::queue &queue, const T *points_ptr, const T *queries_ptr, int64_t num_points, int64_t num_queries, int64_t actual_k, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t subgroups_per_wg, int64_t tile_points)
Launch direct-distance KNN for fixed compile-time NDIM, K, and SG.
Definition KnnSearchSYCLImpl.h:423
void DispatchKnnDirect(sycl::queue &queue, const T *points_ptr, const T *queries_ptr, int64_t dim, int64_t num_points, int64_t num_queries, int64_t actual_k, T *out_dist_ptr, TIndex *out_idx_ptr, int64_t subgroups_per_wg=kKnnDirectSubgroupsPerWG, int64_t tile_points=kKnnDirectTilePoints)
Definition KnnSearchSYCLImpl.h:688
bool UseKnnDirect(int64_t dim, int64_t knn)
True if (dim, knn) qualifies for the direct-distance SYCL KNN path.
Definition KnnSearchSYCLImpl.h:774
int64_t KBucket(int64_t k)
Return the smallest dispatch-bucket value ≥ k.
Definition KnnSearchSYCLImpl.h:297
constexpr int64_t kKnnDirectSubgroupsPerWG
Default sub-groups per work-group (512 work-items at SG=16).
Definition KnnSearchSYCLImpl.h:411
void AddQueryNormsToDistances(const Device &device, int64_t num_queries, int64_t knn, const TIndex *indices_ptr, T *distances_ptr, const T *query_norms_ptr)
Definition KnnSearchSYCLImpl.h:1189
constexpr int64_t kSYCLKnnMidKMax
Definition NeighborSearchCommon.h:73
constexpr int64_t kSYCLKnnSmallKMax
Upper bound of k for the GRF-register heap path (eliminates scratch spill).
Definition NeighborSearchCommon.h:69
sycl::queue GetQueue(const Device &device)
Definition SYCLContext.cpp:183
Definition PinholeCameraIntrinsic.cpp:16