Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions cpp/include/raft/matrix/detail/select_k-inl.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -299,6 +299,10 @@ void select_k(raft::resources const& handle,
return detail::select::warpsort::
select_k_impl<T, IdxT, detail::select::warpsort::warp_sort_distributed_ext, row_layout>(
handle, in_val, in_idx, batch_size, len, k, out_val, out_idx, select_min);
case SelectAlgo::kWarpDistributedShmStable:
return detail::select::warpsort::
select_k_impl<T, IdxT, detail::select::warpsort::warp_sort_distributed_ext, row_layout>(
handle, in_val, in_idx, batch_size, len, k, out_val, out_idx, select_min, nullptr, true);
case SelectAlgo::kWarpAuto:
return detail::select::warpsort::select_k<T, IdxT, row_layout>(
handle, in_val, in_idx, batch_size, len, k, out_val, out_idx, select_min);
Expand Down
64 changes: 50 additions & 14 deletions cpp/include/raft/matrix/detail/select_warpsort.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,16 @@ _RAFT_DEVICE _RAFT_FORCEINLINE auto is_ordered(T left, T right) -> bool
if constexpr (!Ascending) { return left > right; }
}

/** `is_ordered` with an equal-key tie broken on the index, as `swap_needed` does. */
template <bool Ascending, typename T, typename IdxT>
_RAFT_DEVICE _RAFT_FORCEINLINE auto is_ordered(T left, T right, IdxT left_idx, IdxT right_idx)
-> bool
{
if (left != right) { return is_ordered<Ascending>(left, right); }
if constexpr (Ascending) { return left_idx < right_idx; }
if constexpr (!Ascending) { return left_idx > right_idx; }
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

} // namespace

/**
Expand Down Expand Up @@ -146,6 +156,9 @@ class warp_sort {
/** Extra memory required per-block for keeping the state (shared or global). */
constexpr static auto mem_required(uint32_t block_size) -> size_t { return 0; }

/** Order an equal-key tie on the payload index rather than on lane arrival. */
_RAFT_DEVICE _RAFT_FORCEINLINE void set_tie_break(bool tie_break) { tie_break_ = tie_break; }

/**
* Construct the warp_sort empty queue.
*
Expand Down Expand Up @@ -188,15 +201,16 @@ class warp_sort {
for (int i = kMaxArrLen - 1; i >= 0; --i, idx += kWarpWidth) {
if (idx < k) {
T t = in[idx];
if (is_ordered<Ascending>(t, val_arr_[i])) {
if (tie_break_ ? is_ordered<Ascending>(t, val_arr_[i], in_idx[idx], idx_arr_[i])
: is_ordered<Ascending>(t, val_arr_[i])) {
val_arr_[i] = t;
idx_arr_[i] = in_idx[idx];
}
}
}
}
if (kWarpWidth < WarpSize || do_merge) {
util::bitonic<kMaxArrLen>(Ascending, kWarpWidth).merge(val_arr_, idx_arr_);
util::bitonic<kMaxArrLen>(Ascending, kWarpWidth, this->tie_break_).merge(val_arr_, idx_arr_);
}
}

Expand Down Expand Up @@ -234,6 +248,8 @@ class warp_sort {

T val_arr_[kMaxArrLen];
IdxT idx_arr_[kMaxArrLen];
/** Order an equal-key tie on the index. Off by default, so the order is unchanged. */
bool tie_break_ = false;

