Open3D (C++ API)  0.19.0
Loading...
Searching...
No Matches
KnnSearchSYCLImpl.h
Go to the documentation of this file.
1// ----------------------------------------------------------------------------
2// - Open3D: www.open3d.org -
3// ----------------------------------------------------------------------------
4// Copyright (c) 2018-2024 www.open3d.org
5// SPDX-License-Identifier: MIT
6// ----------------------------------------------------------------------------
7
47
48#pragma once
49
50#include <algorithm>
51#include <limits>
52#include <oneapi/dpl/algorithm>
53#include <oneapi/dpl/execution>
54#include <sycl/sycl.hpp>
55#include <type_traits>
56
60
61namespace open3d {
62namespace core {
63namespace nns {
64
67
73inline void ChooseTileSize(int64_t num_queries,
74 int64_t num_points,
75 int64_t element_size,
76 int64_t tile_bytes,
77 int64_t& tile_queries,
78 int64_t& tile_points,
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),
84 int64_t(256));
85 if (tile_points > tile_points_alignment) {
86 tile_points =
87 (tile_points / tile_points_alignment) * tile_points_alignment;
88 }
89 tile_points = std::min<int64_t>(tile_points, num_points);
90 tile_points = std::max<int64_t>(tile_points, 1);
91}
92
94
98
101template <typename T, typename TIndex, int K>
102inline void HeapifyDown(T* d, TIndex* idx, int root) {
103 while (true) {
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])))
107 largest = left;
108 if (right < K && (d[right] > d[largest] || (d[right] == d[largest] &&
109 idx[right] > idx[largest])))
110 largest = right;
111 if (largest == root) break;
112 T td = d[root];
113 d[root] = d[largest];
114 d[largest] = td;
115 TIndex ti = idx[root];
116 idx[root] = idx[largest];
117 idx[largest] = ti;
118 root = largest;
119 }
120}
121
123template <typename T, typename TIndex, int K>
124inline void HeapSort(T* d, TIndex* idx) {
125 for (int end = K - 1; end > 0; --end) {
126 T td = d[0];
127 d[0] = d[end];
128 d[end] = td;
129 TIndex ti = idx[0];
130 idx[0] = idx[end];
131 idx[end] = ti;
132 int root = 0;
133 while (true) {
134 int left = 2 * root + 1, right = 2 * root + 2, largest = root;
135 if (left < end &&
136 (d[left] > d[largest] ||
137 (d[left] == d[largest] && idx[left] > idx[largest])))
138 largest = left;
139 if (right < end &&
140 (d[right] > d[largest] ||
141 (d[right] == d[largest] && idx[right] > idx[largest])))
142 largest = right;
143 if (largest == root) break;
144 T td2 = d[root];
145 d[root] = d[largest];
146 d[largest] = td2;
147 TIndex ti2 = idx[root];
148 idx[root] = idx[largest];
149 idx[largest] = ti2;
150 root = largest;
151 }
152 }
153}
154
156
159
185template <typename T, typename TIndex, int K>
186void UpdateTopKFromTile(sycl::queue& queue,
187 const T* neg2qp_ptr,
188 int64_t distance_stride,
189 const T* point_norms_ptr,
190 int64_t num_queries,
191 int64_t num_points,
192 TIndex point_offset,
193 T* best_dist_ptr,
194 TIndex* best_idx_ptr,
195 bool use_threshold,
196 T threshold) {
197 queue.parallel_for(
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;
204
205 // Load running best into private registers (or scratch for
206 // large K).
207 T d[K];
208 TIndex idx[K];
209 for (int i = 0; i < K; ++i) {
210 d[i] = qd[i];
211 idx[i] = qi[i];
212 }
213
214 // Scan: fused |p|² add, heap insert.
215 // Note: partial_dist = −2qp + |p|² may be negative (|q|² not
216 // yet added). Do NOT clamp here; C1 clamping is applied in
217 // FinalizeTopK / GatherWithinThresholdQueries once |q|² is
218 // added back.
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);
223 // d[0] = heap root = current k-th worst; insert if better.
224 if (dist < d[0] || (dist == d[0] && gp < idx[0])) {
225 d[0] = dist;
226 idx[0] = gp;
227 HeapifyDown<T, TIndex, K>(d, idx, 0);
228 }
229 }
230
231 for (int i = 0; i < K; ++i) {
232 qd[i] = d[i];
233 qi[i] = idx[i];
234 }
235 });
236}
237
248template <typename T, typename TIndex, int K>
249void FinalizeTopK(sycl::queue& queue,
250 int64_t num_queries,
251 const T* running_dist_ptr,
252 const TIndex* running_idx_ptr,
253 T* out_dist_ptr,
254 TIndex* out_idx_ptr,
255 int64_t actual_k,
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];
260 T d[K];
261 TIndex idx[K];
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];
265 }
266 HeapSort<T, TIndex, K>(d, idx);
267
268 const T qnorm =
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) {
273 T dist = d[i];
274 if (query_norms_ptr)
275 dist = sycl::fmax(T(0), dist + qnorm);
276 qout_d[i] = dist;
277 qout_i[i] = idx[i];
278 }
279 });
280}
281
283
287
289inline int64_t KBucket(int64_t k) {
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;
299 return 512;
300}
301
303template <typename T, typename TIndex>
305 const T* neg2qp_ptr,
306 int64_t distance_stride,
307 const T* point_norms_ptr,
308 int64_t num_queries,
309 int64_t num_points,
310 int64_t k_bucket,
311 TIndex point_offset,
312 T* best_dist_ptr,
313 TIndex* best_idx_ptr,
314 bool use_threshold,
315 T threshold) {
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)
321 if (k_bucket <= 1)
322 CALL_UPDATE(1);
323 else if (k_bucket <= 2)
324 CALL_UPDATE(2);
325 else if (k_bucket <= 4)
326 CALL_UPDATE(4);
327 else if (k_bucket <= 8)
328 CALL_UPDATE(8);
329 else if (k_bucket <= 16)
330 CALL_UPDATE(16);
331 else if (k_bucket <= 32)
332 CALL_UPDATE(32);
333 else if (k_bucket <= 64)
334 CALL_UPDATE(64);
335 else if (k_bucket <= 128)
336 CALL_UPDATE(128);
337 else if (k_bucket <= 256)
338 CALL_UPDATE(256);
339 else
340 CALL_UPDATE(512);
341#undef CALL_UPDATE
342}
343
345template <typename T, typename TIndex>
346void DispatchFinalizeTopK(sycl::queue& queue,
347 int64_t num_queries,
348 const T* running_dist_ptr,
349 const TIndex* running_idx_ptr,
350 T* out_dist_ptr,
351 TIndex* out_idx_ptr,
352 int64_t actual_k,
353 int64_t k_bucket,
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)
359 if (k_bucket <= 1)
360 CALL_FINALIZE(1);
361 else if (k_bucket <= 2)
362 CALL_FINALIZE(2);
363 else if (k_bucket <= 4)
364 CALL_FINALIZE(4);
365 else if (k_bucket <= 8)
366 CALL_FINALIZE(8);
367 else if (k_bucket <= 16)
368 CALL_FINALIZE(16);
369 else if (k_bucket <= 32)
370 CALL_FINALIZE(32);
371 else if (k_bucket <= 64)
372 CALL_FINALIZE(64);
373 else if (k_bucket <= 128)
374 CALL_FINALIZE(128);
375 else if (k_bucket <= 256)
376 CALL_FINALIZE(256);
377 else
378 CALL_FINALIZE(512);
379#undef CALL_FINALIZE
380}
381
383
399
401constexpr int64_t kKnnDirectSubgroupSize = 16;
403constexpr int64_t kKnnDirectSubgroupsPerWG = 32;
405constexpr int64_t kKnnDirectTilePoints = 2048;
407constexpr int64_t kKnnDirectMaxDim = 8;
408
410template <typename T, typename TIndex, int NDIM, int K, int SG>
412
414template <typename T, typename TIndex, int NDIM, int K, int SG>
415void KnnDirect(sycl::queue& queue,
416 const T* points_ptr,
417 const T* queries_ptr,
418 int64_t num_points,
419 int64_t num_queries,
420 int64_t actual_k,
421 T* out_dist_ptr,
422 TIndex* out_idx_ptr,
423 int64_t subgroups_per_wg,
424 int64_t tile_points) {
425 if (num_points <= 0 || num_queries <= 0) return;
426
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;
433
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);
447
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;
451
452 // Load this sub-group's query once. Inactive sub-groups
453 // (tail of the last work-group) load row 0 so every
454 // lane in the work-group stays in lock-step for the
455 // shared SLM tile loads / barriers below.
456 T q[NDIM];
457 {
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];
461 }
462 }
463
464 // Private ascending-sorted top-K, sentinel-filled.
465 T d[K];
466 TIndex idx[K];
467 for (int i = 0; i < K; ++i) {
468 d[i] = std::numeric_limits<T>::max();
469 idx[i] = TIndex(-1);
470 }
471
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);
477
478 if (t == 0) {
479 // Cooperative whole-work-group load of tile 0.
480 for (int64_t e = local_lin; e < cur_n * NDIM;
481 e += local_range) {
482 const int64_t p = e / NDIM, dd = e % NDIM;
483 slm[cur * tp * NDIM + e] =
484 points_ptr[(cur_start + p) * NDIM + dd];
485 }
486 sycl::group_barrier(it.get_group());
487 }
488
489 // Prefetch: cooperatively load the NEXT tile into
490 // the other SLM buffer before computing on the
491 // current one, so its global-memory loads are
492 // issued early and can overlap with this tile's
493 // compute below.
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;
500 e += local_range) {
501 const int64_t p = e / NDIM, dd = e % NDIM;
502 slm[nxt * tp * NDIM + e] =
503 points_ptr[(nxt_start + p) * NDIM + dd];
504 }
505 }
506
507 if (active_query) {
508 for (int64_t p_local = lane; p_local < cur_n;
509 p_local += SG) {
510 T dist = T(0);
511 const int64_t base =
512 cur * tp * NDIM + p_local * NDIM;
513 for (int dd = 0; dd < NDIM; ++dd) {
514 const T diff = q[dd] - slm[base + dd];
515 dist += diff * diff;
516 }
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])) {
521 int pos = K - 1;
522 d[pos] = dist;
523 idx[pos] = gp;
524 while (pos > 0 &&
525 (d[pos - 1] > d[pos] ||
526 (d[pos - 1] == d[pos] &&
527 idx[pos - 1] > idx[pos]))) {
528 T td = d[pos - 1];
529 d[pos - 1] = d[pos];
530 d[pos] = td;
531 TIndex ti = idx[pos - 1];
532 idx[pos - 1] = idx[pos];
533 idx[pos] = ti;
534 --pos;
535 }
536 }
537 }
538 }
539
540 // Bottom barrier: (a) the next-tile load issued
541 // above must finish before the following iteration
542 // treats it as "current"; (b) every lane must be
543 // done reading the current buffer before it is
544 // overwritten two iterations from now.
545 sycl::group_barrier(it.get_group());
546 }
547
548 if (!active_query) return;
549
550 // Sub-group all-reduce merge: after log2(SG)
551 // shuffle/merge rounds every lane holds the identical
552 // final top-K for this query, entirely register
553 // resident.
554 for (int step = 1; step < SG; step <<= 1) {
555 const int64_t partner = lane ^ step;
556 T od[K];
557 TIndex oidx[K];
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],
561 partner);
562 }
563 T md[K];
564 TIndex mi[K];
565 int a = 0, b = 0;
566 for (int o = 0; o < K; ++o) {
567 const bool take_a =
568 (b >= K) ||
569 (a < K &&
570 (d[a] < od[b] ||
571 (d[a] == od[b] && idx[a] <= oidx[b])));
572 if (take_a) {
573 md[o] = d[a];
574 mi[o] = idx[a];
575 ++a;
576 } else {
577 md[o] = od[b];
578 mi[o] = oidx[b];
579 ++b;
580 }
581 }
582 for (int o = 0; o < K; ++o) {
583 d[o] = md[o];
584 idx[o] = mi[o];
585 }
586 }
587
588 if (lane == 0) {
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]); // C1
593 oi[i] = idx[i];
594 }
595 }
596 });
597 });
598}
599
604template <typename T, typename TIndex, int NDIM, int SG>
606 const T* points_ptr,
607 const T* queries_ptr,
608 int64_t num_points,
609 int64_t num_queries,
610 int64_t actual_k,
611 T* out_dist_ptr,
612 TIndex* out_idx_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)
620 if (k_bucket <= 1)
621 CALL_DIRECT(1);
622 else if (k_bucket <= 2)
623 CALL_DIRECT(2);
624 else if (k_bucket <= 4)
625 CALL_DIRECT(4);
626 else if (k_bucket <= 8)
627 CALL_DIRECT(8);
628 else if (k_bucket <= 16)
629 CALL_DIRECT(16);
630 else
631 CALL_DIRECT(32);
632#undef CALL_DIRECT
633}
634
637template <typename T, typename TIndex, int NDIM>
638void DispatchKnnDirectK(sycl::queue& queue,
639 const T* points_ptr,
640 const T* queries_ptr,
641 int64_t num_points,
642 int64_t num_queries,
643 int64_t actual_k,
644 T* out_dist_ptr,
645 TIndex* out_idx_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 =
650 queue.get_device()
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)) !=
654 sg_sizes.end();
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,
659 tile_points);
660 } else {
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,
664 tile_points);
665 }
666 } else {
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,
670 tile_points);
671 }
672}
673
679template <typename T, typename TIndex>
680void DispatchKnnDirect(sycl::queue& queue,
681 const T* points_ptr,
682 const T* queries_ptr,
683 int64_t dim,
684 int64_t num_points,
685 int64_t num_queries,
686 int64_t actual_k,
687 T* out_dist_ptr,
688 TIndex* out_idx_ptr,
689 int64_t subgroups_per_wg = kKnnDirectSubgroupsPerWG,
690 int64_t tile_points = kKnnDirectTilePoints) {
691 // kKnnDirectTilePoints is tuned for the common case (dim ≤ 3), where the
692 // resulting per-work-group SLM usage (2 * tile_points * dim * sizeof(T))
693 // is well inside typical device budgets. For larger `dim` (up to
694 // kKnnDirectMaxDim) or double precision, that same tile_points could
695 // exceed the device's actual local memory size, so clamp it down here
696 // using the real device limit (queried once, cheap) rather than baking a
697 // dim/dtype-specific constant into the caller.
698 {
699 const size_t local_mem_bytes =
700 queue.get_device()
701 .get_info<sycl::info::device::local_mem_size>();
702 // Leave 10% headroom for other local allocations / runtime overhead.
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));
707 }
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)
712 switch (dim) {
713 case 1:
714 CALL_DIM(1);
715 break;
716 case 2:
717 CALL_DIM(2);
718 break;
719 case 3:
720 CALL_DIM(3);
721 break;
722 case 4:
723 CALL_DIM(4);
724 break;
725 case 5:
726 CALL_DIM(5);
727 break;
728 case 6:
729 CALL_DIM(6);
730 break;
731 case 7:
732 CALL_DIM(7);
733 break;
734 case 8:
735 CALL_DIM(8);
736 break;
737 default:
738 utility::LogError("DispatchKnnDirect only supports dim 1 to {}.",
740 }
741#undef CALL_DIM
742}
743
745inline bool UseKnnDirect(int64_t dim, int64_t knn) {
746 return dim >= 1 && dim <= kKnnDirectMaxDim && knn <= kSYCLKnnSmallKMax;
747}
748
750
754
755namespace {
756
758template <typename T, typename TIndex, int K>
759inline void HeapifyDownActive(T* local_d,
760 TIndex* local_i,
761 int root,
762 int active_k) {
763 int i = root;
764 while (true) {
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])))
769 largest = left;
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])))
773 largest = right;
774 if (largest == i) break;
775 T td = local_d[i];
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;
781 i = largest;
782 }
783}
784
786template <typename T, typename TIndex, int K>
787void SelectTopKQueriesHeap(sycl::queue& queue,
788 const T* distances_ptr,
789 int64_t distance_query_stride,
790 int64_t num_queries,
791 int64_t num_points,
792 int64_t knn,
793 TIndex index_offset,
794 TIndex* out_indices_ptr,
795 T* out_distances_ptr,
796 int64_t out_query_stride,
797 bool use_threshold,
798 const T* query_norms_ptr,
799 T radius_sq,
800 T scalar_threshold) {
801 const T inf = std::numeric_limits<T>::max();
802 const int64_t actual_knn = std::min(knn, num_points);
803
804 queue.parallel_for(
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;
811
812 const T thr = (use_threshold && query_norms_ptr)
813 ? (radius_sq - query_norms_ptr[q])
814 : scalar_threshold;
815
816 T local_d[K];
817 TIndex local_i[K];
818 for (int k = 0; k < actual_knn; ++k) {
819 local_d[k] = inf;
820 local_i[k] = TIndex(-1);
821 }
822
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])) {
829 local_d[0] = dist;
830 local_i[0] = p;
831 HeapifyDownActive<T, TIndex, K>(
832 local_d, local_i, 0,
833 static_cast<int>(actual_knn));
834 }
835 }
836
837 for (int i = 1; i < actual_knn; ++i) {
838 T key_d = local_d[i];
839 TIndex key_i = local_i[i];
840 int j = i - 1;
841 while (j >= 0 &&
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];
846 j--;
847 }
848 local_d[j + 1] = key_d;
849 local_i[j + 1] = key_i;
850 }
851
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);
855 qout_d[k] = inf;
856 } else {
857 qout_i[k] = index_offset + local_i[k];
858 qout_d[k] = local_d[k];
859 }
860 }
861 });
862}
863
864} // namespace
865
868template <typename T, typename TIndex>
870 const T* distances_ptr,
871 int64_t distance_query_stride,
872 int64_t num_queries,
873 int64_t num_points,
874 int64_t knn,
875 int64_t k_bucket,
876 TIndex index_offset,
877 TIndex* out_indices_ptr,
878 T* out_distances_ptr,
879 int64_t out_query_stride,
880 bool use_threshold,
881 const T* query_norms_ptr,
882 T radius_sq,
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, \
889 scalar_threshold)
890 if (k_bucket <= 1)
891 CALL_SELECT(1);
892 else if (k_bucket <= 2)
893 CALL_SELECT(2);
894 else if (k_bucket <= 4)
895 CALL_SELECT(4);
896 else if (k_bucket <= 8)
897 CALL_SELECT(8);
898 else if (k_bucket <= 16)
899 CALL_SELECT(16);
900 else if (k_bucket <= 32)
901 CALL_SELECT(32);
902 else if (k_bucket <= 64)
903 CALL_SELECT(64);
904 else if (k_bucket <= 128)
905 CALL_SELECT(128);
906 else if (k_bucket <= 256)
907 CALL_SELECT(256);
908 else
909 CALL_SELECT(512);
910#undef CALL_SELECT
911}
912
921template <typename T, typename TIndex>
922void SelectTopKQueries(const Device& device,
923 const T* distances_ptr,
924 int64_t distance_query_stride,
925 int64_t num_queries,
926 int64_t num_points,
927 int64_t knn,
928 TIndex index_offset,
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,
936 T radius_sq = T(0),
937 T scalar_threshold = T(0)) {
938 if (num_queries == 0 || num_points == 0 || knn <= 0) return;
939
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);
943
944 if (knn <= kSYCLKnnMidKMax) {
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);
951 } else {
952 // oneDPL partial_sort fallback (P8: serial per query).
953 auto policy = oneapi::dpl::execution::make_device_policy(queue);
954 queue.parallel_for(
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]);
959 });
960 queue.wait_and_throw();
961
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;
965 // C1's clamp is for the *reported* distance only (applied below,
966 // when writing qout_d). Clamping here in the comparator would
967 // tie together every point whose true partial distance is
968 // slightly negative from P2 cancellation (common for
969 // widely-spread float32 data), corrupting the selected/sorted
970 // *set* of neighbors -- not just their reported distance value.
971 // Comparing the raw (unclamped) values preserves the true
972 // relative order even when cancellation makes some values
973 // slightly negative.
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;
981 return lhs < rhs; // C4
982 });
983 }
984
985 queue.parallel_for(
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);
993 qout_d[k] = inf;
994 return;
995 }
996 const TIndex li =
997 scratch_indices_ptr[qi * scratch_query_stride + k];
998 // P2/C1: this is the *partial* distance (−2qp+|p|², |q|²
999 // not yet added by the caller). Do not clamp ≥ 0 here --
1000 // the partial value can be legitimately very negative
1001 // (missing +|q|²), especially for widely-spread float32
1002 // data; clamping it here (before |q|² is added) ties
1003 // together every such point at exactly 0, corrupting the
1004 // reported distance for many neighbors at once. The
1005 // final clamp is applied once |q|² has been added (see
1006 // AddQueryNormsToDistances / FinalizeTopK's C1).
1007 const T dist =
1008 distances_ptr[qi * distance_query_stride + li];
1009 const T thr = (use_threshold && query_norms_ptr)
1010 ? (radius_sq - query_norms_ptr[qi])
1011 : scalar_threshold;
1012 if (use_threshold && dist > thr) {
1013 qout_i[k] = TIndex(-1);
1014 qout_d[k] = inf;
1015 return;
1016 }
1017 qout_i[k] = index_offset + li;
1018 qout_d[k] = dist;
1019 });
1020 }
1021}
1022
1025template <typename T, typename TIndex>
1026void MergeTopKQueries(const Device& device,
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,
1034 int64_t knn,
1035 TIndex* scratch_ptr,
1036 int64_t scratch_stride,
1037 TIndex* out_idx_ptr,
1038 T* out_dist_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);
1043
1044 if (knn <= kSYCLKnnMidKMax) {
1045 queue.parallel_for(
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;
1055
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);
1062 qout_d[k] = inf;
1063 continue;
1064 }
1065 bool take_curr;
1066 if (ci < 0) {
1067 take_curr = false;
1068 } else if (ai < 0) {
1069 take_curr = true;
1070 } else {
1071 const T cd = qcd[ic], ad = qad[ia];
1072 if (cd < ad)
1073 take_curr = true;
1074 else if (ad < cd)
1075 take_curr = false;
1076 else
1077 take_curr = (ci < ai); // C4
1078 }
1079 if (take_curr) {
1080 qout_d[k] = qcd[ic];
1081 qout_i[k] = ci;
1082 ++ic;
1083 } else {
1084 qout_d[k] = qad[ia];
1085 qout_i[k] = ai;
1086 ++ia;
1087 }
1088 }
1089 });
1090 } else {
1091 // oneDPL merge sort fallback for large knn (P8).
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]);
1098 });
1099 queue.wait_and_throw();
1100
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;
1107 std::partial_sort(
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;
1118 return li < ri; // C4
1119 });
1120 }
1121
1122 queue.parallel_for(
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;
1129 const TIndex ii =
1130 is_curr ? curr_idx_ptr[qi * curr_stride + off]
1131 : cand_idx_ptr[qi * cand_stride + off];
1132 const T dd =
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;
1137 if (ii < 0) {
1138 qout_i[k] = TIndex(-1);
1139 qout_d[k] = inf;
1140 } else {
1141 qout_i[k] = ii;
1142 qout_d[k] = dd;
1143 }
1144 });
1145 }
1146}
1147
1149
1152
1155template <typename T, typename TIndex>
1157 int64_t num_queries,
1158 int64_t knn,
1159 const TIndex* indices_ptr,
1160 T* distances_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]);
1170 });
1171}
1172
1174
1175} // namespace nns
1176} // namespace core
1177} // namespace open3d
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
Definition Device.h:18
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