From d1ebaf81228d520309e9d31e41b53dfbc0446691 Mon Sep 17 00:00:00 2001 From: Nikita Grigorian Date: Tue, 29 Sep 2026 10:46:22 -0700 Subject: [PATCH] add radix selection algorithm a few variants: one which handles multiple rows per wg, one which handles one row per wg, one which handles multiple wgs per row --- .../include/kernels/sorting/radix_select.hpp | 1038 +++++++++++++++++ .../include/kernels/sorting/radix_sort.hpp | 263 +---- .../include/kernels/sorting/radix_utils.hpp | 371 ++++++ .../include/kernels/sorting/topk.hpp | 27 + dpnp/tensor/libtensor/source/sorting/topk.cpp | 29 +- dpnp/tests/tensor/test_usm_ndarray_top_k.py | 92 ++ 6 files changed, 1548 insertions(+), 272 deletions(-) create mode 100644 dpnp/tensor/libtensor/include/kernels/sorting/radix_select.hpp create mode 100644 dpnp/tensor/libtensor/include/kernels/sorting/radix_utils.hpp diff --git a/dpnp/tensor/libtensor/include/kernels/sorting/radix_select.hpp b/dpnp/tensor/libtensor/include/kernels/sorting/radix_select.hpp new file mode 100644 index 000000000000..f5dbe8f41bea --- /dev/null +++ b/dpnp/tensor/libtensor/include/kernels/sorting/radix_select.hpp @@ -0,0 +1,1038 @@ +//***************************************************************************** +// Copyright (c) 2026, Intel Corporation +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions are met: +// - Redistributions of source code must retain the above copyright notice, +// this list of conditions and the following disclaimer. +// - Redistributions in binary form must reproduce the above copyright notice, +// this list of conditions and the following disclaimer in the documentation +// and/or other materials provided with the distribution. +// - Neither the name of the copyright holder nor the names of its contributors +// may be used to endorse or promote products derived from this software +// without specific prior written permission. +// +// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +// LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +// CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +// SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +// CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +// ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF +// THE POSSIBILITY OF SUCH DAMAGE. +//***************************************************************************** +// +//===----------------------------------------------------------------------===// +/// +/// \file +/// This file defines radix select kernels for tensor topk operation. +//===----------------------------------------------------------------------===// + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "kernels/sorting/radix_utils.hpp" +#include "utils/sycl_alloc_utils.hpp" + +namespace dpnp::tensor::kernels::radix_select_details +{ + +inline constexpr std::uint32_t radix_bits = 8; +inline constexpr std::uint32_t radix_states = std::uint32_t(1) << radix_bits; +inline constexpr std::uint32_t radix_mask = radix_states - 1; + +/*! @brief Number of radix passes which resolve a key of type `KeyT` */ +template +constexpr std::uint32_t n_radix_passes() +{ + return radix_utils::number_of_buckets_in_type(radix_bits); +} + +/*! @brief Bucket holding the element of a given rank */ +struct bucket_info +{ + std::uint32_t bucket_id; + // number of elements in buckets preceding `bucket_id` + std::uint64_t before; + // number of elements in `bucket_id` + std::uint64_t count; +}; + +//----------------------------------------------------------------------- +// radix select: work-group level building blocks +//----------------------------------------------------------------------- + +/*! @brief Adds to the zeroed local histogram `hist` the digits at `shift` of + * the keys in `[begin, end)` whose prefix under `mask` is `desired` */ +template +void count_digits(const sycl::nd_item<1> &ndit, + const KeyFnT &key_fn, + std::size_t begin, + std::size_t end, + KeyT desired, + KeyT mask, + std::uint32_t shift, + const HistAccT &hist) +{ + using AtomicT = sycl::atomic_ref; + + const auto &sg = ndit.get_sub_group(); + const std::uint32_t sg_size = sg.get_local_linear_range(); + const std::size_t lid = ndit.get_local_linear_id(); + const std::size_t wg_size = ndit.get_local_range(0); + + for (std::size_t i0 = begin; i0 < end; i0 += wg_size) { + const std::size_t i = i0 + lid; + + bool match = false; + std::uint32_t bucket_id = 0; + if (i < end) { + const KeyT key = key_fn(i); + match = ((key & mask) == desired); + bucket_id = radix_utils::get_bucket_id(key, shift); + } + + // a sub-group whose keys all land in one bucket, as for runs of + // equal values, makes one update instead of contending sg_size times + const std::uint32_t leader_bucket_id = + sycl::group_broadcast(sg, bucket_id); + if (sycl::all_of_group(sg, match && (bucket_id == leader_bucket_id))) { + if (sg.leader()) { + AtomicT(hist[bucket_id]).fetch_add(sg_size); + } + } + else if (match) { + AtomicT(hist[bucket_id]).fetch_add(std::uint32_t(1)); + } + } +} + +/*! @brief Finds the bucket of the complete local histogram `hist` holding + * the element of 1-based rank `rank` */ +template +bucket_info find_bucket(const sycl::nd_item<1> &ndit, + const HistAccT &hist, + std::uint64_t rank, + const ResultAccT &result) +{ + const std::uint32_t lid = ndit.get_local_linear_id(); + const std::uint32_t wg_size = ndit.get_local_range(0); + + const std::uint32_t buckets_per_wi = (radix_states + wg_size - 1) / wg_size; + const std::uint32_t bucket_begin = + std::min(lid * buckets_per_wi, radix_states); + const std::uint32_t bucket_end = + std::min(bucket_begin + buckets_per_wi, radix_states); + + std::uint64_t wi_count = 0; + for (std::uint32_t b = bucket_begin; b < bucket_end; ++b) { + wi_count += hist[b]; + } + + std::uint64_t prefix = sycl::exclusive_scan_over_group( + ndit.get_group(), wi_count, sycl::plus()); + + // exactly one work-item owns the bucket in which the prefix crosses rank + for (std::uint32_t b = bucket_begin; b < bucket_end; ++b) { + const std::uint64_t c = hist[b]; + if (prefix < rank && rank <= prefix + c) { + result[0] = bucket_info{b, prefix, c}; + } + prefix += c; + } + + sycl::group_barrier(ndit.get_group()); + + return result[0]; +} + +/*! @brief Passes the elements `i` of `[begin, end)` whose prefix under `mask` + * is less than `desired` to `less_out(j, i)`, and the first `n_ties` of the + * row equal to it to `tie_out(j, i)`, `j` being their rank in index order */ +template +void gather_selected(const sycl::nd_item<1> &ndit, + const KeyFnT &key_fn, + std::size_t begin, + std::size_t end, + KeyT desired, + KeyT mask, + std::uint64_t less_offset, // rank of first less + std::uint64_t n_less, // less in range + std::uint64_t n_ties_before, // ties before begin + std::uint64_t n_ties, + const LessOutT &less_out, + const TieOutT &tie_out) +{ + // an element's kind and its rank within a tile are counted in the low + // (less) and high (tie) halves of one integer + static constexpr std::uint32_t less_flag = 1; + static constexpr std::uint32_t tie_flag = std::uint32_t(1) << 16; + static constexpr std::uint32_t half_mask = tie_flag - 1; + + const auto &wg = ndit.get_group(); + const auto &sg = ndit.get_sub_group(); + const std::uint32_t lane_id = sg.get_local_linear_id(); + const std::uint32_t sg_size = sg.get_local_linear_range(); + const std::size_t lid = ndit.get_local_linear_id(); + const std::size_t wg_size = ndit.get_local_range(0); + + // each sub-group processes a contiguous piece of the tile, so that + // elements are ranked in index order + const std::size_t tile_size = wg_size * elems_per_wi; + const std::size_t sg_tile_offset = (lid - lane_id) * elems_per_wi; + + std::uint64_t less_done = 0; + std::uint64_t ties_done = 0; + for (std::size_t t0 = begin; t0 < end; t0 += tile_size) { + if (less_done >= n_less && n_ties_before + ties_done >= n_ties) { + break; + } + + const std::size_t sg_begin = t0 + sg_tile_offset; + + std::uint32_t flags[elems_per_wi]; + std::uint32_t ranks[elems_per_wi]; + std::uint32_t sg_count = 0; +#pragma unroll + for (std::uint32_t j = 0; j < elems_per_wi; ++j) { + const std::size_t i = sg_begin + j * sg_size + lane_id; + + std::uint32_t f = 0; + if (i < end) { + const KeyT prefix = key_fn(i) & mask; + f = (prefix < desired) ? less_flag + : (prefix == desired) ? tie_flag + : 0; + } + const std::uint32_t pre = sycl::exclusive_scan_over_group( + sg, f, sycl::plus()); + flags[j] = f; + ranks[j] = sg_count + pre; + sg_count += + sycl::reduce_over_group(sg, f, sycl::plus()); + } + + // offset of the sub-group's piece within the tile + const std::uint32_t contrib = (lane_id == 0) ? sg_count : 0; + const std::uint32_t sg_offset = sycl::group_broadcast( + sg, sycl::exclusive_scan_over_group(wg, contrib, + sycl::plus())); + const std::uint32_t tile_count = + sycl::reduce_over_group(wg, contrib, sycl::plus()); + +#pragma unroll + for (std::uint32_t j = 0; j < elems_per_wi; ++j) { + if (flags[j] == 0) { + continue; + } + const std::size_t i = sg_begin + j * sg_size + lane_id; + const std::uint32_t r = sg_offset + ranks[j]; + if (flags[j] == less_flag) { + less_out(less_offset + less_done + (r & half_mask), i); + } + else { + const std::uint64_t tie_rank = + n_ties_before + ties_done + (r >> 16); + if (tie_rank < n_ties) { + tie_out(tie_rank, i); + } + } + } + + less_done += (tile_count & half_mask); + ties_done += (tile_count >> 16); + } +} + +template +struct RowKey +{ + const ValueT *arg; + + KeyT operator()(std::size_t i) const + { + return radix_utils::ordered_radix_key(arg[i]); + } +}; + +/*! @brief Key of candidate `i` of a row, candidates listed by row index */ +template +struct CandidateKey +{ + const ValueT *arg; + const std::uint32_t *cand; + + KeyT operator()(std::size_t i) const + { + return radix_utils::ordered_radix_key(arg[cand[i]]); + } +}; + +/*! @brief Stores the value and the index of element `i` of a row as the + * `j`-th selected element of the row */ +template +struct SelectedOut +{ + const ValueT *arg; + ValueT *vals; + IndexT *inds; + + void operator()(std::uint64_t j, std::size_t i) const + { + vals[j] = arg[i]; + inds[j] = static_cast(i); + } +}; + +/*! @brief Stores the value and the index of candidate `i` of a row as the + * `j`-th selected element of the row */ +template +struct CandidateSelectedOut +{ + const ValueT *arg; + const std::uint32_t *cand; + ValueT *vals; + IndexT *inds; + + void operator()(std::uint64_t j, std::size_t i) const + { + const std::uint32_t c = cand[i]; + vals[j] = arg[c]; + inds[j] = static_cast(c); + } +}; + +/*! @brief Lists element `i` of a row as its `j`-th candidate */ +struct CandidateListOut +{ + std::uint32_t *cand; + + void operator()(std::uint64_t j, std::size_t i) const + { + cand[j] = static_cast(i); + } +}; + +//----------------------------------------------------------------------- +// radix select: one work-group per row +//----------------------------------------------------------------------- + +template +class radix_select_one_group_krn; + +/*! @brief Largest tile size `gather_selected` can rank within 16 bits */ +inline constexpr std::size_t max_tile_size = (std::size_t(1) << 16) - 1; + +template +sycl::event + radix_select_one_group_submit(sycl::queue &exec_q, + std::size_t n_iters, + std::size_t n_values, + std::size_t k, + const ValueT *arg_ptr, + ValueT *vals_ptr, + IndexT *inds_ptr, + std::size_t wg_size, + const std::vector &depends) +{ + using KeyT = radix_utils::radix_key_t; + using KernelName = + radix_select_one_group_krn; + + static constexpr std::uint32_t n_passes = n_radix_passes(); + static constexpr std::uint32_t key_bits = + radix_utils::number_of_bits_in_type(); + + if (n_values > std::numeric_limits::max() || + wg_size * elems_per_wi > max_tile_size) { + throw std::runtime_error("Invalid parameters for radix select"); + } + + return exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(depends); + + sycl::local_accessor hist(radix_states, cgh); + sycl::local_accessor bucket(1, cgh); + + sycl::nd_range<1> ndRange(n_iters * wg_size, wg_size); + + cgh.parallel_for(ndRange, [=](sycl::nd_item<1> ndit) { + const std::size_t iter_id = ndit.get_group(0); + const std::size_t lid = ndit.get_local_linear_id(); + + const RowKey key_fn{arg_ptr + + iter_id * n_values}; + + KeyT desired{0}; + KeyT mask{0}; + std::uint64_t k_rem = k; + for (std::uint32_t pass = 0; pass < n_passes; ++pass) { + const std::uint32_t shift = key_bits - (pass + 1) * radix_bits; + + for (std::size_t b = lid; b < radix_states; b += wg_size) { + hist[b] = 0; + } + sycl::group_barrier(ndit.get_group()); + + count_digits(ndit, key_fn, 0, n_values, desired, mask, shift, + hist); + sycl::group_barrier(ndit.get_group()); + + const bucket_info info = find_bucket(ndit, hist, k_rem, bucket); + + desired |= static_cast(KeyT(info.bucket_id) << shift); + mask |= static_cast(KeyT(radix_mask) << shift); + k_rem -= info.before; + + // all keys with the prefix are selected, the rest of the + // digits does not matter + if (info.count == k_rem) { + break; + } + } + + const std::uint64_t n_less = k - k_rem; + using OutT = SelectedOut; + const ValueT *row_arg = arg_ptr + iter_id * n_values; + ValueT *row_vals = vals_ptr + iter_id * k; + IndexT *row_inds = inds_ptr + iter_id * k; + gather_selected( + ndit, key_fn, 0, n_values, desired, mask, 0, n_less, 0, k_rem, + OutT{row_arg, row_vals, row_inds}, + OutT{row_arg, row_vals + n_less, row_inds + n_less}); + }); + }); +} + +//----------------------------------------------------------------------- +// radix select: one sub-group per row +//----------------------------------------------------------------------- + +template +class radix_select_sub_group_krn; + +/*! @brief Selection with a sub-group per row of at most `max_chunks` times the + * sub-group size, each element is ranked by counting the elements ordering + * before it and written at its rank, so the selection comes out sorted */ +template +sycl::event + radix_select_sub_group_submit(sycl::queue &exec_q, + std::size_t n_iters, + std::size_t n_values, + std::size_t k, + const ValueT *arg_ptr, + ValueT *vals_ptr, + IndexT *inds_ptr, + std::size_t n_groups, + std::size_t wg_size, + const std::vector &depends) +{ + using KeyT = radix_utils::radix_key_t; + using KernelName = + radix_select_sub_group_krn; + + return exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(depends); + + sycl::nd_range<1> ndRange(n_groups * wg_size, wg_size); + + cgh.parallel_for(ndRange, [=](sycl::nd_item<1> ndit) { + const auto &sg = ndit.get_sub_group(); + const std::uint32_t lane_id = sg.get_local_linear_id(); + const std::uint32_t sg_size = sg.get_local_linear_range(); + const std::size_t sgs_per_group = sg.get_group_linear_range(); + const std::size_t n_sub_groups = + ndit.get_group_range(0) * sgs_per_group; + const std::size_t sg_id = + ndit.get_group(0) * sgs_per_group + sg.get_group_linear_id(); + + const std::uint32_t n = static_cast(n_values); + const std::uint32_t n_chunks = (n + sg_size - 1) / sg_size; + + // rows are strided over sub-groups, so any sub-group size of at + // least n_values / max_chunks covers all of them + for (std::size_t iter_id = sg_id; iter_id < n_iters; + iter_id += n_sub_groups) { + const ValueT *row_arg = arg_ptr + iter_id * n_values; + + // element c * sg_size + lane_id of the row + KeyT keys[max_chunks]; + std::uint32_t ranks[max_chunks]; +#pragma unroll + for (std::uint32_t c = 0; c < max_chunks; ++c) { + const std::uint32_t i = c * sg_size + lane_id; + keys[c] = + (i < n) ? radix_utils::ordered_radix_key( + row_arg[i]) + : KeyT{0}; + ranks[c] = 0; + } + + // the chunks are unrolled so that the keys stay in registers, + // the loops over chunks past the row's end are skipped by the + // whole sub-group +#pragma unroll + for (std::uint32_t cj = 0; cj < max_chunks; ++cj) { + if (cj >= n_chunks) { + break; + } + const std::uint32_t j0 = cj * sg_size; + const std::uint32_t l_end = std::min(sg_size, n - j0); + for (std::uint32_t l = 0; l < l_end; ++l) { + const KeyT kj = sycl::group_broadcast(sg, keys[cj], l); + const std::uint32_t j = j0 + l; +#pragma unroll + for (std::uint32_t c = 0; c < max_chunks; ++c) { + if (c < n_chunks) { + const std::uint32_t i = c * sg_size + lane_id; + ranks[c] += + (kj < keys[c]) || (kj == keys[c] && j < i); + } + } + } + } + + ValueT *row_vals = vals_ptr + iter_id * k; + IndexT *row_inds = inds_ptr + iter_id * k; +#pragma unroll + for (std::uint32_t c = 0; c < max_chunks; ++c) { + const std::uint32_t i = c * sg_size + lane_id; + if (i < n && ranks[c] < k) { + row_vals[ranks[c]] = row_arg[i]; + row_inds[ranks[c]] = static_cast(i); + } + } + } + }); + }); +} + +//----------------------------------------------------------------------- +// radix select: several work-groups per row +//----------------------------------------------------------------------- + +/*! @brief Selection state of a row, carried from pass to pass */ +template +struct row_state +{ + KeyT desired; + KeyT mask; + // rank of the sought element among the keys having the prefix + std::uint64_t k_rem; + // when nonzero, the passes after the first one go over this many + // candidates, the elements which share the prefix of the first pass + std::uint64_t n_cand; + // number of elements selected by the first pass + std::uint64_t n_less_first; + std::uint32_t done; + // number of work-groups of the row which completed the current pass + std::uint32_t n_arrived; +}; + +template +class radix_select_init_krn; + +template +class radix_select_count_krn; + +template +class radix_select_filter_krn; + +template +class radix_select_gather_krn; + +/*! @brief Range `[begin, end)` of segment `segment_id` of `n_segments` + * covering `[0, nelems)` */ +inline std::pair segment_range(std::size_t segment_id, + std::size_t n_segments, + std::size_t nelems) +{ + const std::size_t elems_per_segment = + (nelems + n_segments - 1) / n_segments; + const std::size_t begin = std::min(segment_id * elems_per_segment, nelems); + return {begin, std::min(begin + elems_per_segment, nelems)}; +} + +// candidates are listed only when they are at most this fraction of their +// row, a larger buffer costs more to allocate than it saves +inline constexpr std::size_t filter_cap_divisor = 16; + +/*! @brief Radix select with `n_segments` work-groups per row and a kernel per + * pass, the last work-group of a row to finish a pass resolves its digit */ +template +sycl::event + radix_select_multi_group_impl(sycl::queue &exec_q, + std::size_t n_iters, + std::size_t n_values, + std::size_t k, + const ValueT *arg_ptr, + ValueT *vals_ptr, + IndexT *inds_ptr, + std::size_t n_segments, + std::size_t wg_size, + bool filter, + const std::vector &depends) +{ + using KeyT = radix_utils::radix_key_t; + using StateT = row_state; + using RowKeyT = RowKey; + using CandKeyT = CandidateKey; + + static constexpr std::uint32_t n_passes = n_radix_passes(); + static constexpr std::uint32_t key_bits = + radix_utils::number_of_bits_in_type(); + + const std::size_t elems_per_segment = + (n_values + n_segments - 1) / n_segments; + if (elems_per_segment > std::numeric_limits::max() || + n_segments > std::numeric_limits::max() || + wg_size * elems_per_wi > max_tile_size) { + throw std::runtime_error("Invalid parameters for radix select"); + } + // when filtering, the elements sharing the prefix of the first pass are + // listed, by 32-bit indices within their row, for the later passes to go + // over + filter = filter && (n_passes > 1) && + (n_values <= std::numeric_limits::max()); + // rows with more candidates than this are not filtered + const std::size_t cand_cap = + (n_values + filter_cap_divisor - 1) / filter_cap_divisor; + + const std::size_t n_all_segments = n_iters * n_segments; + + auto state_owner = + dpnp::tensor::alloc_utils::smart_malloc_device(n_iters, exec_q); + StateT *state_ptr = state_owner.get(); + + // per work-group digit histograms of the current pass + auto hist_owner = + dpnp::tensor::alloc_utils::smart_malloc_device( + n_all_segments * radix_states, exec_q); + std::uint32_t *segment_hist_ptr = hist_owner.get(); + + // per work-group counts of less and tie elements, and their exclusive + // scans over the work-groups of a row + auto counts_owner = + dpnp::tensor::alloc_utils::smart_malloc_device( + 4 * n_all_segments, exec_q); + std::uint64_t *less_count_ptr = counts_owner.get(); + std::uint64_t *less_offset_ptr = less_count_ptr + n_all_segments; + std::uint64_t *tie_count_ptr = less_offset_ptr + n_all_segments; + std::uint64_t *tie_offset_ptr = tie_count_ptr + n_all_segments; + + // candidates of each row, in index order + auto cand_owner = + dpnp::tensor::alloc_utils::smart_malloc_device( + (filter) ? n_iters * cand_cap : 1, exec_q); + std::uint32_t *cand_ptr = cand_owner.get(); + + sycl::event init_ev = exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(depends); + + using KernelName = radix_select_init_krn; + cgh.parallel_for( + sycl::range<1>(n_iters), [=](sycl::id<1> id) { + state_ptr[id[0]] = StateT{KeyT{0}, KeyT{0}, k, 0, 0, 0, 0}; + }); + }); + + const sycl::nd_range<1> ndRange(n_all_segments * wg_size, wg_size); + + sycl::event pass_ev = init_ev; + for (std::uint32_t pass = 0; pass < n_passes; ++pass) { + const std::uint32_t shift = key_bits - (pass + 1) * radix_bits; + const bool last_pass = (pass + 1 == n_passes); + const bool filter_after = filter && (pass == 0); + + pass_ev = exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(pass_ev); + + sycl::local_accessor hist(radix_states, cgh); + sycl::local_accessor row_hist(radix_states, cgh); + sycl::local_accessor bucket(1, cgh); + sycl::local_accessor is_last(1, cgh); + + using KernelName = + radix_select_count_krn; + cgh.parallel_for(ndRange, [=](sycl::nd_item<1> ndit) { + const auto &wg = ndit.get_group(); + const std::size_t group_id = ndit.get_group(0); + const std::size_t iter_id = group_id / n_segments; + const std::size_t segment_id = group_id - iter_id * n_segments; + const std::size_t lid = ndit.get_local_linear_id(); + + StateT &state = state_ptr[iter_id]; + if (state.done) { + return; + } + const KeyT desired = state.desired; + const KeyT mask = state.mask; + const std::uint64_t n_cand = state.n_cand; + + for (std::size_t b = lid; b < radix_states; b += wg_size) { + hist[b] = 0; + } + sycl::group_barrier(wg); + + const ValueT *row_arg = arg_ptr + iter_id * n_values; + if (n_cand) { + const auto [begin, end] = + segment_range(segment_id, n_segments, n_cand); + count_digits( + ndit, CandKeyT{row_arg, cand_ptr + iter_id * cand_cap}, + begin, end, desired, mask, shift, hist); + } + else { + const auto [begin, end] = + segment_range(segment_id, n_segments, n_values); + count_digits(ndit, RowKeyT{row_arg}, begin, end, desired, + mask, shift, hist); + } + sycl::group_barrier(wg); + + std::uint32_t *row_segment_hist = + segment_hist_ptr + iter_id * n_segments * radix_states; + for (std::size_t b = lid; b < radix_states; b += wg_size) { + row_segment_hist[segment_id * radix_states + b] = hist[b]; + } + + // publish the histogram, then count this work-group in + sycl::atomic_fence(sycl::memory_order::release, + sycl::memory_scope::device); + sycl::group_barrier(wg, sycl::memory_scope::device); + if (lid == 0) { + sycl::atomic_ref + n_arrived(state.n_arrived); + is_last[0] = (n_arrived.fetch_add(std::uint32_t(1)) + 1 == + n_segments); + } + sycl::group_barrier(wg, sycl::memory_scope::device); + if (!is_last[0]) { + return; + } + sycl::atomic_fence(sycl::memory_order::acquire, + sycl::memory_scope::device); + + // the last work-group of the row resolves this pass' digit + for (std::size_t b = lid; b < radix_states; b += wg_size) { + std::uint64_t s = 0; + for (std::size_t j = 0; j < n_segments; ++j) { + s += row_segment_hist[j * radix_states + b]; + } + row_hist[b] = s; + } + sycl::group_barrier(wg); + + const std::uint64_t k_rem = state.k_rem; + const bucket_info info = + find_bucket(ndit, row_hist, k_rem, bucket); + const std::uint64_t new_k_rem = k_rem - info.before; + const bool done = last_pass || (info.count == new_k_rem); + // the elements of this pass' bucket become the candidates + const bool to_filter = + filter_after && !done && (info.count <= cand_cap); + + // the counts restart when the segments switch to candidates + const bool restart = (pass == 0) || (n_cand && pass == 1); + std::uint64_t *row_less_count = + less_count_ptr + iter_id * n_segments; + std::uint64_t *row_less_offset = + less_offset_ptr + iter_id * n_segments; + std::uint64_t *row_tie_count = + tie_count_ptr + iter_id * n_segments; + std::uint64_t *row_tie_offset = + tie_offset_ptr + iter_id * n_segments; + for (std::size_t j = lid; j < n_segments; j += wg_size) { + const std::uint32_t *h = + row_segment_hist + j * radix_states; + std::uint64_t s = 0; + for (std::uint32_t b = 0; b < info.bucket_id; ++b) { + s += h[b]; + } + row_less_count[j] = ((restart) ? 0 : row_less_count[j]) + s; + if (done || to_filter) { + row_tie_count[j] = h[info.bucket_id]; + } + } + + if (done || to_filter) { + sycl::group_barrier(wg); + sycl::joint_exclusive_scan( + wg, row_less_count, row_less_count + n_segments, + row_less_offset, std::uint64_t(0), + sycl::plus()); + sycl::joint_exclusive_scan(wg, row_tie_count, + row_tie_count + n_segments, + row_tie_offset, std::uint64_t(0), + sycl::plus()); + } + + if (lid == 0) { + state.desired = + desired | + static_cast(KeyT(info.bucket_id) << shift); + state.mask = + mask | static_cast(KeyT(radix_mask) << shift); + state.k_rem = new_k_rem; + state.done = done; + state.n_arrived = 0; + if (to_filter) { + state.n_cand = info.count; + state.n_less_first = info.before; + } + } + }); + }); + + if (filter_after) { + // writes out the less elements of the first pass and lists its + // candidates + pass_ev = exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(pass_ev); + + using KernelName = + radix_select_filter_krn; + cgh.parallel_for( + ndRange, [=](sycl::nd_item<1> ndit) { + const std::size_t group_id = ndit.get_group(0); + const std::size_t iter_id = group_id / n_segments; + const std::size_t segment_id = + group_id - iter_id * n_segments; + + const StateT state = state_ptr[iter_id]; + if (state.n_cand == 0) { + return; + } + + const auto [begin, end] = + segment_range(segment_id, n_segments, n_values); + gather_selected( + ndit, RowKeyT{arg_ptr + iter_id * n_values}, begin, + end, state.desired, state.mask, + less_offset_ptr[group_id], less_count_ptr[group_id], + tie_offset_ptr[group_id], state.n_cand, + SelectedOut{ + arg_ptr + iter_id * n_values, + vals_ptr + iter_id * k, inds_ptr + iter_id * k}, + CandidateListOut{cand_ptr + iter_id * cand_cap}); + }); + }); + } + } + + sycl::event gather_ev = exec_q.submit([&](sycl::handler &cgh) { + cgh.depends_on(pass_ev); + + using KernelName = + radix_select_gather_krn; + cgh.parallel_for(ndRange, [=](sycl::nd_item<1> ndit) { + const std::size_t group_id = ndit.get_group(0); + const std::size_t iter_id = group_id / n_segments; + const std::size_t segment_id = group_id - iter_id * n_segments; + + const StateT state = state_ptr[iter_id]; + + const ValueT *row_arg = arg_ptr + iter_id * n_values; + // the first pass wrote out its less elements already + ValueT *row_vals = vals_ptr + iter_id * k + state.n_less_first; + IndexT *row_inds = inds_ptr + iter_id * k + state.n_less_first; + const std::uint64_t n_less = k - state.n_less_first - state.k_rem; + + if (state.n_cand) { + using OutT = CandidateSelectedOut; + const std::uint32_t *row_cand = cand_ptr + iter_id * cand_cap; + const auto [begin, end] = + segment_range(segment_id, n_segments, state.n_cand); + gather_selected( + ndit, CandKeyT{row_arg, row_cand}, begin, end, + state.desired, state.mask, less_offset_ptr[group_id], + less_count_ptr[group_id], tie_offset_ptr[group_id], + state.k_rem, OutT{row_arg, row_cand, row_vals, row_inds}, + OutT{row_arg, row_cand, row_vals + n_less, + row_inds + n_less}); + } + else { + using OutT = SelectedOut; + const auto [begin, end] = + segment_range(segment_id, n_segments, n_values); + gather_selected( + ndit, RowKeyT{row_arg}, begin, end, state.desired, + state.mask, less_offset_ptr[group_id], + less_count_ptr[group_id], tie_offset_ptr[group_id], + state.k_rem, OutT{row_arg, row_vals, row_inds}, + OutT{row_arg, row_vals + n_less, row_inds + n_less}); + } + }); + }); + + return dpnp::tensor::alloc_utils::async_smart_free( + exec_q, {gather_ev}, state_owner, hist_owner, counts_owner, cand_owner); +} + +//----------------------------------------------------------------------- +// radix select: main function +//----------------------------------------------------------------------- + +inline constexpr std::uint32_t gather_elems_per_wi = 4; + +// longest rows ranked by a sub-group each, and the most sub-group sizes they +// may span; the keys are held in registers, so shorter rows get a kernel with +// fewer of them for better occupancy +inline constexpr std::size_t sub_group_max_n = 64; +inline constexpr std::uint32_t sub_group_max_chunks = 8; +inline constexpr std::uint32_t sub_group_few_chunks = 4; + +template +sycl::event radix_select_dispatch(sycl::queue &exec_q, + std::size_t n_iters, + std::size_t n_values, + std::size_t k, + const ValueT *arg_ptr, + ValueT *vals_ptr, + IndexT *inds_ptr, + const std::vector &depends) +{ + const auto &dev = exec_q.get_device(); + const std::size_t max_wg_size = + dev.get_info(); + const std::size_t n_cus = + dev.get_info(); + + const std::size_t wg_size = std::min(256, max_wg_size); + + // short rows spend their time scheduling work-groups, so they are ranked + // by a sub-group each; the kernel may be compiled for any of the device's + // sub-group sizes, so the smallest one bounds the rows it takes + const auto sg_sizes = dev.get_info(); + const std::size_t min_sg_size = + (sg_sizes.empty()) + ? 1 + : *std::min_element(sg_sizes.begin(), sg_sizes.end()); + if (n_values <= + std::min(sub_group_max_n, sub_group_max_chunks * min_sg_size) && + min_sg_size <= wg_size) { + const std::size_t rows_per_group = wg_size / min_sg_size; + const std::size_t n_groups = + (n_iters + rows_per_group - 1) / rows_per_group; + if (n_values <= sub_group_few_chunks * min_sg_size) { + return radix_select_sub_group_submit( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, + n_groups, wg_size, depends); + } + return radix_select_sub_group_submit( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, n_groups, + wg_size, depends); + } + + // enough work-groups to occupy the device + const std::size_t target_groups = 4 * n_cus; + const std::size_t min_segment_size = 8 * wg_size * gather_elems_per_wi; + // below this row size a kernel per pass costs more than it saves + static constexpr std::size_t multi_group_min_n = std::size_t(1) << 16; + + std::size_t n_segments = 1; + if (n_iters < target_groups && n_values >= multi_group_min_n) { + n_segments = + std::min((target_groups + n_iters - 1) / n_iters, + (n_values + min_segment_size - 1) / min_segment_size); + } + // counts of a work-group's elements are 32-bit + static constexpr std::size_t max_segment_size = + std::numeric_limits::max(); + n_segments = std::max(n_segments, + (n_values + max_segment_size - 1) / max_segment_size); + + if (n_segments > 1) { + using KeyT = radix_utils::radix_key_t; + // with fewer passes re-reading the rows costs less than listing the + // candidates + static constexpr bool filter = (sizeof(KeyT) >= 4); + return radix_select_multi_group_impl( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, + n_segments, wg_size, filter, depends); + } + // short rows leave most of a large work-group idle + std::size_t row_wg_size = 64; + while (row_wg_size < wg_size && row_wg_size * 8 < n_values) { + row_wg_size *= 2; + } + row_wg_size = std::min(row_wg_size, wg_size); + return radix_select_one_group_submit( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, row_wg_size, + depends); +} + +/*! @brief Writes the `k` smallest (largest if not `is_ascending`) elements of + * each row of C-contiguous `(n_iters, n_values)` array and their indices into + * `(n_iters, k)` arrays, in unspecified order; ties are resolved in favor of + * smaller indices, NaNs order last, and -0.0 and +0.0 compare equal */ +template +sycl::event radix_select_impl(sycl::queue &exec_q, + std::size_t n_iters, + std::size_t n_values, + std::size_t k, + bool is_ascending, + const ValueT *arg_ptr, + ValueT *vals_ptr, + IndexT *inds_ptr, + const std::vector &depends) +{ + if (k == 0 || k > n_values) { + throw std::runtime_error("Invalid value of k for radix select"); + } + + if (is_ascending) { + return radix_select_dispatch( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, depends); + } + return radix_select_dispatch( + exec_q, n_iters, n_values, k, arg_ptr, vals_ptr, inds_ptr, depends); +} + +} // namespace dpnp::tensor::kernels::radix_select_details diff --git a/dpnp/tensor/libtensor/include/kernels/sorting/radix_sort.hpp b/dpnp/tensor/libtensor/include/kernels/sorting/radix_sort.hpp index 163f2ae64dcc..6950890af85d 100644 --- a/dpnp/tensor/libtensor/include/kernels/sorting/radix_sort.hpp +++ b/dpnp/tensor/libtensor/include/kernels/sorting/radix_sort.hpp @@ -40,7 +40,6 @@ #include #include #include -#include #include #include #include @@ -48,9 +47,9 @@ #include #include "kernels/dpnp_tensor_types.hpp" +#include "kernels/sorting/radix_utils.hpp" #include "kernels/sorting/sort_utils.hpp" #include "utils/sycl_alloc_utils.hpp" -#include "utils/type_utils.hpp" namespace dpnp::tensor::kernels { @@ -70,213 +69,16 @@ class radix_sort_reorder_peer_kernel; template class radix_sort_reorder_kernel; -/*! @brief Computes smallest exponent such that `n <= (1 << exponent)` */ -template && - sizeof(SizeT) == sizeof(std::uint64_t), - int> = 0> -std::uint32_t ceil_log2(SizeT n) -{ - // if n > 2^b, n = q * 2^b + r for q > 0 and 0 <= r < 2^b - // floor_log2(q * 2^b + r) == floor_log2(q * 2^b) == q + floor_log2(n1) - // ceil_log2(n) == 1 + floor_log2(n-1) - if (n <= 1) - return std::uint32_t{1}; - - std::uint32_t exp{1}; - --n; - if (n >= (SizeT{1} << 32)) { - n >>= 32; - exp += 32; - } - if (n >= (SizeT{1} << 16)) { - n >>= 16; - exp += 16; - } - if (n >= (SizeT{1} << 8)) { - n >>= 8; - exp += 8; - } - if (n >= (SizeT{1} << 4)) { - n >>= 4; - exp += 4; - } - if (n >= (SizeT{1} << 2)) { - n >>= 2; - exp += 2; - } - if (n >= (SizeT{1} << 1)) { - n >>= 1; - ++exp; - } - return exp; -} - -//---------------------------------------------------------- -// bitwise order-preserving conversions to unsigned integers -//---------------------------------------------------------- - -template -bool order_preserving_cast(const bool &val) -{ - // by reference: a bool copy lets the compiler assume a 0/1 byte, and the - // bucket index below reads only the low radix bits, see gh-2121 - const bool v = dpnp::tensor::type_utils::normalize_bool(val); - if constexpr (is_ascending) - return v; - else - return !v; -} - -template , int> = 0> -UIntT order_preserving_cast(UIntT val) -{ - if constexpr (is_ascending) { - return val; - } - else { - // bitwise invert - return (~val); - } -} - -template && std::is_signed_v, - int> = 0> -std::make_unsigned_t order_preserving_cast(IntT val) -{ - using UIntT = std::make_unsigned_t; - const UIntT uint_val = sycl::bit_cast(val); - - if constexpr (is_ascending) { - // ascending_mask: 100..0 - static constexpr UIntT ascending_mask = - (UIntT(1) << std::numeric_limits::digits); - return (uint_val ^ ascending_mask); - } - else { - // descending_mask: 011..1 - static constexpr UIntT descending_mask = - (std::numeric_limits::max() >> 1); - return (uint_val ^ descending_mask); - } -} - -template -std::uint16_t order_preserving_cast(sycl::half val) -{ - using UIntT = std::uint16_t; - - const UIntT uint_val = sycl::bit_cast( - (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() - : val); - UIntT mask; - - // test the sign bit of the original value - const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 15)); - - static constexpr UIntT zero_mask = UIntT(0x8000u); - static constexpr UIntT nonzero_mask = UIntT(0xFFFFu); - - static constexpr UIntT inv_zero_mask = static_cast(~zero_mask); - static constexpr UIntT inv_nonzero_mask = static_cast(~nonzero_mask); - - if constexpr (is_ascending) { - mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; - } - else { - mask = (zero_fp_sign_bit) ? (inv_zero_mask) : (inv_nonzero_mask); - } - - return (uint_val ^ mask); -} - -template && - sizeof(FloatT) == sizeof(std::uint32_t), - int> = 0> -std::uint32_t order_preserving_cast(FloatT val) -{ - using UIntT = std::uint32_t; - - UIntT uint_val = sycl::bit_cast( - (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() : val); +// utilities shared with other radix-based algorithms +using radix_utils::ceil_log2; +using radix_utils::get_bucket_id; +using radix_utils::number_of_bits_in_type; +using radix_utils::number_of_buckets_in_type; +using radix_utils::order_preserving_cast; - UIntT mask; - - // test the sign bit of the original value - const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 31)); - - static constexpr UIntT zero_mask = UIntT(0x80000000u); - static constexpr UIntT nonzero_mask = UIntT(0xFFFFFFFFu); - - if constexpr (is_ascending) - mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; - else - mask = (zero_fp_sign_bit) ? (~zero_mask) : (~nonzero_mask); - - return (uint_val ^ mask); -} - -template && - sizeof(FloatT) == sizeof(std::uint64_t), - int> = 0> -std::uint64_t order_preserving_cast(FloatT val) -{ - using UIntT = std::uint64_t; - - UIntT uint_val = sycl::bit_cast( - (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() : val); - UIntT mask; - - // test the sign bit of the original value - const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 63)); - - static constexpr UIntT zero_mask = UIntT(0x8000000000000000u); - static constexpr UIntT nonzero_mask = UIntT(0xFFFFFFFFFFFFFFFFu); - - if constexpr (is_ascending) - mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; - else - mask = (zero_fp_sign_bit) ? (~zero_mask) : (~nonzero_mask); - - return (uint_val ^ mask); -} - -//----------------- -// bucket functions -//----------------- - -template -constexpr std::size_t number_of_bits_in_type() -{ - constexpr std::size_t type_bits = - (sizeof(T) * std::numeric_limits::digits); - return type_bits; -} - -// the number of buckets (size of radix bits) in T -template -constexpr std::uint32_t number_of_buckets_in_type(std::uint32_t radix_bits) -{ - constexpr std::size_t type_bits = number_of_bits_in_type(); - return (type_bits + radix_bits - 1) / radix_bits; -} - -// get bits value (bucket) in a certain radix position -template -std::uint32_t get_bucket_id(T val, std::uint32_t radix_offset) -{ - static_assert(std::is_unsigned_v); - - return (val >> radix_offset) & T(radix_mask); -} +using radix_utils::IdentityProj; +using radix_utils::IndexedProj; +using radix_utils::ValueProj; //-------------------------------- // count kernel (single iteration) @@ -1757,51 +1559,6 @@ sycl::event parallel_radix_sort_impl(sycl::queue &exec_q, return sort_ev; } -struct IdentityProj -{ - constexpr IdentityProj() {} - - template - constexpr T operator()(T val) const - { - return val; - } -}; - -template -struct ValueProj -{ - constexpr ValueProj() {} - - constexpr ValueT operator()(const std::pair &pair) const - { - return pair.first; - } -}; - -template -struct IndexedProj -{ - IndexedProj(const ValueT *arg_ptr) : ptr(arg_ptr), value_projector{} {} - - IndexedProj(const ValueT *arg_ptr, const ProjT &proj_op) - : ptr(arg_ptr), value_projector(proj_op) - { - } - - auto operator()(IndexT i) const - { - // normalize the value read from memory: for bool a byte other than - // 0x00/0x01 would otherwise order by its raw value, see gh-2121 - return value_projector( - dpnp::tensor::type_utils::normalize_bool(ptr[i])); - } - -private: - const ValueT *ptr; - ProjT value_projector; -}; - } // namespace radix_sort_details using dpnp::tensor::ssize_t; diff --git a/dpnp/tensor/libtensor/include/kernels/sorting/radix_utils.hpp b/dpnp/tensor/libtensor/include/kernels/sorting/radix_utils.hpp new file mode 100644 index 000000000000..52d097e89c91 --- /dev/null +++ b/dpnp/tensor/libtensor/include/kernels/sorting/radix_utils.hpp @@ -0,0 +1,371 @@ +//***************************************************************************** +// Copyright (c) 2026, Intel Corporation +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions are met: +// - Redistributions of source code must retain the above copyright notice, +// this list of conditions and the following disclaimer. +// - Redistributions in binary form must reproduce the above copyright notice, +// this list of conditions and the following disclaimer in the documentation +// and/or other materials provided with the distribution. +// - Neither the name of the copyright holder nor the names of its contributors +// may be used to endorse or promote products derived from this software +// without specific prior written permission. +// +// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +// ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +// LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +// CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +// SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +// INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +// CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +// ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF +// THE POSSIBILITY OF SUCH DAMAGE. +//***************************************************************************** +// +//===----------------------------------------------------------------------===// +/// +/// \file +/// This file defines utilities shared by radix sort and select kernels. +//===----------------------------------------------------------------------===// + +#pragma once + +#include +#include +#include +#include +#include + +#include + +#include "utils/type_utils.hpp" + +namespace dpnp::tensor::kernels::radix_utils +{ + +/*! @brief Computes smallest exponent such that `n <= (1 << exponent)` */ +template && + sizeof(SizeT) == sizeof(std::uint64_t), + int> = 0> +std::uint32_t ceil_log2(SizeT n) +{ + // if n > 2^b, n = q * 2^b + r for q > 0 and 0 <= r < 2^b + // floor_log2(q * 2^b + r) == floor_log2(q * 2^b) == q + floor_log2(n1) + // ceil_log2(n) == 1 + floor_log2(n-1) + if (n <= 1) + return std::uint32_t{1}; + + std::uint32_t exp{1}; + --n; + if (n >= (SizeT{1} << 32)) { + n >>= 32; + exp += 32; + } + if (n >= (SizeT{1} << 16)) { + n >>= 16; + exp += 16; + } + if (n >= (SizeT{1} << 8)) { + n >>= 8; + exp += 8; + } + if (n >= (SizeT{1} << 4)) { + n >>= 4; + exp += 4; + } + if (n >= (SizeT{1} << 2)) { + n >>= 2; + exp += 2; + } + if (n >= (SizeT{1} << 1)) { + n >>= 1; + ++exp; + } + return exp; +} + +//---------------------------------------------------------- +// bitwise order-preserving conversions to unsigned integers +//---------------------------------------------------------- + +template +bool order_preserving_cast(const bool &val) +{ + // by reference: a bool copy lets the compiler assume a 0/1 byte, and the + // bucket index below reads only the low radix bits, see gh-2121 + const bool v = dpnp::tensor::type_utils::normalize_bool(val); + if constexpr (is_ascending) + return v; + else + return !v; +} + +template , int> = 0> +UIntT order_preserving_cast(UIntT val) +{ + if constexpr (is_ascending) { + return val; + } + else { + // bitwise invert + return (~val); + } +} + +template && std::is_signed_v, + int> = 0> +std::make_unsigned_t order_preserving_cast(IntT val) +{ + using UIntT = std::make_unsigned_t; + const UIntT uint_val = sycl::bit_cast(val); + + if constexpr (is_ascending) { + // ascending_mask: 100..0 + static constexpr UIntT ascending_mask = + (UIntT(1) << std::numeric_limits::digits); + return (uint_val ^ ascending_mask); + } + else { + // descending_mask: 011..1 + static constexpr UIntT descending_mask = + (std::numeric_limits::max() >> 1); + return (uint_val ^ descending_mask); + } +} + +template +std::uint16_t order_preserving_cast(sycl::half val) +{ + using UIntT = std::uint16_t; + + const UIntT uint_val = sycl::bit_cast( + (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() + : val); + UIntT mask; + + // test the sign bit of the original value + const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 15)); + + static constexpr UIntT zero_mask = UIntT(0x8000u); + static constexpr UIntT nonzero_mask = UIntT(0xFFFFu); + + static constexpr UIntT inv_zero_mask = static_cast(~zero_mask); + static constexpr UIntT inv_nonzero_mask = static_cast(~nonzero_mask); + + if constexpr (is_ascending) { + mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; + } + else { + mask = (zero_fp_sign_bit) ? (inv_zero_mask) : (inv_nonzero_mask); + } + + return (uint_val ^ mask); +} + +template && + sizeof(FloatT) == sizeof(std::uint32_t), + int> = 0> +std::uint32_t order_preserving_cast(FloatT val) +{ + using UIntT = std::uint32_t; + + UIntT uint_val = sycl::bit_cast( + (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() : val); + + UIntT mask; + + // test the sign bit of the original value + const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 31)); + + static constexpr UIntT zero_mask = UIntT(0x80000000u); + static constexpr UIntT nonzero_mask = UIntT(0xFFFFFFFFu); + + if constexpr (is_ascending) + mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; + else + mask = (zero_fp_sign_bit) ? (~zero_mask) : (~nonzero_mask); + + return (uint_val ^ mask); +} + +template && + sizeof(FloatT) == sizeof(std::uint64_t), + int> = 0> +std::uint64_t order_preserving_cast(FloatT val) +{ + using UIntT = std::uint64_t; + + UIntT uint_val = sycl::bit_cast( + (sycl::isnan(val)) ? std::numeric_limits::quiet_NaN() : val); + UIntT mask; + + // test the sign bit of the original value + const bool zero_fp_sign_bit = (UIntT(0) == (uint_val >> 63)); + + static constexpr UIntT zero_mask = UIntT(0x8000000000000000u); + static constexpr UIntT nonzero_mask = UIntT(0xFFFFFFFFFFFFFFFFu); + + if constexpr (is_ascending) + mask = (zero_fp_sign_bit) ? zero_mask : nonzero_mask; + else + mask = (zero_fp_sign_bit) ? (~zero_mask) : (~nonzero_mask); + + return (uint_val ^ mask); +} + +//----------------- +// bucket functions +//----------------- + +template +constexpr std::size_t number_of_bits_in_type() +{ + constexpr std::size_t type_bits = + (sizeof(T) * std::numeric_limits::digits); + return type_bits; +} + +// the number of buckets (size of radix bits) in T +template +constexpr std::uint32_t number_of_buckets_in_type(std::uint32_t radix_bits) +{ + constexpr std::size_t type_bits = number_of_bits_in_type(); + return (type_bits + radix_bits - 1) / radix_bits; +} + +// get bits value (bucket) in a certain radix position +template +std::uint32_t get_bucket_id(T val, std::uint32_t radix_offset) +{ + static_assert(std::is_unsigned_v); + + return (val >> radix_offset) & T(radix_mask); +} + +//-------------------------------------------------------------- +// keys which are equal if and only if the values compare equal +//-------------------------------------------------------------- + +template +struct uint_of_size; + +template <> +struct uint_of_size<1> +{ + using type = std::uint8_t; +}; + +template <> +struct uint_of_size<2> +{ + using type = std::uint16_t; +}; + +template <> +struct uint_of_size<4> +{ + using type = std::uint32_t; +}; + +template <> +struct uint_of_size<8> +{ + using type = std::uint64_t; +}; + +/*! @brief Unsigned integer type of radix keys for values of type `T` */ +template +using radix_key_t = typename uint_of_size::type; + +template +T normalize_signed_zero(T val) +{ + if constexpr (std::is_floating_point_v || + std::is_same_v) { + return (val == T(0)) ? T(0) : val; + } + else { + return val; + } +} + +/*! @brief Order-preserving unsigned key of `val`, unlike + * `order_preserving_cast` equal exactly when the values compare equal, so that + * -0.0 and +0.0 as well as all NaNs share a key */ +template +radix_key_t ordered_radix_key(const T &val) +{ + using KeyT = radix_key_t; + if constexpr (std::is_same_v) { + // by reference, see order_preserving_cast for bool + return KeyT(order_preserving_cast(val)); + } + else { + return KeyT( + order_preserving_cast(normalize_signed_zero(val))); + } +} + +//----------- +// projections +//----------- + +struct IdentityProj +{ + constexpr IdentityProj() {} + + template + constexpr T operator()(T val) const + { + return val; + } +}; + +template +struct ValueProj +{ + constexpr ValueProj() {} + + constexpr ValueT operator()(const std::pair &pair) const + { + return pair.first; + } +}; + +template +struct IndexedProj +{ + IndexedProj(const ValueT *arg_ptr) : ptr(arg_ptr), value_projector{} {} + + IndexedProj(const ValueT *arg_ptr, const ProjT &proj_op) + : ptr(arg_ptr), value_projector(proj_op) + { + } + + auto operator()(IndexT i) const + { + // normalize the value read from memory: for bool a byte other than + // 0x00/0x01 would otherwise order by its raw value, see gh-2121 + return value_projector( + dpnp::tensor::type_utils::normalize_bool(ptr[i])); + } + +private: + const ValueT *ptr; + ProjT value_projector; +}; + +} // namespace dpnp::tensor::kernels::radix_utils diff --git a/dpnp/tensor/libtensor/include/kernels/sorting/topk.hpp b/dpnp/tensor/libtensor/include/kernels/sorting/topk.hpp index b403b3f20fa5..d7ca98f2ae35 100644 --- a/dpnp/tensor/libtensor/include/kernels/sorting/topk.hpp +++ b/dpnp/tensor/libtensor/include/kernels/sorting/topk.hpp @@ -45,6 +45,7 @@ #include #include "kernels/sorting/merge_sort.hpp" +#include "kernels/sorting/radix_select.hpp" #include "kernels/sorting/radix_sort.hpp" #include "kernels/sorting/search_sorted_detail.hpp" #include "kernels/sorting/sort_utils.hpp" @@ -505,4 +506,30 @@ sycl::event topk_radix_impl(sycl::queue &exec_q, return cleanup_ev; } +// the order of the k elements of each row is unspecified +template +sycl::event + topk_radix_select_impl(sycl::queue &exec_q, + std::size_t iter_nelems, // number of sub-arrays + std::size_t axis_nelems, // size of each sub-array + std::size_t k, + bool ascending, + const char *arg_cp, + char *vals_cp, + char *inds_cp, + const std::vector &depends) +{ + if (axis_nelems < k) { + throw std::runtime_error("Invalid sort axis size for value of k"); + } + + const argTy *arg_tp = reinterpret_cast(arg_cp); + argTy *vals_tp = reinterpret_cast(vals_cp); + IndexTy *inds_tp = reinterpret_cast(inds_cp); + + return radix_select_details::radix_select_impl( + exec_q, iter_nelems, axis_nelems, k, ascending, arg_tp, vals_tp, + inds_tp, depends); +} + } // namespace dpnp::tensor::kernels diff --git a/dpnp/tensor/libtensor/source/sorting/topk.cpp b/dpnp/tensor/libtensor/source/sorting/topk.cpp index db4b99ab879b..6e5239953d72 100644 --- a/dpnp/tensor/libtensor/source/sorting/topk.cpp +++ b/dpnp/tensor/libtensor/source/sorting/topk.cpp @@ -51,6 +51,7 @@ #include "utils/output_validation.hpp" #include "utils/rich_comparisons.hpp" #include "utils/type_dispatch.hpp" +#include "utils/type_utils.hpp" #include "topk.hpp" @@ -74,20 +75,10 @@ static topk_impl_fn_ptr_t topk_dispatch_vector[td_ns::num_types]; namespace { -template -struct use_radix_sort : public std::false_type -{ -}; - +// types with an order-preserving radix key template -struct use_radix_sort< - T, - std::enable_if_t, - std::is_same, - std::is_same, - std::is_same, - std::is_same>::value>> - : public std::true_type +struct use_radix_select + : public std::negation> { }; @@ -102,12 +93,12 @@ sycl::event topk_caller(sycl::queue &exec_q, char *inds_cp, const std::vector &depends) { - if constexpr (use_radix_sort::value) { - using dpnp::tensor::kernels::topk_radix_impl; - auto ascending = !largest; - return topk_radix_impl(exec_q, iter_nelems, axis_nelems, - k, ascending, arg_cp, vals_cp, - inds_cp, depends); + if constexpr (use_radix_select::value) { + using dpnp::tensor::kernels::topk_radix_select_impl; + const bool ascending = !largest; + return topk_radix_select_impl( + exec_q, iter_nelems, axis_nelems, k, ascending, arg_cp, vals_cp, + inds_cp, depends); } else { using dpnp::tensor::kernels::topk_merge_impl; diff --git a/dpnp/tests/tensor/test_usm_ndarray_top_k.py b/dpnp/tests/tensor/test_usm_ndarray_top_k.py index 1c04c1fff57a..7484669e3b0f 100644 --- a/dpnp/tests/tensor/test_usm_ndarray_top_k.py +++ b/dpnp/tests/tensor/test_usm_ndarray_top_k.py @@ -26,7 +26,9 @@ # THE POSSIBILITY OF SUCH DAMAGE. # ***************************************************************************** +import numpy as np import pytest +from numpy.testing import assert_array_equal import dpnp.tensor as dpt @@ -267,6 +269,73 @@ def test_top_k_2d_smallest(dtype, n): assert dpt.all(dpt.sort(r.values, axis=1) == dpt.sort(x[:, :k], axis=1)) +def _reference_top_k_inds(x_np, k, mode): + "indices of the top k along the last axis, equal elements by lower index" + if mode == "largest": + n = x_np.shape[-1] + inds = np.argsort(x_np[..., ::-1], axis=-1, kind="stable")[..., ::-1] + inds = n - 1 - inds + else: + inds = np.argsort(x_np, axis=-1, kind="stable") + return inds[..., :k] + + +def _check_top_k(r, x_np, k, mode): + "the result is exact, but its elements may come in any order" + inds = dpt.asnumpy(r.indices) + vals = dpt.asnumpy(r.values) + expected_inds = _reference_top_k_inds(x_np, k, mode) + assert_array_equal(np.sort(inds, axis=-1), np.sort(expected_inds, axis=-1)) + expected_vals = np.take_along_axis(x_np, inds, axis=-1) + assert_array_equal(vals, expected_vals) + if vals.dtype.kind == "f": + # the sign of zeros is kept + assert_array_equal(np.signbit(vals), np.signbit(expected_vals)) + + +@pytest.mark.parametrize( + "dtype", + ["?", "i1", "u1", "i2", "u2", "i4", "u4", "i8", "u8", "f2", "f4", "f8"], +) +@pytest.mark.parametrize( + "shape", [(3000, 5), (2000, 15), (1000, 50), (4, 3001), (2, 100003)] +) +@pytest.mark.parametrize("mode", ["largest", "smallest"]) +def test_top_k_ties(dtype, shape, mode): + q = get_queue_or_skip() + skip_if_dtype_not_supported(dtype, q) + + # few distinct values, so that the k-th value is repeated + rng = np.random.default_rng(42) + x_np = rng.integers(0, 4, size=shape).astype(dtype) + x = dpt.asarray(x_np, sycl_queue=q) + + n = shape[-1] + for k in sorted({1, min(n, 7), min(n, 300), max(1, n // 3), n}): + r = dpt.top_k(x, k, axis=-1, mode=mode) + _check_top_k(r, x_np, k, mode) + + +@pytest.mark.parametrize("dtype", ["f2", "f4", "f8"]) +@pytest.mark.parametrize("n", [11, 257, 100003]) +@pytest.mark.parametrize("mode", ["largest", "smallest"]) +def test_top_k_nans_and_signed_zeros(dtype, n, mode): + q = get_queue_or_skip() + skip_if_dtype_not_supported(dtype, q) + + # NaNs compare equal and order after other values, -0.0 == 0.0 + special = np.array( + [np.nan, -np.nan, 0.0, -0.0, np.inf, -np.inf, 1.0, -1.0], dtype=dtype + ) + rng = np.random.default_rng(7) + x_np = rng.choice(special, size=n) + x = dpt.asarray(x_np, sycl_queue=q) + + for k in sorted({1, min(n, 5), min(n, 200), n // 2, n}): + r = dpt.top_k(x, k, mode=mode) + _check_top_k(r, x_np, k, mode) + + def test_top_k_0d(): get_queue_or_skip() @@ -329,3 +398,26 @@ def test_top_k_validation(): with pytest.raises(ValueError): # mode must be "largest", or "smallest" dpt.top_k(x, 2, mode="invalid") + + +@pytest.mark.parametrize("dtype", ["i4", "u4", "i8", "f4", "f8"]) +@pytest.mark.parametrize("mode", ["largest", "smallest"]) +def test_top_k_long_rows(dtype, mode): + q = get_queue_or_skip() + skip_if_dtype_not_supported(dtype, q) + + # long rows of many distinct values, with repeats of some of them + rng = np.random.default_rng(3) + shape = (3, 300001) + if np.dtype(dtype).kind == "f": + x_np = rng.standard_normal(size=shape).astype(dtype) + else: + info = np.iinfo(dtype) + x_np = rng.integers(info.min, info.max, size=shape, dtype=dtype) + x_np[:, ::7] = x_np[:, 1:2] + x_np[:, 5::11] = x_np[:, 2:3] + x = dpt.asarray(x_np, sycl_queue=q) + + for k in [1, 10, 2000, 60000]: + r = dpt.top_k(x, k, axis=-1, mode=mode) + _check_top_k(r, x_np, k, mode)