/**
* Merge another array (sorted in the opposite direction) in the queue.
Expand All @@ -259,12 +275,16 @@ class warp_sort {
for (int i = std::min(kMaxArrLen, PerThreadSizeIn); i > 0; i--) {
T& key = val_arr_[kMaxArrLen - i];
T other = keys_in[PerThreadSizeIn - i];
if (is_ordered<Ascending>(other, key)) {
if (tie_break_ ? is_ordered<Ascending>(other,
key,
ids_in[PerThreadSizeIn - i],
idx_arr_[kMaxArrLen - i])
: is_ordered<Ascending>(other, key)) {
key = other;
idx_arr_[kMaxArrLen - i] = ids_in[PerThreadSizeIn - i];
}
}
util::bitonic<kMaxArrLen>(Ascending, kWarpWidth).merge(val_arr_, idx_arr_);
util::bitonic<kMaxArrLen>(Ascending, kWarpWidth, this->tie_break_).merge(val_arr_, idx_arr_);
}
};

Expand Down Expand Up @@ -334,7 +354,7 @@ class warp_sort_filtered : public warp_sort<Capacity, Ascending, T, IdxT> {

_RAFT_DEVICE _RAFT_FORCEINLINE void merge_buf_()
{
util::bitonic<kMaxBufLen>(!Ascending, kWarpWidth).sort(val_buf_, idx_buf_);
util::bitonic<kMaxBufLen>(!Ascending, kWarpWidth, this->tie_break_).sort(val_buf_, idx_buf_);
this->template merge_in<kMaxBufLen>(val_buf_, idx_buf_);
buf_len_ = 0;
set_k_th_(); // contains warp sync
Expand Down Expand Up @@ -459,7 +479,7 @@ class warp_sort_distributed : public warp_sort<Capacity, Ascending, T, IdxT> {

_RAFT_DEVICE _RAFT_FORCEINLINE void merge_buf_()
{
util::bitonic<1>(!Ascending, kWarpWidth).sort(buf_val_, buf_idx_);
util::bitonic<1>(!Ascending, kWarpWidth, this->tie_break_).sort(buf_val_, buf_idx_);
this->template merge_in<1>(&buf_val_, &buf_idx_);
set_k_th_(); // contains warp sync
buf_val_ = kDummy;
Expand Down Expand Up @@ -571,7 +591,7 @@ class warp_sort_distributed_ext : public warp_sort<Capacity, Ascending, T, IdxT>
T buf_val = val_buf_[laneId()];
IdxT buf_idx = idx_buf_[laneId()];
val_buf_[laneId()] = kDummy;
util::bitonic<1>(!Ascending, kWarpWidth).sort(buf_val, buf_idx);
util::bitonic<1>(!Ascending, kWarpWidth, this->tie_break_).sort(buf_val, buf_idx);
this->template merge_in<1>(&buf_val, &buf_idx);
set_k_th_(); // contains warp sync
}
Expand Down Expand Up @@ -630,7 +650,7 @@ class warp_sort_immediate : public warp_sort<Capacity, Ascending, T, IdxT> {

++buf_len_;
if (buf_len_ == kMaxArrLen) {
util::bitonic<kMaxArrLen>(!Ascending, kWarpWidth).sort(val_buf_, idx_buf_);
util::bitonic<kMaxArrLen>(!Ascending, kWarpWidth, this->tie_break_).sort(val_buf_, idx_buf_);
this->template merge_in<kMaxArrLen>(val_buf_, idx_buf_);
#pragma unroll
for (int i = 0; i < kMaxArrLen; i++) {
Expand All @@ -643,7 +663,7 @@ class warp_sort_immediate : public warp_sort<Capacity, Ascending, T, IdxT> {
_RAFT_DEVICE void done()
{
if (buf_len_ != 0) {
util::bitonic<kMaxArrLen>(!Ascending, kWarpWidth).sort(val_buf_, idx_buf_);
util::bitonic<kMaxArrLen>(!Ascending, kWarpWidth, this->tie_break_).sort(val_buf_, idx_buf_);
this->template merge_in<kMaxArrLen>(val_buf_, idx_buf_);
}
}
Expand Down Expand Up @@ -679,6 +699,8 @@ class block_sort {
{
}

_RAFT_DEVICE void set_tie_break(bool tie_break) { queue_.set_tie_break(tie_break); }

_RAFT_DEVICE void add(T val, IdxT idx) { queue_.add(val, idx); }

/**
Expand Down Expand Up @@ -758,7 +780,8 @@ __launch_bounds__(256) RAFT_KERNEL block_kernel(const T* in,
IdxT len,
int k,
T* out,
IdxT* out_idx)
IdxT* out_idx,
bool tie_break)
{
// * per-block output
{
Expand Down Expand Up @@ -788,6 +811,7 @@ __launch_bounds__(256) RAFT_KERNEL block_kernel(const T* in,
using bq_t = block_sort<WarpSortClass, Capacity, Ascending, bits_t, IdxT>;
uint8_t* warp_smem = bq_t::queue_t::mem_required(blockDim.x) > 0 ? smem_buf_bytes : nullptr;
bq_t queue(k, warp_smem);
queue.set_tie_break(tie_break);

// * main loop
const IdxT stride = gridDim.x * blockDim.x;
Expand Down Expand Up @@ -865,6 +889,7 @@ struct launch_setup {
const IdxT* in_indptr,
T* out_key,
IdxT* out_idx,
bool tie_break,
cuda::stream_ref stream)
{
const int capacity = bound_by_power_of_two(k);
Expand All @@ -882,6 +907,7 @@ struct launch_setup {
in_indptr,
out_key,
out_idx,
tie_break,
stream);
}
}
Expand All @@ -906,7 +932,8 @@ struct launch_setup {
IdxT(len),
k,
out_key,
out_idx);
out_idx,
tie_break);
} else {
raft::launch_kernel({stream, smem_size},
gs,
Expand All @@ -919,7 +946,8 @@ struct launch_setup {
IdxT(len),
k,
out_key,
out_idx);
out_idx,
tie_break);
}
RAFT_CUDA_TRY(cudaPeekAtLastError());
out_key += batch_chunk * num_blocks * k;
Expand Down Expand Up @@ -1075,6 +1103,7 @@ void select_k_(bool dry_run,
T* out,
IdxT* out_idx,
bool select_min,
bool tie_break,
cuda::stream_ref stream,
rmm::device_async_resource_ref mr)
{
Expand Down Expand Up @@ -1105,6 +1134,7 @@ void select_k_(bool dry_run,
in_indptr,
result_val,
result_idx,
tie_break,
stream);

if (num_of_block > 1) {
Expand All @@ -1121,6 +1151,7 @@ void select_k_(bool dry_run,
nullptr,
out,
out_idx,
tie_break,
stream);
}
}
Expand All @@ -1138,7 +1169,8 @@ void select_k_impl(raft::resources const& res,
T* out,
IdxT* out_idx,
bool select_min,
const IdxT* in_indptr = nullptr)
const IdxT* in_indptr = nullptr,
bool tie_break = false)
{
int num_of_block = 0;
int num_of_warp = 0;
Expand All @@ -1157,6 +1189,7 @@ void select_k_impl(raft::resources const& res,
out,
out_idx,
select_min,
tie_break,
resource::get_cuda_stream(res),
resource::get_workspace_resource_ref(res));
}
Expand Down Expand Up @@ -1208,7 +1241,8 @@ void select_k(raft::resources const& res,
T* out,
IdxT* out_idx,
bool select_min,
const IdxT* in_indptr = nullptr)
const IdxT* in_indptr = nullptr,
bool tie_break = false)
{
ASSERT(k <= kMaxCapacity, "Current max k is %d (requested %d)", kMaxCapacity, k);
ASSERT(len <= size_t(std::numeric_limits<IdxT>::max()),
Expand All @@ -1235,6 +1269,7 @@ void select_k(raft::resources const& res,
out,
out_idx,
select_min,
tie_break,
resource::get_cuda_stream(res),
resource::get_workspace_resource_ref(res));
} else {
Expand All @@ -1252,6 +1287,7 @@ void select_k(raft::resources const& res,
out,
out_idx,
select_min,
tie_break,
resource::get_cuda_stream(res),
resource::get_workspace_resource_ref(res));
}
Expand Down
7 changes: 7 additions & 0 deletions cpp/include/raft/matrix/select_k_types.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -67,13 +67,20 @@ enum class SelectAlgo : uint8_t {
* in the given external pointers (normally, a shared memory pointer should be passed in).
*/
kWarpDistributedShm = 8,
/**
* The same as `kWarpDistributedShm`, but orders an equal-key tie by the payload index instead of
* by lane arrival, so a repeated selection over the same input returns the same elements.
*/
kWarpDistributedShmStable = 9,
};

inline auto operator<<(std::ostream& os, const SelectAlgo& algo) -> std::ostream&
{
auto underlying_value = static_cast<std::underlying_type<SelectAlgo>::type>(algo);

switch (algo) {
case SelectAlgo::kWarpDistributedShmStable:
return os << "kWarpDistributedShmStable=" << underlying_value;
case SelectAlgo::kAuto: return os << "kAuto=" << underlying_value;
case SelectAlgo::kRadix8bits: return os << "kRadix8bits=" << underlying_value;
case SelectAlgo::kRadix11bits: return os << "kRadix11bits=" << underlying_value;
Expand Down
Loading
Loading