52#include <oneapi/dpl/algorithm>
53#include <oneapi/dpl/execution>
54#include <sycl/sycl.hpp>
77 int64_t& tile_queries,
79 int64_t max_tile_queries = 128,
80 int64_t tile_points_alignment = 128) {
81 tile_queries = std::min<int64_t>(num_queries, max_tile_queries);
82 tile_queries = std::max<int64_t>(tile_queries, 1);
83 tile_points = std::max<int64_t>(tile_bytes / (tile_queries * element_size),
85 if (tile_points > tile_points_alignment) {
87 (tile_points / tile_points_alignment) * tile_points_alignment;
89 tile_points = std::min<int64_t>(tile_points, num_points);
90 tile_points = std::max<int64_t>(tile_points, 1);
101template <
typename T,
typename TIndex,
int K>
104 int left = 2 * root + 1, right = 2 * root + 2, largest = root;
105 if (left < K && (d[left] > d[largest] ||
106 (d[left] == d[largest] && idx[left] > idx[largest])))
108 if (right < K && (d[right] > d[largest] || (d[right] == d[largest] &&
109 idx[right] > idx[largest])))
111 if (largest == root)
break;
113 d[root] = d[largest];
115 TIndex ti = idx[root];
116 idx[root] = idx[largest];
123template <
typename T,
typename TIndex,
int K>
125 for (
int end = K - 1; end > 0; --end) {
134 int left = 2 * root + 1, right = 2 * root + 2, largest = root;
136 (d[left] > d[largest] ||
137 (d[left] == d[largest] && idx[left] > idx[largest])))
140 (d[right] > d[largest] ||
141 (d[right] == d[largest] && idx[right] > idx[largest])))
143 if (largest == root)
break;
145 d[root] = d[largest];
147 TIndex ti2 = idx[root];
148 idx[root] = idx[largest];
185template <
typename T,
typename TIndex,
int K>
188 int64_t distance_stride,
189 const T* point_norms_ptr,
194 TIndex* best_idx_ptr,
198 sycl::range<1>(num_queries),
199 [=](sycl::id<1>
id) [[intel::kernel_args_restrict]] {
200 const int64_t q =
id[0];
201 const T* qrow = neg2qp_ptr + q * distance_stride;
202 T* qd = best_dist_ptr + q * K;
203 TIndex* qi = best_idx_ptr + q * K;
209 for (
int i = 0; i < K; ++i) {
219 for (int64_t p = 0; p < num_points; ++p) {
220 const T
dist = qrow[p] + point_norms_ptr[p];
221 if (use_threshold &&
dist > threshold)
continue;
222 const TIndex gp = point_offset +
static_cast<TIndex
>(p);
224 if (
dist < d[0] || (
dist == d[0] && gp < idx[0])) {
227 HeapifyDown<T, TIndex, K>(d, idx, 0);
231 for (
int i = 0; i < K; ++i) {
248template <
typename T,
typename TIndex,
int K>
251 const T* running_dist_ptr,
252 const TIndex* running_idx_ptr,
256 const T* query_norms_ptr) {
257 queue.parallel_for(sycl::range<1>(num_queries),
258 [=](sycl::id<1>
id) [[intel::kernel_args_restrict]] {
259 const int64_t q =
id[0];
262 for (
int i = 0; i < K; ++i) {
263 d[i] = running_dist_ptr[q * K + i];
264 idx[i] = running_idx_ptr[q * K + i];
266 HeapSort<T, TIndex, K>(d, idx);
269 query_norms_ptr ? query_norms_ptr[q] : T(0);
270 T* qout_d = out_dist_ptr + q * actual_k;
271 TIndex* qout_i = out_idx_ptr + q * actual_k;
272 for (int64_t i = 0; i < actual_k; ++i) {
275 dist = sycl::fmax(T(0),
dist + qnorm);
290 if (k <= 1)
return 1;
291 if (k <= 2)
return 2;
292 if (k <= 4)
return 4;
293 if (k <= 8)
return 8;
294 if (k <= 16)
return 16;
295 if (k <= 32)
return 32;
296 if (k <= 64)
return 64;
297 if (k <= 128)
return 128;
298 if (k <= 256)
return 256;
303template <
typename T,
typename TIndex>
306 int64_t distance_stride,
307 const T* point_norms_ptr,
313 TIndex* best_idx_ptr,
316#define CALL_UPDATE(Kval) \
317 UpdateTopKFromTile<T, TIndex, Kval>( \
318 queue, neg2qp_ptr, distance_stride, point_norms_ptr, num_queries, \
319 num_points, point_offset, best_dist_ptr, best_idx_ptr, \
320 use_threshold, threshold)
323 else if (k_bucket <= 2)
325 else if (k_bucket <= 4)
327 else if (k_bucket <= 8)
329 else if (k_bucket <= 16)
331 else if (k_bucket <= 32)
333 else if (k_bucket <= 64)
335 else if (k_bucket <= 128)
337 else if (k_bucket <= 256)
345template <
typename T,
typename TIndex>
348 const T* running_dist_ptr,
349 const TIndex* running_idx_ptr,
354 const T* query_norms_ptr) {
355#define CALL_FINALIZE(Kval) \
356 FinalizeTopK<T, TIndex, Kval>(queue, num_queries, running_dist_ptr, \
357 running_idx_ptr, out_dist_ptr, out_idx_ptr, \
358 actual_k, query_norms_ptr)
361 else if (k_bucket <= 2)
363 else if (k_bucket <= 4)
365 else if (k_bucket <= 8)
367 else if (k_bucket <= 16)
369 else if (k_bucket <= 32)
371 else if (k_bucket <= 64)
373 else if (k_bucket <= 128)
375 else if (k_bucket <= 256)
410template <
typename T,
typename TIndex,
int NDIM,
int K,
int SG>
414template <
typename T,
typename TIndex,
int NDIM,
int K,
int SG>
417 const T* queries_ptr,
423 int64_t subgroups_per_wg,
424 int64_t tile_points) {
425 if (num_points <= 0 || num_queries <= 0)
return;
427 const int64_t wg_size = subgroups_per_wg * SG;
428 const int64_t num_wgs =
429 (num_queries + subgroups_per_wg - 1) / subgroups_per_wg;
430 const int64_t global_size = num_wgs * wg_size;
431 const int64_t tp = std::min<int64_t>(tile_points, num_points);
432 const int64_t num_tiles = (num_points + tp - 1) / tp;
434 queue.submit([&](sycl::handler& h) {
435 sycl::local_accessor<T, 1> slm(sycl::range<1>(2 * tp * NDIM), h);
437 sycl::nd_range<1>(sycl::range<1>(global_size),
438 sycl::range<1>(wg_size)),
439 [=](sycl::nd_item<1> it) [[sycl::reqd_sub_group_size(
440 SG)]] [[intel::kernel_args_restrict]] {
441 const auto sg = it.get_sub_group();
442 const int64_t lane = sg.get_local_id()[0];
443 const int64_t sg_id_in_wg = sg.get_group_id()[0];
444 const int64_t wg_id = it.get_group(0);
445 const int64_t local_lin = it.get_local_linear_id();
446 const int64_t local_range = it.get_local_range(0);
448 const int64_t query_idx =
449 wg_id * subgroups_per_wg + sg_id_in_wg;
450 const bool active_query = query_idx < num_queries;
458 const int64_t qrow = active_query ? query_idx : 0;
459 for (
int d = 0; d < NDIM; ++d) {
460 q[d] = queries_ptr[qrow * NDIM + d];
467 for (
int i = 0; i < K; ++i) {
468 d[i] = std::numeric_limits<T>::max();
472 for (int64_t
t = 0;
t < num_tiles; ++
t) {
473 const int64_t cur =
t & 1;
474 const int64_t cur_start =
t * tp;
475 const int64_t cur_n =
476 std::min<int64_t>(tp, num_points - cur_start);
480 for (int64_t e = local_lin; e < cur_n * NDIM;
482 const int64_t p = e / NDIM, dd = e % NDIM;
483 slm[cur * tp * NDIM + e] =
484 points_ptr[(cur_start + p) * NDIM + dd];
486 sycl::group_barrier(it.get_group());
494 if (
t + 1 < num_tiles) {
495 const int64_t nxt = 1 - cur;
496 const int64_t nxt_start = (
t + 1) * tp;
497 const int64_t nxt_n = std::min<int64_t>(
498 tp, num_points - nxt_start);
499 for (int64_t e = local_lin; e < nxt_n * NDIM;
501 const int64_t p = e / NDIM, dd = e % NDIM;
502 slm[nxt * tp * NDIM + e] =
503 points_ptr[(nxt_start + p) * NDIM + dd];
508 for (int64_t p_local = lane; p_local < cur_n;
512 cur * tp * NDIM + p_local * NDIM;
513 for (
int dd = 0; dd < NDIM; ++dd) {
514 const T diff = q[dd] - slm[base + dd];
517 const TIndex gp =
static_cast<TIndex
>(
518 cur_start + p_local);
519 if (
dist < d[K - 1] ||
520 (
dist == d[K - 1] && gp < idx[K - 1])) {
525 (d[pos - 1] > d[pos] ||
526 (d[pos - 1] == d[pos] &&
527 idx[pos - 1] > idx[pos]))) {
531 TIndex ti = idx[pos - 1];
532 idx[pos - 1] = idx[pos];
545 sycl::group_barrier(it.get_group());
548 if (!active_query)
return;
554 for (
int step = 1; step < SG; step <<= 1) {
555 const int64_t partner = lane ^ step;
558 for (
int i = 0; i < K; ++i) {
559 od[i] = sycl::select_from_group(sg, d[i], partner);
560 oidx[i] = sycl::select_from_group(sg, idx[i],
566 for (
int o = 0; o < K; ++o) {
571 (d[a] == od[b] && idx[a] <= oidx[b])));
582 for (
int o = 0; o < K; ++o) {
589 T* od = out_dist_ptr + query_idx * actual_k;
590 TIndex* oi = out_idx_ptr + query_idx * actual_k;
591 for (int64_t i = 0; i < actual_k; ++i) {
592 od[i] = sycl::fmax(T(0), d[i]);
604template <
typename T,
typename TIndex,
int NDIM,
int SG>
607 const T* queries_ptr,
613 int64_t subgroups_per_wg,
614 int64_t tile_points) {
615 const int64_t k_bucket =
KBucket(actual_k);
616#define CALL_DIRECT(Kval) \
617 KnnDirect<T, TIndex, NDIM, Kval, SG>( \
618 queue, points_ptr, queries_ptr, num_points, num_queries, actual_k, \
619 out_dist_ptr, out_idx_ptr, subgroups_per_wg, tile_points)
622 else if (k_bucket <= 2)
624 else if (k_bucket <= 4)
626 else if (k_bucket <= 8)
628 else if (k_bucket <= 16)
637template <
typename T,
typename TIndex,
int NDIM>
640 const T* queries_ptr,
646 int64_t subgroups_per_wg,
647 int64_t tile_points) {
648 if constexpr (std::is_same_v<T, double>) {
649 const auto sg_sizes =
651 .get_info<sycl::info::device::sub_group_sizes>();
652 const bool supports_subgroup_8 =
653 std::find(sg_sizes.begin(), sg_sizes.end(),
size_t(8)) !=
655 if (supports_subgroup_8) {
656 DispatchKnnDirectKForSG<T, TIndex, NDIM, 8>(
657 queue, points_ptr, queries_ptr, num_points, num_queries,
658 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
661 DispatchKnnDirectKForSG<T, TIndex, NDIM, 16>(
662 queue, points_ptr, queries_ptr, num_points, num_queries,
663 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
667 DispatchKnnDirectKForSG<T, TIndex, NDIM, 16>(
668 queue, points_ptr, queries_ptr, num_points, num_queries,
669 actual_k, out_dist_ptr, out_idx_ptr, subgroups_per_wg,
679template <
typename T,
typename TIndex>
682 const T* queries_ptr,
699 const size_t local_mem_bytes =
701 .get_info<sycl::info::device::local_mem_size>();
703 const int64_t max_tile_points_by_slm =
static_cast<int64_t
>(
704 (local_mem_bytes * 9 / 10) / (2 * dim *
sizeof(T)));
705 tile_points = std::min(tile_points,
706 std::max<int64_t>(max_tile_points_by_slm, 1));
708#define CALL_DIM(NDIMVAL) \
709 DispatchKnnDirectK<T, TIndex, NDIMVAL>( \
710 queue, points_ptr, queries_ptr, num_points, num_queries, actual_k, \
711 out_dist_ptr, out_idx_ptr, subgroups_per_wg, tile_points)
738 utility::LogError(
"DispatchKnnDirect only supports dim 1 to {}.",
758template <
typename T,
typename TIndex,
int K>
759inline void HeapifyDownActive(T* local_d,
765 int left = 2 * i + 1, right = 2 * i + 2, largest = i;
766 if (left < active_k && (local_d[left] > local_d[largest] ||
767 (local_d[left] == local_d[largest] &&
768 local_i[left] > local_i[largest])))
770 if (right < active_k && (local_d[right] > local_d[largest] ||
771 (local_d[right] == local_d[largest] &&
772 local_i[right] > local_i[largest])))
774 if (largest == i)
break;
776 local_d[i] = local_d[largest];
777 local_d[largest] = td;
778 TIndex ti = local_i[i];
779 local_i[i] = local_i[largest];
780 local_i[largest] = ti;
786template <
typename T,
typename TIndex,
int K>
787void SelectTopKQueriesHeap(sycl::queue&
queue,
788 const T* distances_ptr,
789 int64_t distance_query_stride,
794 TIndex* out_indices_ptr,
795 T* out_distances_ptr,
796 int64_t out_query_stride,
798 const T* query_norms_ptr,
800 T scalar_threshold) {
801 const T inf = std::numeric_limits<T>::max();
802 const int64_t actual_knn = std::min(knn, num_points);
805 sycl::range<1>(num_queries),
806 [=](sycl::id<1>
id) [[intel::kernel_args_restrict]] {
807 const int64_t q =
id[0];
808 const T* qd = distances_ptr + q * distance_query_stride;
809 TIndex* qout_i = out_indices_ptr + q * out_query_stride;
810 T* qout_d = out_distances_ptr + q * out_query_stride;
812 const T thr = (use_threshold && query_norms_ptr)
813 ? (radius_sq - query_norms_ptr[q])
818 for (
int k = 0; k < actual_knn; ++k) {
820 local_i[k] = TIndex(-1);
823 for (TIndex p = 0; p < static_cast<TIndex>(num_points); ++p) {
824 const T
dist = qd[p];
825 if (use_threshold &&
dist > thr)
continue;
826 if (
dist < local_d[0] ||
827 (
dist == local_d[0] &&
828 index_offset + p < index_offset + local_i[0])) {
831 HeapifyDownActive<T, TIndex, K>(
833 static_cast<int>(actual_knn));
837 for (
int i = 1; i < actual_knn; ++i) {
838 T key_d = local_d[i];
839 TIndex key_i = local_i[i];
842 (local_d[j] > key_d ||
843 (local_d[j] == key_d && local_i[j] > key_i))) {
844 local_d[j + 1] = local_d[j];
845 local_i[j + 1] = local_i[j];
848 local_d[j + 1] = key_d;
849 local_i[j + 1] = key_i;
852 for (int64_t k = 0; k < knn; ++k) {
853 if (k >= actual_knn || local_i[k] == TIndex(-1)) {
854 qout_i[k] = TIndex(-1);
857 qout_i[k] = index_offset + local_i[k];
858 qout_d[k] = local_d[k];
868template <
typename T,
typename TIndex>
870 const T* distances_ptr,
871 int64_t distance_query_stride,
877 TIndex* out_indices_ptr,
878 T* out_distances_ptr,
879 int64_t out_query_stride,
881 const T* query_norms_ptr,
883 T scalar_threshold) {
884#define CALL_SELECT(Kval) \
885 SelectTopKQueriesHeap<T, TIndex, Kval>( \
886 queue, distances_ptr, distance_query_stride, num_queries, \
887 num_points, knn, index_offset, out_indices_ptr, out_distances_ptr, \
888 out_query_stride, use_threshold, query_norms_ptr, radius_sq, \
892 else if (k_bucket <= 2)
894 else if (k_bucket <= 4)
896 else if (k_bucket <= 8)
898 else if (k_bucket <= 16)
900 else if (k_bucket <= 32)
902 else if (k_bucket <= 64)
904 else if (k_bucket <= 128)
906 else if (k_bucket <= 256)
921template <
typename T,
typename TIndex>
923 const T* distances_ptr,
924 int64_t distance_query_stride,
929 TIndex* scratch_indices_ptr,
930 int64_t scratch_query_stride,
931 TIndex* out_indices_ptr,
932 T* out_distances_ptr,
933 int64_t out_query_stride,
934 bool use_threshold =
false,
935 const T* query_norms_ptr =
nullptr,
937 T scalar_threshold = T(0)) {
938 if (num_queries == 0 || num_points == 0 || knn <= 0)
return;
940 const T inf = std::numeric_limits<T>::max();
941 const int64_t actual_knn = std::min(knn, num_points);
942 sycl::queue
queue = sy::SYCLContext::GetInstance().GetDefaultQueue(device);
945 const int64_t k_bucket =
KBucket(knn);
946 DispatchSelectTopKQueries<T, TIndex>(
947 queue, distances_ptr, distance_query_stride, num_queries,
948 num_points, knn, k_bucket, index_offset, out_indices_ptr,
949 out_distances_ptr, out_query_stride, use_threshold,
950 query_norms_ptr, radius_sq, scalar_threshold);
953 auto policy = oneapi::dpl::execution::make_device_policy(
queue);
955 sycl::range<2>(num_queries, num_points),
956 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
957 scratch_indices_ptr[
id[0] * scratch_query_stride +
id[1]] =
958 static_cast<TIndex
>(
id[1]);
960 queue.wait_and_throw();
962 for (int64_t qi = 0; qi < num_queries; ++qi) {
963 TIndex* q_scratch = scratch_indices_ptr + qi * scratch_query_stride;
964 const T* q_dist = distances_ptr + qi * distance_query_stride;
974 std::partial_sort(policy, q_scratch, q_scratch + actual_knn,
975 q_scratch + num_points,
976 [q_dist](TIndex lhs, TIndex rhs) {
977 const T ld = q_dist[lhs];
978 const T rd = q_dist[rhs];
979 if (ld < rd)
return true;
980 if (rd < ld)
return false;
986 sycl::range<2>(num_queries, knn),
987 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
988 const int64_t qi =
id[0], k =
id[1];
989 TIndex* qout_i = out_indices_ptr + qi * out_query_stride;
990 T* qout_d = out_distances_ptr + qi * out_query_stride;
991 if (k >= actual_knn) {
992 qout_i[k] = TIndex(-1);
997 scratch_indices_ptr[qi * scratch_query_stride + k];
1008 distances_ptr[qi * distance_query_stride + li];
1009 const T thr = (use_threshold && query_norms_ptr)
1010 ? (radius_sq - query_norms_ptr[qi])
1012 if (use_threshold &&
dist > thr) {
1013 qout_i[k] = TIndex(-1);
1017 qout_i[k] = index_offset + li;
1025template <
typename T,
typename TIndex>
1027 const T* curr_dist_ptr,
1028 const TIndex* curr_idx_ptr,
1029 int64_t curr_stride,
1030 const T* cand_dist_ptr,
1031 const TIndex* cand_idx_ptr,
1032 int64_t cand_stride,
1033 int64_t num_queries,
1035 TIndex* scratch_ptr,
1036 int64_t scratch_stride,
1037 TIndex* out_idx_ptr,
1039 int64_t out_stride) {
1040 if (num_queries == 0 || knn <= 0)
return;
1041 const T inf = std::numeric_limits<T>::max();
1042 sycl::queue
queue = sy::SYCLContext::GetInstance().GetDefaultQueue(device);
1046 sycl::range<1>(num_queries),
1047 [=](sycl::id<1>
id) [[intel::kernel_args_restrict]] {
1048 const int64_t q =
id[0];
1049 const T* qcd = curr_dist_ptr + q * curr_stride;
1050 const TIndex* qci = curr_idx_ptr + q * curr_stride;
1051 const T* qad = cand_dist_ptr + q * cand_stride;
1052 const TIndex* qai = cand_idx_ptr + q * cand_stride;
1053 TIndex* qout_i = out_idx_ptr + q * out_stride;
1054 T* qout_d = out_dist_ptr + q * out_stride;
1056 int64_t ic = 0, ia = 0;
1057 for (int64_t k = 0; k < knn; ++k) {
1058 const TIndex ci = (ic < knn) ? qci[ic] : TIndex(-1);
1059 const TIndex ai = (ia < knn) ? qai[ia] : TIndex(-1);
1060 if (ci < 0 && ai < 0) {
1061 qout_i[k] = TIndex(-1);
1068 }
else if (ai < 0) {
1071 const T cd = qcd[ic], ad = qad[ia];
1077 take_curr = (ci < ai);
1080 qout_d[k] = qcd[ic];
1084 qout_d[k] = qad[ia];
1092 const int64_t combined = 2 * knn;
1093 auto policy = oneapi::dpl::execution::make_device_policy(
queue);
1094 queue.parallel_for(sycl::range<2>(num_queries, combined),
1095 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
1096 scratch_ptr[
id[0] * scratch_stride +
id[1]] =
1097 static_cast<TIndex
>(
id[1]);
1099 queue.wait_and_throw();
1101 for (int64_t qi = 0; qi < num_queries; ++qi) {
1102 TIndex* qs = scratch_ptr + qi * scratch_stride;
1103 const T* qcd = curr_dist_ptr + qi * curr_stride;
1104 const TIndex* qci = curr_idx_ptr + qi * curr_stride;
1105 const T* qad = cand_dist_ptr + qi * cand_stride;
1106 const TIndex* qai = cand_idx_ptr + qi * cand_stride;
1108 policy, qs, qs + knn, qs + combined,
1109 [qcd, qci, qad, qai, knn](TIndex lhs, TIndex rhs) {
1110 const bool lc = (lhs < knn), rc = (rhs < knn);
1111 const TIndex li = lc ? qci[lhs] : qai[lhs - knn];
1112 const TIndex ri = rc ? qci[rhs] : qai[rhs - knn];
1113 const T ld = lc ? qcd[lhs] : qad[lhs - knn];
1114 const T rd = rc ? qcd[rhs] : qad[rhs - knn];
1115 if ((li >= 0) != (ri >= 0))
return li >= 0;
1116 if (ld < rd)
return true;
1117 if (rd < ld)
return false;
1123 sycl::range<2>(num_queries, knn),
1124 [=](sycl::id<2> id) [[intel::kernel_args_restrict]] {
1125 const int64_t qi =
id[0], k =
id[1];
1126 const TIndex src = scratch_ptr[qi * scratch_stride + k];
1127 const bool is_curr = (src < knn);
1128 const int64_t off = is_curr ? src : src - knn;
1130 is_curr ? curr_idx_ptr[qi * curr_stride + off]
1131 : cand_idx_ptr[qi * cand_stride + off];
1133 is_curr ? curr_dist_ptr[qi * curr_stride + off]
1134 : cand_dist_ptr[qi * cand_stride + off];
1135 TIndex* qout_i = out_idx_ptr + qi * out_stride;
1136 T* qout_d = out_dist_ptr + qi * out_stride;
1138 qout_i[k] = TIndex(-1);
1155template <
typename T,
typename TIndex>
1157 int64_t num_queries,
1159 const TIndex* indices_ptr,
1161 const T* query_norms_ptr) {
1162 sycl::queue
queue = sy::SYCLContext::GetInstance().GetDefaultQueue(device);
1163 queue.parallel_for(sycl::range<2>(num_queries, knn),
1164 [=](sycl::id<2>
id) [[intel::kernel_args_restrict]] {
1165 const int64_t q =
id[0], k =
id[1];
1166 if (indices_ptr[q * knn + k] < 0)
return;
1167 distances_ptr[q * knn + k] =
1168 sycl::fmax(T(0), distances_ptr[q * knn + k] +
1169 query_norms_ptr[q]);
T dist
Definition FixedRadiusSearchSYCLImpl.h:163
#define CALL_SELECT(Kval)
#define CALL_FINALIZE(Kval)
#define CALL_DIM(NDIMVAL)
#define CALL_DIRECT(Kval)
#define CALL_UPDATE(Kval)
sycl::queue queue
Definition SYCLContext.cpp:51
SYCL device properties and (when built) queue manager.
double t
Definition SurfaceReconstructionPoisson.cpp:172
Named kernel tag for KnnDirect (SYCL kernel naming).
Definition KnnSearchSYCLImpl.h:411
Shared types and SYCL nearest-neighbor search tuning defaults.
constexpr int64_t kKnnDirectMaxDim
Maximum point dimension compiled for DispatchKnnDirect.
Definition KnnSearchSYCLImpl.h:407
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:249
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:304
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:346
constexpr int64_t kKnnDirectSubgroupSize
Default sub-group width for the direct KNN kernel (float path).
Definition KnnSearchSYCLImpl.h:401
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:605
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:922
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:186
void HeapifyDown(T *d, TIndex *idx, int root)
Definition KnnSearchSYCLImpl.h:102
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:638
constexpr int64_t kKnnDirectTilePoints
Default point tile size for SLM staging.
Definition KnnSearchSYCLImpl.h:405
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:869
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:1026
void HeapSort(T *d, TIndex *idx)
Heap-sort a compile-time max-heap of size K into ascending order.
Definition KnnSearchSYCLImpl.h:124
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:73
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:415
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:680
bool UseKnnDirect(int64_t dim, int64_t knn)
True if (dim, knn) qualifies for the direct-distance SYCL KNN path.
Definition KnnSearchSYCLImpl.h:745
int64_t KBucket(int64_t k)
Return the smallest dispatch-bucket value ≥ k.
Definition KnnSearchSYCLImpl.h:289
constexpr int64_t kKnnDirectSubgroupsPerWG
Default sub-groups per work-group (512 work-items at SG=16).
Definition KnnSearchSYCLImpl.h:403
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:1156
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
Definition PinholeCameraIntrinsic.cpp:16