feat(CCCL): device-level CUB algorithms for MoE dispatch

Add complete CCCL CUB header tree (1394 files) to cccl_preload/include/:
- cub/device/ — DeviceRadixSort, DeviceScan, DeviceHistogram, DeviceReduce, DeviceSelect
- cub/agent/ — all agent implementations (sort, scan, reduce, histogram, etc)
- cub/block/ — BlockScan, BlockReduce, BlockExchange, BlockLoad, BlockStore, etc
- cub/warp/ — WarpScan, WarpReduce, WarpExchange, WarpMergeSort
- cub/thread/ — thread-level operators
- thrust/ — sort_by_key, iterator utilities
- cuda/ — execution, stream, memory_resource, functional

New kernel: cccl_moe_sort_scatter.cu
- Uses CUB DeviceRadixSort::SortPairs to sort (expert_id, token_idx) pairs
- O(n) radix sort replaces O(n log n) torch.argsort in MoE prefill path
- Boundary detection + fill for expert offsets/sizes
- Compiled against CCCL upstream headers (not corex CUB) to avoid BI-V100 bugs

Previously only 288 CCCL headers (CachingDeviceAllocator only).
Now 1394 headers — full CUB device-level algorithm stack available for
all future kernels.
This commit is contained in:
project6-dev
2026-08-13 11:18:52 +00:00
parent 7ba97f7977
commit 4c365b8c03
1108 changed files with 294533 additions and 0 deletions

View File

@@ -0,0 +1,245 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_adjacent_difference.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_store.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_namespace.cuh>
#include <cub/util_type.cuh>
#include <thrust/system/cuda/detail/core/util.h>
CUB_NAMESPACE_BEGIN
namespace detail
{
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
cub::BlockLoadAlgorithm LoadAlgorithm = cub::BLOCK_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifier = cub::LOAD_LDG,
cub::BlockStoreAlgorithm StoreAlgorithm = cub::BLOCK_STORE_DIRECT>
struct agent_adjacent_difference_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
static constexpr int ITEMS_PER_TILE = BLOCK_THREADS * ITEMS_PER_THREAD;
static constexpr cub::BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr cub::CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr cub::BlockStoreAlgorithm STORE_ALGORITHM = StoreAlgorithm;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
cub::BlockLoadAlgorithm LoadAlgorithm = cub::BLOCK_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifier = cub::LOAD_LDG,
cub::BlockStoreAlgorithm StoreAlgorithm = cub::BLOCK_STORE_DIRECT>
using AgentAdjacentDifferencePolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceAdjacentDifference") =
detail::agent_adjacent_difference_policy<ThreadsPerBlock, ItemsPerThread, LoadAlgorithm, LoadModifier, StoreAlgorithm>;
namespace detail::adjacent_difference
{
template <typename Policy,
typename InputIteratorT,
typename OutputIteratorT,
typename DifferenceOpT,
typename OffsetT,
typename InputT,
typename OutputT,
bool MayAlias,
bool ReadLeft>
struct AgentDifference
{
using LoadIt = try_make_cache_modified_iterator_t<Policy::LOAD_MODIFIER, InputIteratorT>;
using BlockLoad = typename cub::BlockLoadType<Policy, LoadIt>::type;
using BlockStore = typename cub::BlockStoreType<Policy, OutputIteratorT, OutputT>::type;
using BlockAdjacentDifferenceT = cub::BlockAdjacentDifference<InputT, Policy::BLOCK_THREADS>;
union _TempStorage
{
typename BlockLoad::TempStorage load;
typename BlockStore::TempStorage store;
typename BlockAdjacentDifferenceT::TempStorage adjacent_difference;
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
static constexpr int BLOCK_THREADS = Policy::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = Policy::ITEMS_PER_THREAD;
static constexpr int ITEMS_PER_TILE = Policy::ITEMS_PER_TILE;
static constexpr int SHARED_MEMORY_SIZE = static_cast<int>(sizeof(TempStorage));
_TempStorage& temp_storage;
InputIteratorT input_it;
LoadIt load_it;
InputT* first_tile_previous;
OutputIteratorT result;
DifferenceOpT difference_op;
OffsetT num_items;
_CCCL_DEVICE _CCCL_FORCEINLINE AgentDifference(
TempStorage& temp_storage,
InputIteratorT input_it,
InputT* first_tile_previous,
OutputIteratorT result,
DifferenceOpT difference_op,
OffsetT num_items)
: temp_storage(temp_storage.Alias())
, input_it(input_it)
, load_it(try_make_cache_modified_iterator<Policy::LOAD_MODIFIER>(input_it))
, first_tile_previous(first_tile_previous)
, result(result)
, difference_op(difference_op)
, num_items(num_items)
{}
template <bool IS_LAST_TILE, bool IS_FIRST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile_impl(int num_remaining, int tile_idx, OffsetT tile_base)
{
InputT input[ITEMS_PER_THREAD];
OutputT output[ITEMS_PER_THREAD];
if (IS_LAST_TILE)
{
// Fill last elements with the first element
// because collectives are not suffix guarded
BlockLoad(temp_storage.load).Load(load_it + tile_base, input, num_remaining, *(load_it + tile_base));
}
else
{
BlockLoad(temp_storage.load).Load(load_it + tile_base, input);
}
__syncthreads();
if (ReadLeft)
{
if (IS_FIRST_TILE)
{
if (IS_LAST_TILE)
{
BlockAdjacentDifferenceT(temp_storage.adjacent_difference)
.SubtractLeftPartialTile(input, output, difference_op, num_remaining);
}
else
{
BlockAdjacentDifferenceT(temp_storage.adjacent_difference).SubtractLeft(input, output, difference_op);
}
}
else
{
InputT tile_prev_input = MayAlias ? first_tile_previous[tile_idx] : *(input_it + tile_base - 1);
if (IS_LAST_TILE)
{
BlockAdjacentDifferenceT(temp_storage.adjacent_difference)
.SubtractLeftPartialTile(input, output, difference_op, num_remaining, tile_prev_input);
}
else
{
BlockAdjacentDifferenceT(temp_storage.adjacent_difference)
.SubtractLeft(input, output, difference_op, tile_prev_input);
}
}
}
else
{
if (IS_LAST_TILE)
{
BlockAdjacentDifferenceT(temp_storage.adjacent_difference)
.SubtractRightPartialTile(input, output, difference_op, num_remaining);
}
else
{
InputT tile_next_input = MayAlias ? first_tile_previous[tile_idx] : *(input_it + tile_base + ITEMS_PER_TILE);
BlockAdjacentDifferenceT(temp_storage.adjacent_difference)
.SubtractRight(input, output, difference_op, tile_next_input);
}
}
__syncthreads();
if (IS_LAST_TILE)
{
BlockStore(temp_storage.store).Store(result + tile_base, output, num_remaining);
}
else
{
BlockStore(temp_storage.store).Store(result + tile_base, output);
}
}
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile(int num_remaining, int tile_idx, OffsetT tile_base)
{
if (tile_idx == 0)
{
consume_tile_impl<IS_LAST_TILE, true>(num_remaining, tile_idx, tile_base);
}
else
{
consume_tile_impl<IS_LAST_TILE, false>(num_remaining, tile_idx, tile_base);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process(int tile_idx, OffsetT tile_base)
{
OffsetT num_remaining = num_items - tile_base;
if (num_remaining > ITEMS_PER_TILE) // not a last tile
{
consume_tile<false>(num_remaining, tile_idx, tile_base);
}
else
{
consume_tile<true>(num_remaining, tile_idx, tile_base);
}
}
};
template <typename InputIteratorT, typename InputT, typename OffsetT, bool ReadLeft>
struct AgentDifferenceInit
{
static constexpr int BLOCK_THREADS = 128;
static _CCCL_DEVICE _CCCL_FORCEINLINE void
Process(int tile_idx, InputIteratorT first, InputT* result, OffsetT num_tiles, int items_per_tile)
{
OffsetT tile_base = static_cast<OffsetT>(tile_idx) * items_per_tile;
if (tile_base > 0 && tile_idx < num_tiles)
{
if (ReadLeft)
{
result[tile_idx] = first[tile_base - 1];
}
else
{
result[tile_idx - 1] = first[tile_base];
}
}
}
};
} // namespace detail::adjacent_difference
CUB_NAMESPACE_END

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,375 @@
// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/block/block_topk.cuh>
#include <cub/detail/choose_offset.cuh>
#include <cub/detail/segmented_params.cuh>
#include <cub/device/dispatch/dispatch_common.cuh>
#include <cub/device/dispatch/tuning/tuning_batched_topk.cuh>
#include <cub/util_type.cuh>
#include <cuda/__cmath/ceil_div.h>
#include <cuda/argument>
CUB_NAMESPACE_BEGIN
namespace detail::batched_topk
{
// Atomic counters used by the small-segment kernel to (a) enqueue large segments into the large-segment work queue
// and (b) elect the last block to run the epilogue scan over the queued tile counts. `alignas(128)` isolates each
// counter on its own cache line for performance.
template <class NumSegmentsT>
struct batched_topk_counters
{
// Force unsigned integer type for segment count.
using segment_count_t = detail::choose_offset_t<NumSegmentsT>;
// Number of segments enqueued in the large-segment work queue. Atomically incremented (by 1) by the first thread
// of each block that decides its segment is large.
alignas(128) segment_count_t large_segments_count;
// Block retirement counter. Each block atomically increments by 1 when it has finished processing its segment, and
// the block that observes `gridDim.x - 1` runs the epilogue on the queued large segments tile counts.
// Assumption: Future support for more than 2^31 - 1 segments will use multiple launches of a slightly modified
// small-segment kernel instead of additional grid dimensions. Therefore each grid will handle a maximum of 2^31 - 1
// segments per launch. The counter would not even have to be reset to 0 after each launch if we cleverly make use of
// its modulo arithmetic.
alignas(128) unsigned retirement_count;
};
template <typename PolicyGetter, // TODO(bgruber): pass worker_policy as NTTP in C++20
typename KeyInputItItT,
typename KeyOutputItItT,
typename ValueInputItItT,
typename ValueOutputItItT,
typename SegmentSizeParameterT,
typename KParameterT,
typename SelectDirectionParameterT,
typename NumSegmentsParameterT,
typename LargeSegmentTileOffsetT>
struct agent_batched_topk_worker_per_segment
{
// -------------------------------------------------------------------------
// Types and Constants
// -------------------------------------------------------------------------
// Derive inner types from Iterator of Iterators
using key_it_t = it_value_t<KeyInputItItT>;
using value_it_t = it_value_t<ValueInputItItT>;
using key_t = it_value_t<key_it_t>;
using value_t = it_value_t<value_it_t>;
using segment_size_val_t = typename ::cuda::args::__traits<SegmentSizeParameterT>::element_type;
using num_segments_val_t = typename ::cuda::args::__traits<NumSegmentsParameterT>::element_type;
using counters_t = batched_topk_counters<num_segments_val_t>;
static constexpr auto policy = PolicyGetter{}();
static constexpr worker_policy active_policy = policy.worker_per_segment_policy;
// For block-topk (and keys/values load/store):
static constexpr int threads_per_block = active_policy.threads_per_block;
static constexpr int items_per_thread = active_policy.items_per_thread;
static constexpr int tile_size = threads_per_block * items_per_thread;
// For block-scan (and offsets load/store):
static constexpr int epilogue_items_per_thread = active_policy.epilogue.items_per_thread;
static constexpr int epilogue_tile_size = threads_per_block * epilogue_items_per_thread;
// Number used for preprocessing segment-size data, not for tuning => should not affect performance of this agent.
static constexpr multi_worker_policy multi_worker_per_segment_policy = policy.multi_worker_per_segment_policy;
static constexpr int multi_worker_per_segment_tile_size =
multi_worker_per_segment_policy.threads_per_block * multi_worker_per_segment_policy.items_per_thread;
// Check if there could be large segments present
static constexpr bool only_small_segments = ::cuda::args::__traits<SegmentSizeParameterT>::highest <= tile_size;
// Check if we are dealing with keys-only or key-value pairs
static constexpr bool is_keys_only = ::cuda::std::is_same_v<value_t, cub::NullType>;
// -------------------------------------------------------------------------
// Primitive Types
// -------------------------------------------------------------------------
using block_load_keys_t = BlockLoad<key_t, threads_per_block, items_per_thread, active_policy.load_algorithm>;
using block_load_vals_t = BlockLoad<value_t, threads_per_block, items_per_thread, active_policy.load_algorithm>;
using block_topk_t = block_topk<key_t, threads_per_block, items_per_thread, value_t>;
// TODO (elstehle): Specialize for the case that we statically know k and we can skip passing num_valid_items to
// Store()
using block_store_keys_t = BlockStore<key_t, threads_per_block, items_per_thread, active_policy.store_algorithm>;
using block_store_vals_t = BlockStore<value_t, threads_per_block, items_per_thread, active_policy.store_algorithm>;
using block_load_epilogue_t =
BlockLoad<segment_size_val_t, threads_per_block, epilogue_items_per_thread, active_policy.epilogue.load_algorithm>;
using block_scan_epilogue_t = BlockScan<int, threads_per_block, active_policy.epilogue.scan_algorithm>;
using block_store_epilogue_t =
BlockStore<segment_size_val_t, threads_per_block, epilogue_items_per_thread, active_policy.epilogue.store_algorithm>;
// -------------------------------------------------------------------------
// Shared Memory Storage
// -------------------------------------------------------------------------
struct TempStorage_
{
union
{
typename block_load_keys_t::TempStorage load_keys;
typename block_load_vals_t::TempStorage load_vals;
typename block_topk_t::TempStorage topk;
typename block_store_keys_t::TempStorage store_keys;
typename block_store_vals_t::TempStorage store_vals;
typename block_load_epilogue_t::TempStorage load_epilogue;
typename block_scan_epilogue_t::TempStorage scan_epilogue;
typename block_store_epilogue_t::TempStorage store_epilogue;
};
};
using TempStorage = Uninitialized<TempStorage_>;
// -------------------------------------------------------------------------
// Members
// -------------------------------------------------------------------------
TempStorage_& temp_storage;
KeyInputItItT d_key_segments_it;
KeyOutputItItT d_key_segments_out_it;
ValueInputItItT d_value_segments_it;
ValueOutputItItT d_value_segments_out_it;
SegmentSizeParameterT segment_sizes;
KParameterT k_param;
SelectDirectionParameterT select_directions;
NumSegmentsParameterT num_segments;
counters_t* d_counters;
num_segments_val_t* d_large_segments_ids;
LargeSegmentTileOffsetT* d_large_segments_tile_offsets;
// -------------------------------------------------------------------------
// Constructor
// -------------------------------------------------------------------------
_CCCL_DEVICE_API _CCCL_FORCEINLINE agent_batched_topk_worker_per_segment(
TempStorage& temp_storage,
KeyInputItItT d_key_segments_it,
KeyOutputItItT d_key_segments_out_it,
ValueInputItItT d_value_segments_it,
ValueOutputItItT d_value_segments_out_it,
SegmentSizeParameterT segment_sizes,
KParameterT k_param,
SelectDirectionParameterT select_directions,
NumSegmentsParameterT num_segments,
counters_t* d_counters,
num_segments_val_t* d_large_segments_ids,
LargeSegmentTileOffsetT* d_large_segments_tile_offsets)
: temp_storage(temp_storage.Alias())
, d_key_segments_it(d_key_segments_it)
, d_key_segments_out_it(d_key_segments_out_it)
, d_value_segments_it(d_value_segments_it)
, d_value_segments_out_it(d_value_segments_out_it)
, segment_sizes(segment_sizes)
, k_param(k_param)
, select_directions(select_directions)
, num_segments(num_segments)
, d_counters(d_counters)
, d_large_segments_ids(d_large_segments_ids)
, d_large_segments_tile_offsets(d_large_segments_tile_offsets)
{}
_CCCL_DEVICE_API _CCCL_FORCEINLINE void Process()
{
// Identify Segment
const int segment_id = static_cast<int>(blockIdx.x);
// Boundary check
// TODO (elstehle): consider skipping boundary check if we can safely assume the right grid dimensions
if (segment_id >= params::get_param(num_segments, 0))
{
return;
}
constexpr bool is_full_tile = ::cuda::args::__traits<SegmentSizeParameterT>::is_constant
&& ::cuda::args::__traits<SegmentSizeParameterT>::lowest == tile_size;
// Resolve Segment Parameters
const auto segment_size = params::get_param(segment_sizes, segment_id);
if (!only_small_segments && segment_size > tile_size)
{
// Enqueue large segment
if (threadIdx.x == 0u)
{
// Add to large segment queue
const auto large_segment_queue_idx = atomicAdd(&d_counters->large_segments_count, 1ull);
d_large_segments_ids[large_segment_queue_idx] = static_cast<num_segments_val_t>(segment_id);
d_large_segments_tile_offsets[large_segment_queue_idx] =
static_cast<LargeSegmentTileOffsetT>(::cuda::ceil_div(segment_size, multi_worker_per_segment_tile_size));
}
}
else
{
// Process small segment
const auto k = (::cuda::std::min) (params::get_param(k_param, segment_id),
static_cast<decltype(params::get_param(k_param, segment_id))>(segment_size));
const auto direction = select_directions.get_param(segment_id);
// Determine padding key based on direction
const key_t padding_key =
(direction == detail::topk::select::max)
? ::cuda::std::numeric_limits<key_t>::lowest()
: (::cuda::std::numeric_limits<key_t>::max)();
// Dereference iterator-of-iterators to get the segment specific iterator
auto block_keys_in = d_key_segments_it[segment_id];
// Load Keys
key_t thread_keys[items_per_thread];
if constexpr (is_full_tile)
{
// No padding needed
block_load_keys_t(temp_storage.load_keys).Load(block_keys_in, thread_keys);
}
else
{
// Potentially partial final load with padding
// TODO (elstehle): explore whether a runtime check for segment_size == tile_size improves performance
block_load_keys_t(temp_storage.load_keys).Load(block_keys_in, thread_keys, segment_size);
}
// Load Values (if applicable)
[[maybe_unused]] value_t thread_values[items_per_thread];
if constexpr (!is_keys_only)
{
__syncthreads();
auto block_vals_in = d_value_segments_it[segment_id];
if constexpr (is_full_tile)
{
// No padding needed
block_load_vals_t(temp_storage.load_vals).Load(block_vals_in, thread_values);
}
else
{
// Potentially partial final load with padding
// TODO (elstehle): explore whether a runtime check for segment_size == tile_size improves performance
block_load_vals_t(temp_storage.load_vals).Load(block_vals_in, thread_values, segment_size);
}
}
__syncthreads();
// Perform Block Top-K
if constexpr (is_keys_only)
{
const bool is_successful_dispatch = cub::detail::params::dispatch_discrete(
select_directions, segment_id, [this, &thread_keys, k, segment_size](auto direction_tag) {
if constexpr (decltype(direction_tag)::value == detail::topk::select::max)
{
block_topk_t(temp_storage.topk).template max_keys<is_full_tile>(thread_keys, k, segment_size);
}
else
{
block_topk_t(temp_storage.topk).template min_keys<is_full_tile>(thread_keys, k, segment_size);
}
});
_CCCL_ASSERT(is_successful_dispatch, "Error: Unsupported select direction");
}
else
{
// Pass both keys and values
const bool is_successful_dispatch = cub::detail::params::dispatch_discrete(
select_directions, segment_id, [this, &thread_keys, &thread_values, k, segment_size](auto direction_tag) {
if constexpr (decltype(direction_tag)::value == detail::topk::select::max)
{
block_topk_t(temp_storage.topk)
.template max_pairs<is_full_tile>(thread_keys, thread_values, k, segment_size);
}
else
{
block_topk_t(temp_storage.topk)
.template min_pairs<is_full_tile>(thread_keys, thread_values, k, segment_size);
}
});
_CCCL_ASSERT(is_successful_dispatch, "Error: Unsupported select direction");
}
__syncthreads();
auto block_keys_out = d_key_segments_out_it[segment_id];
block_store_keys_t(temp_storage.store_keys)
.Store(block_keys_out,
thread_keys,
k // Only store K items
);
if constexpr (!is_keys_only)
{
__syncthreads();
auto block_vals_out = d_value_segments_out_it[segment_id];
block_store_vals_t(temp_storage.store_vals).Store(block_vals_out, thread_values, k);
}
}
// Epilogue: Scan queued large segment sizes (in tiles not elements) for load balancing search in the large segment
// agent
if constexpr (!only_small_segments)
{
// Determine last block trying to retire.
bool is_last_block = false;
if (threadIdx.x == 0u)
{
__threadfence();
const auto retirement_count = atomicAdd(&d_counters->retirement_count, 1u);
is_last_block = retirement_count == (gridDim.x - 1u);
}
// This sync also makes sure that the shared memory can be reused.
is_last_block = static_cast<bool>(__syncthreads_or(static_cast<int>(is_last_block)));
if (!is_last_block)
{
return;
}
const auto num_large_segments = d_counters->large_segments_count;
// For tracking the running total across tiles (loop iterations).
// Caution: The functor is only invoked by the first warp in the block, and the value returned by lane 0 in that
// warp is used as the initial value.
const auto prefix_callback_op =
[running_total = segment_size_val_t{0}](segment_size_val_t block_aggregate) mutable {
auto old_running_total = running_total;
running_total += block_aggregate;
return old_running_total;
};
_CCCL_PRAGMA_NOUNROLL()
for (int large_segment_offset = 0; large_segment_offset < num_large_segments;
large_segment_offset += epilogue_tile_size)
{
segment_size_val_t segment_tile_offsets[epilogue_items_per_thread];
block_load_epilogue_t(temp_storage.load_epilogue)
.Load(d_large_segments_tile_offsets + large_segment_offset,
segment_tile_offsets,
num_large_segments - large_segment_offset,
0);
__syncthreads();
block_scan_epilogue_t(temp_storage.scan_epilogue)
.ExclusiveSum(segment_tile_offsets, segment_tile_offsets, prefix_callback_op);
__syncthreads();
block_store_epilogue_t(temp_storage.store_epilogue)
.Store(d_large_segments_tile_offsets + large_segment_offset,
segment_tile_offsets,
num_large_segments - large_segment_offset);
__syncthreads();
}
}
}
};
} // namespace detail::batched_topk
CUB_NAMESPACE_END

View File

@@ -0,0 +1,201 @@
// SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
#pragma once
#include <cub/config.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/thread/thread_load.cuh>
#include <cub/util_arch.cuh>
#include <cub/util_type.cuh>
#include <thrust/detail/raw_reference_cast.h>
#include <thrust/type_traits/is_trivially_relocatable.h>
#include <cuda/__memory/is_aligned.h>
#if !_CCCL_HAS_NV_ATOMIC_BUILTINS()
# include <cuda/atomic>
#endif // !_CCCL_HAS_NV_ATOMIC_BUILTINS()
#include <cuda/std/__type_traits/integral_constant.h>
CUB_NAMESPACE_BEGIN
namespace detail::find
{
template <int ThreadsPerBlock,
int ItemsPerThread,
int VecSize,
CacheLoadModifier LoadModifier,
typename InputIteratorT,
typename OffsetT,
typename PredicateT>
struct agent_t
{
// The input value type
using InputT = typename ::cuda::std::iterator_traits<InputIteratorT>::value_type;
// Vector type of InputT for data movement
using VectorT = typename CubVector<InputT, VecSize>::Type;
static constexpr int tile_size = ThreadsPerBlock * ItemsPerThread;
// Can vectorize according to the policy if the input iterator is a native pointer to a primitive type
static constexpr bool attempt_vectorization =
(VecSize > 1) && (ItemsPerThread % VecSize == 0) && (::cuda::std::contiguous_iterator<InputIteratorT>)
&& THRUST_NS_QUALIFIER::is_trivially_relocatable_v<InputT>;
static constexpr CacheLoadModifier load_modifier = LoadModifier;
// Shared memory type required by this thread block
struct _TempStorage
{
OffsetT global_result;
OffsetT block_result;
};
// Alias wrapper allowing storage to be unioned
using TempStorage = Uninitialized<_TempStorage>;
_TempStorage& temp_storage;
InputIteratorT d_in;
PredicateT predicate;
OffsetT* found_pos_ptr;
OffsetT num_items;
template <typename Iterator = InputIteratorT, bool CanVectorize = attempt_vectorization>
_CCCL_DEVICE _CCCL_FORCEINLINE bool is_aligned_and_full_tile(OffsetT tile_offset)
{
if constexpr (CanVectorize)
{
static_assert(::cuda::std::is_pointer_v<Iterator>);
// Retrieve the value type from the iterator to determine the vector type
using InputT = typename ::cuda::std::iterator_traits<Iterator>::value_type;
using VectorT = typename CubVector<InputT, VecSize>::Type;
const bool full_tile = (tile_offset + tile_size) <= num_items;
// Check alignment at the actual load position (d_in + tile_offset)
return full_tile && ::cuda::is_aligned(d_in + tile_offset, sizeof(VectorT));
}
else
{
return false;
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE bool
ConsumeTile(OffsetT tile_offset, ::cuda::std::integral_constant<bool, true> /*CAN_VECTORIZE*/)
{
using InputT = typename ::cuda::std::iterator_traits<InputIteratorT>::value_type;
using VectorT = typename CubVector<InputT, VecSize>::Type;
// vectorized loads begin
auto load_ptr = reinterpret_cast<const VectorT*>(d_in + tile_offset + (threadIdx.x * VecSize));
CacheModifiedInputIterator<LoadModifier, VectorT> d_vec_in(load_ptr);
alignas(InputT) unsigned char input_bytes[ItemsPerThread * sizeof(InputT)];
auto* vec_items = reinterpret_cast<VectorT*>(input_bytes);
constexpr int number_of_vectors = ItemsPerThread / VecSize;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < number_of_vectors; ++i)
{
vec_items[i] = d_vec_in[ThreadsPerBlock * i];
}
for (int i = 0; i < ItemsPerThread; ++i)
{
OffsetT nth_vector_of_thread = i / VecSize;
OffsetT element_in_vector = i % VecSize;
OffsetT vector_of_tile = nth_vector_of_thread * ThreadsPerBlock + threadIdx.x;
OffsetT index = tile_offset + vector_of_tile * VecSize + element_in_vector;
auto* input_items = reinterpret_cast<InputT*>(input_bytes);
if (index < num_items && predicate(input_items[i]))
{
atomicMin(&temp_storage.block_result, index);
// every thread goes over multiple elements per thread for every tile. If a thread finds a local minimum it
// doesn't need to proceed further (inner early exit).
return true;
}
}
return false;
}
_CCCL_DEVICE _CCCL_FORCEINLINE bool
ConsumeTile(OffsetT tile_offset, ::cuda::std::integral_constant<bool, false> /*CAN_VECTORIZE*/)
{
for (int i = 0; i < ItemsPerThread; ++i)
{
const auto index = tile_offset + threadIdx.x + i * blockDim.x;
if (index < num_items)
{
// using raw_reference_cast and passing directly to predicate should avoid creating a copy, and thus prevent
// bugs like: http://github.com/NVIDIA/cccl/issues/3591
if (predicate(THRUST_NS_QUALIFIER::raw_reference_cast(d_in[index])))
{
atomicMin(&temp_storage.block_result, index);
return true;
}
}
}
return false;
}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
if (threadIdx.x == 0)
{
temp_storage.block_result = num_items;
}
__syncthreads();
// use a grid strided loop
OffsetT grid_stride = static_cast<OffsetT>(tile_size) * static_cast<OffsetT>(gridDim.x);
for (OffsetT tile_offset = static_cast<OffsetT>(blockIdx.x) * static_cast<OffsetT>(tile_size);
tile_offset < num_items;
tile_offset += grid_stride)
{
// Only one thread reads atomically and propagates it to other threads of the block through shared memory
if (threadIdx.x == 0)
{
#if _CCCL_HAS_NV_ATOMIC_BUILTINS()
// __nv_atomic_load is a compiler build-in and compiles a lot faster
__nv_atomic_load(found_pos_ptr, &temp_storage.global_result, __NV_ATOMIC_RELAXED, __NV_THREAD_SCOPE_DEVICE);
#else // ^^^ _CCCL_HAS_NV_ATOMIC_BUILTINS() ^^^ / vvv !_CCCL_HAS_NV_ATOMIC_BUILTINS() vvv
temp_storage.global_result = ::cuda::atomic_ref<OffsetT, ::cuda::std::thread_scope_device>{*found_pos_ptr}.load(
::cuda::std::memory_order_relaxed);
#endif // !_CCCL_HAS_NV_ATOMIC_BUILTINS()
}
__syncthreads();
// early exit
if (temp_storage.global_result < tile_offset)
{
return;
}
const bool found_thread =
is_aligned_and_full_tile(tile_offset)
? ConsumeTile(tile_offset, ::cuda::std::bool_constant<attempt_vectorization>{})
: ConsumeTile(tile_offset, ::cuda::std::false_type{});
const bool found_block = __syncthreads_or(found_thread);
if (found_block)
{
// our block found it, update global position and exit
if (threadIdx.x == 0 && temp_storage.block_result < num_items)
{
atomicMin(found_pos_ptr, temp_storage.block_result);
}
return;
}
}
}
};
} // namespace detail::find
CUB_NAMESPACE_END

View File

@@ -0,0 +1,222 @@
// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_merge_sort.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_namespace.cuh>
#include <cub/util_type.cuh>
#include <cuda/std/__algorithm/min.h>
#include <cuda/std/__utility/forward.h>
CUB_NAMESPACE_BEGIN
namespace detail::find_bound_sorted_values
{
// lower_bound vs upper_bound: partition comparator and per-step advance differ.
struct lower_bound_mode
{
// Wrap user comp so the merge path partitions identically to std::lower_bound.
template <typename CompareOp>
struct partition_comp_t
{
CompareOp comp;
template <typename A, typename B>
_CCCL_HOST_DEVICE_API _CCCL_FORCEINLINE bool operator()(A&& a, B&& b) const
{
return !comp(::cuda::std::forward<B>(b), ::cuda::std::forward<A>(a));
}
};
template <typename CompareOp>
_CCCL_HOST_DEVICE_API static partition_comp_t<CompareOp> make_partition_comp(CompareOp compare_op)
{
return partition_comp_t<CompareOp>{compare_op};
}
template <typename HaystackT, typename NeedlesT, typename CompareOp>
_CCCL_DEVICE_API _CCCL_FORCEINLINE static bool
should_advance(const HaystackT& haystack_value, const NeedlesT& needle_value, CompareOp compare_op)
{
return compare_op(haystack_value, needle_value);
}
};
struct upper_bound_mode
{
template <typename CompareOp>
using partition_comp_t = CompareOp;
template <typename CompareOp>
_CCCL_HOST_DEVICE_API static CompareOp make_partition_comp(CompareOp compare_op)
{
return compare_op;
}
template <typename HaystackT, typename NeedlesT, typename CompareOp>
_CCCL_DEVICE_API _CCCL_FORCEINLINE static bool
should_advance(const HaystackT& haystack_value, const NeedlesT& needle_value, CompareOp compare_op)
{
return !compare_op(needle_value, haystack_value);
}
};
template <int ThreadsPerBlock,
int ItemsPerThread,
CacheLoadModifier LoadModifier,
typename Mode,
typename HaystackIt,
typename NeedlesIt,
typename OutputIt,
typename Offset,
typename CompareOp>
struct agent_t
{
static constexpr int tile_size = ThreadsPerBlock * ItemsPerThread;
using haystack_type = it_value_t<HaystackIt>;
using needles_type = it_value_t<NeedlesIt>;
// Separate buffers because haystack and needles may have different value types.
struct _TempStorage
{
haystack_type haystack[tile_size];
needles_type needles[tile_size];
};
using TempStorage = Uninitialized<_TempStorage>;
_TempStorage& storage;
HaystackIt d_range;
NeedlesIt d_values;
OutputIt d_output;
Offset range_count;
Offset values_count;
Offset* range_beg_offsets;
CompareOp compare_op;
template <bool IsFullTile>
_CCCL_DEVICE_API _CCCL_FORCEINLINE void consume_tile(int tile_idx, Offset diag0, int total_in_tile)
{
const Offset range_beg = range_beg_offsets[tile_idx];
const Offset range_end = range_beg_offsets[tile_idx + 1];
_CCCL_ASSERT(range_end >= range_beg, "");
_CCCL_ASSERT(diag0 >= range_beg, "");
const Offset values_beg = diag0 - range_beg;
const int haystack_count = static_cast<int>(range_end - range_beg);
const int needles_count = total_in_tile - haystack_count;
{
const auto d_range_cm = cub::detail::try_make_cache_modified_iterator<LoadModifier>(d_range + range_beg);
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ItemsPerThread; ++item)
{
const int idx = ThreadsPerBlock * item + threadIdx.x;
if (idx < haystack_count)
{
storage.haystack[idx] = d_range_cm[idx];
}
}
}
{
auto d_values_cm = cub::detail::try_make_cache_modified_iterator<LoadModifier>(d_values + values_beg);
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ItemsPerThread; ++item)
{
const int idx = ThreadsPerBlock * item + threadIdx.x;
if (idx < needles_count)
{
storage.needles[idx] = d_values_cm[idx];
}
}
}
__syncthreads();
#ifdef CCCL_ENABLE_DEVICE_ASSERTIONS
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ItemsPerThread; ++item)
{
const int idx = ThreadsPerBlock * item + threadIdx.x;
if (idx < needles_count && (values_beg + idx) > 0)
{
const needles_type prev = (idx == 0) ? d_values[values_beg - 1] : storage.needles[idx - 1];
_CCCL_ASSERT(!compare_op(storage.needles[idx], prev), "d_values must be sorted consistently with comp");
}
}
#endif // CCCL_ENABLE_DEVICE_ASSERTIONS
const auto partition_comp = Mode::make_partition_comp(compare_op);
int d0_thread = ItemsPerThread * static_cast<int>(threadIdx.x);
if constexpr (!IsFullTile)
{
d0_thread = ::cuda::std::min(d0_thread, total_in_tile);
}
const int i0 =
cub::MergePath(storage.haystack, storage.needles, haystack_count, needles_count, d0_thread, partition_comp);
const int j0 = d0_thread - i0;
int i = i0;
int j = j0;
int haystack_remaining = haystack_count - i0;
int needles_remaining = needles_count - j0;
const int steps = IsFullTile ? ItemsPerThread : ::cuda::std::min(total_in_tile - d0_thread, ItemsPerThread);
_CCCL_PRAGMA_UNROLL(ItemsPerThread)
for (int step = 0; step < steps; ++step)
{
const bool advance_haystack =
(needles_remaining == 0)
|| (haystack_remaining > 0 && Mode::should_advance(storage.haystack[i], storage.needles[j], compare_op));
if (advance_haystack)
{
++i;
--haystack_remaining;
}
else
{
using output_value_t = cub::detail::non_void_value_t<OutputIt, Offset>;
d_output[values_beg + j] = static_cast<output_value_t>(range_beg + i);
++j;
--needles_remaining;
}
}
}
_CCCL_DEVICE_API _CCCL_FORCEINLINE void operator()()
{
const int tile_idx = static_cast<int>(blockIdx.x);
const Offset diag0 = static_cast<Offset>(tile_size) * tile_idx;
const Offset diag1 = ::cuda::std::min(diag0 + static_cast<Offset>(tile_size), range_count + values_count);
const int total_in_tile = static_cast<int>(diag1 - diag0);
if (total_in_tile == tile_size)
{
consume_tile</* IsFullTile = */ true>(tile_idx, diag0, tile_size);
}
else
{
consume_tile</* IsFullTile = */ false>(tile_idx, diag0, total_in_tile);
}
}
};
} // namespace detail::find_bound_sorted_values
CUB_NAMESPACE_END

View File

@@ -0,0 +1,56 @@
// SPDX-FileCopyrightText: Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/util_ptx.cuh>
#include <cub/util_type.cuh>
CUB_NAMESPACE_BEGIN
namespace detail::for_each
{
template <int ThreadsPerBlock, int ItemsPerThread>
struct policy_t
{
static constexpr int threads_per_block = ThreadsPerBlock;
static constexpr int items_per_thread = ItemsPerThread;
};
template <class PolicyT, class OffsetT, class OpT>
struct agent_block_striped_t
{
static constexpr int items_per_thread = PolicyT::items_per_thread;
OffsetT tile_base;
OpT op;
template <bool IsFullTile>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile(int items_in_tile, int threads_per_block)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < items_per_thread; item++)
{
const auto idx =
static_cast<OffsetT>(threads_per_block * item + threadIdx.x); // NOLINT(bugprone-misplaced-widening-cast)
if (IsFullTile || idx < items_in_tile)
{
(void) op(tile_base + idx);
}
}
}
};
} // namespace detail::for_each
CUB_NAMESPACE_END

View File

@@ -0,0 +1,725 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2018, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
//! \file
//! cub::AgentHistogram implements a stateful abstraction of CUDA thread blocks for participating in device-wide
//! histogram.
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/grid/grid_queue.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_type.cuh>
#include <cuda/std/__concepts/same_as.h>
#include <cuda/std/__fwd/format.h>
#include <cuda/std/__host_stdlib/ostream>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/integral_constant.h>
#include <cuda/std/__type_traits/is_pointer.h>
#include <cuda/std/cstdint>
CUB_NAMESPACE_BEGIN
enum BlockHistogramMemoryPreference
{
GMEM,
SMEM,
BLEND
};
#if _CCCL_HOSTED()
namespace detail
{
[[nodiscard]] _CCCL_HOST_DEVICE_API constexpr const char* to_string(BlockHistogramMemoryPreference mempref) noexcept
{
switch (mempref)
{
case GMEM:
return "GMEM";
case SMEM:
return "SMEM";
case BLEND:
return "BLEND";
}
return "<unknown BlockHistogramMemoryPreference>";
}
} // namespace detail
inline ::std::ostream& operator<<(::std::ostream& os, BlockHistogramMemoryPreference mempref)
{
return os << CUB_NS_QUALIFIER::detail::to_string(mempref);
}
#endif // _CCCL_HOSTED()
CUB_NAMESPACE_END
#if __cpp_lib_format >= 201907L && !defined(_CCCL_DOXYGEN_INVOKED)
template <::cuda::std::same_as<char> CharT>
struct std::formatter<CUB_NS_QUALIFIER::BlockHistogramMemoryPreference, CharT> : formatter<const CharT*, CharT>
{
template <class FmtCtx>
auto format(const CUB_NS_QUALIFIER::BlockHistogramMemoryPreference& mempref, FmtCtx& ctx) const
{
return formatter<const CharT*, CharT>::format(CUB_NS_QUALIFIER::detail::to_string(mempref), ctx);
}
};
#endif // __cpp_lib_format >= 201907L && !defined(_CCCL_DOXYGEN_INVOKED)
CUB_NAMESPACE_BEGIN
namespace detail
{
//! Parameterizable tuning policy type for AgentHistogram
template <int ThreadsPerBlock,
int PixelsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
bool RleCompress,
BlockHistogramMemoryPreference MemoryPreference,
bool WorkStealing,
int VecSize = 4>
struct agent_histogram_policy
{
/// Threads per thread block
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
/// Pixels per thread (per tile of input)
static constexpr int PIXELS_PER_THREAD = PixelsPerThread;
/// Whether to perform localized RLE to compress samples before histogramming
static constexpr bool IS_RLE_COMPRESS = RleCompress;
/// Whether to prefer privatized shared-memory bins (versus privatized global-memory bins)
static constexpr BlockHistogramMemoryPreference MEM_PREFERENCE = MemoryPreference;
/// Whether to dequeue tiles from a global work queue
static constexpr bool IS_WORK_STEALING = WorkStealing;
/// Vector size for samples loading (1, 2, 4)
static constexpr int VEC_SIZE = VecSize;
static_assert(VEC_SIZE == 1 || VEC_SIZE == 2 || VEC_SIZE == 4);
///< The BlockLoad algorithm to use
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
///< Cache load modifier for reading input elements
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int PixelsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
bool RleCompress,
BlockHistogramMemoryPreference MemoryPreference,
bool WorkStealing,
int VecSize = 4>
using AgentHistogramPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceHistogram") = detail::agent_histogram_policy<
ThreadsPerBlock,
PixelsPerThread,
LoadAlgorithm,
LoadModifier,
RleCompress,
MemoryPreference,
WorkStealing,
VecSize>;
namespace detail::histogram
{
// Return a native pixel pointer (specialized for CacheModifiedInputIterator types)
template <CacheLoadModifier Modifier, typename ValueT, typename OffsetT>
_CCCL_DEVICE _CCCL_FORCEINLINE auto NativePointer(CacheModifiedInputIterator<Modifier, ValueT, OffsetT> itr)
{
return itr.ptr;
}
// Return a native pixel pointer (specialized for other types)
template <typename IteratorT>
_CCCL_DEVICE _CCCL_FORCEINLINE auto NativePointer(IteratorT itr)
{
return nullptr;
}
//! @brief AgentHistogram implements a stateful abstraction of CUDA thread blocks for participating
//! in device-wide histogram .
//!
//! @tparam AgentHistogramPolicyT
//! Parameterized AgentHistogramPolicy tuning policy type
//!
//! @tparam PrivatizedSmemBins
//! Number of privatized shared-memory histogram bins of any channel. Zero indicates privatized
//! counters to be maintained in device-accessible memory.
//!
//! @tparam NumChannels
//! Number of channels interleaved in the input data. Supports up to four channels.
//!
//! @tparam NumActiveChannels
//! Number of channels actively being histogrammed
//!
//! @tparam SampleIteratorT
//! Random-access input iterator type for reading samples
//!
//! @tparam CounterT
//! Integer type for counting sample occurrences per histogram bin
//!
//! @tparam PrivatizedDecodeOpT
//! The transform operator type for determining privatized counter indices from samples, one for
//! each channel
//!
//! @tparam OutputDecodeOpT
//! The transform operator type for determining output bin-ids from privatized counter indices, one
//! for each channel
//!
//! @tparam OffsetT
//! Signed integer type for global offsets
template <typename AgentHistogramPolicyT,
int PrivatizedSmemBins,
int NumChannels,
int NumActiveChannels,
typename SampleIteratorT,
typename CounterT,
typename PrivatizedDecodeOpT,
typename OutputDecodeOpT,
typename OffsetT>
struct AgentHistogram
{
static constexpr int vec_size = AgentHistogramPolicyT::VEC_SIZE;
static constexpr int threads_per_block = AgentHistogramPolicyT::BLOCK_THREADS;
static constexpr int pixels_per_thread = AgentHistogramPolicyT::PIXELS_PER_THREAD;
static constexpr int samples_per_thread = pixels_per_thread * NumChannels;
static constexpr int vecs_per_thread = samples_per_thread / vec_size;
static constexpr int tile_pixels = pixels_per_thread * threads_per_block;
static constexpr int tile_samples = samples_per_thread * threads_per_block;
static constexpr bool is_rle_compress = AgentHistogramPolicyT::IS_RLE_COMPRESS;
static constexpr bool is_work_stealing = AgentHistogramPolicyT::IS_WORK_STEALING;
static constexpr CacheLoadModifier load_modifier = AgentHistogramPolicyT::LOAD_MODIFIER;
static constexpr auto mem_preference =
(PrivatizedSmemBins > 0) ? BlockHistogramMemoryPreference{AgentHistogramPolicyT::MEM_PREFERENCE} : GMEM;
using SampleT = it_value_t<SampleIteratorT>;
using PixelT = typename CubVector<SampleT, NumChannels>::Type;
using VecT = typename CubVector<SampleT, vec_size>::Type;
/// Input iterator wrapper type (for applying cache modifier)
// Wrap the native input pointer with CacheModifiedInputIterator or directly use the supplied input iterator type
// TODO(bgruber): we can wrap all contiguous iterators, not just pointers
using WrappedSampleIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<SampleIteratorT>,
CacheModifiedInputIterator<load_modifier, SampleT, OffsetT>,
SampleIteratorT>;
using WrappedPixelIteratorT = CacheModifiedInputIterator<load_modifier, PixelT, OffsetT>;
using WrappedVecsIteratorT = CacheModifiedInputIterator<load_modifier, VecT, OffsetT>;
using BlockLoadSampleT =
BlockLoad<SampleT, threads_per_block, samples_per_thread, AgentHistogramPolicyT::LOAD_ALGORITHM>;
using BlockLoadPixelT =
BlockLoad<PixelT, threads_per_block, pixels_per_thread, AgentHistogramPolicyT::LOAD_ALGORITHM>;
using BlockLoadVecT = BlockLoad<VecT, threads_per_block, vecs_per_thread, AgentHistogramPolicyT::LOAD_ALGORITHM>;
struct _TempStorage
{
// Smem needed for block-privatized smem histogram (with 1 word of padding)
CounterT histograms[NumActiveChannels][PrivatizedSmemBins + 1];
int tile_idx;
union
{
typename BlockLoadSampleT::TempStorage sample_load;
typename BlockLoadPixelT::TempStorage pixel_load;
typename BlockLoadVecT::TempStorage vec_load;
};
};
using TempStorage = Uninitialized<_TempStorage>;
_TempStorage& temp_storage;
WrappedSampleIteratorT d_wrapped_samples; // with cache modifier applied, if possible
SampleT* d_native_samples; // possibly nullptr if unavailable
const int* num_output_bins; // one for each channel
const int* num_privatized_bins; // one for each channel
CounterT* d_privatized_histograms[NumActiveChannels]; // one for each channel
CounterT** d_output_histograms; // in global memory
const OutputDecodeOpT* output_decode_op; // determines output bin-id from privatized counter index, one for each
// channel
const PrivatizedDecodeOpT* privatized_decode_op; // determines privatized counter index from sample, one for each
// channel
bool prefer_smem; // for privatized counterss
template <typename TwoDimSubscriptableCounterT>
_CCCL_DEVICE _CCCL_FORCEINLINE void ZeroBinCounters(TwoDimSubscriptableCounterT& privatized_histograms)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ch = 0; ch < NumActiveChannels; ++ch)
{
for (int bin = static_cast<int>(threadIdx.x); bin < num_privatized_bins[ch]; bin += threads_per_block)
{
privatized_histograms[ch][bin] = 0;
}
}
// TODO(bgruber): do we also need the __syncthreads() when prefer_smem is false?
// Barrier to make sure all threads are done updating counters
__syncthreads();
}
// Update final output histograms from privatized histograms
template <typename TwoDimSubscriptableCounterT>
_CCCL_DEVICE _CCCL_FORCEINLINE void StoreOutput(TwoDimSubscriptableCounterT& privatized_histograms)
{
// Barrier to make sure all threads are done updating counters
__syncthreads();
// Apply privatized bin counts to output bin counts
_CCCL_PRAGMA_UNROLL_FULL()
for (int ch = 0; ch < NumActiveChannels; ++ch)
{
const int channel_bins = num_privatized_bins[ch];
for (int bin = static_cast<int>(threadIdx.x); bin < channel_bins; bin += threads_per_block)
{
int output_bin = -1;
const CounterT count = privatized_histograms[ch][bin];
const bool is_valid = count > 0;
output_decode_op[ch].template BinSelect<load_modifier>(static_cast<SampleT>(bin), output_bin, is_valid);
if (output_bin >= 0)
{
atomicAdd(&d_output_histograms[ch][output_bin], count);
}
}
}
}
// Accumulate pixels. Specialized for RLE compression.
template <typename TwoDimSubscriptableCounterT>
_CCCL_DEVICE _CCCL_FORCEINLINE void AccumulatePixels(
SampleT samples[pixels_per_thread][NumChannels],
bool is_valid[pixels_per_thread],
TwoDimSubscriptableCounterT& privatized_histograms,
::cuda::std::true_type is_rle_compress)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ch = 0; ch < NumActiveChannels; ++ch)
{
// Bin pixels
int bins[pixels_per_thread];
_CCCL_PRAGMA_UNROLL_FULL()
for (int pixel = 0; pixel < pixels_per_thread; ++pixel)
{
bins[pixel] = -1;
privatized_decode_op[ch].template BinSelect<load_modifier>(samples[pixel][ch], bins[pixel], is_valid[pixel]);
}
CounterT accumulator = 1;
_CCCL_PRAGMA_UNROLL_FULL()
for (int pixel = 0; pixel < pixels_per_thread - 1; ++pixel)
{
if (bins[pixel] != bins[pixel + 1])
{
if (bins[pixel] >= 0)
{
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_60,
(atomicAdd_block(privatized_histograms[ch] + bins[pixel], accumulator);),
(atomicAdd(privatized_histograms[ch] + bins[pixel], accumulator);));
}
accumulator = 0;
}
accumulator++;
}
// Last pixel
if (bins[pixels_per_thread - 1] >= 0)
{
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_60,
(atomicAdd_block(privatized_histograms[ch] + bins[pixels_per_thread - 1], accumulator);),
(atomicAdd(privatized_histograms[ch] + bins[pixels_per_thread - 1], accumulator);));
}
}
}
// Accumulate pixels. Specialized for individual accumulation of each pixel.
template <typename TwoDimSubscriptableCounterT>
_CCCL_DEVICE _CCCL_FORCEINLINE void AccumulatePixels(
SampleT samples[pixels_per_thread][NumChannels],
bool is_valid[pixels_per_thread],
TwoDimSubscriptableCounterT& privatized_histograms,
::cuda::std::false_type is_rle_compress)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int pixel = 0; pixel < pixels_per_thread; ++pixel)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ch = 0; ch < NumActiveChannels; ++ch)
{
int bin = -1;
privatized_decode_op[ch].template BinSelect<load_modifier>(samples[pixel][ch], bin, is_valid[pixel]);
if (bin >= 0)
{
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_60,
(atomicAdd_block(privatized_histograms[ch] + bin, 1);),
(atomicAdd(privatized_histograms[ch] + bin, 1);));
}
}
}
}
// Load full, aligned tile using pixel iterator
_CCCL_DEVICE _CCCL_FORCEINLINE void
LoadFullAlignedTile(OffsetT block_offset, SampleT (&samples)[pixels_per_thread][NumChannels])
{
if constexpr (NumActiveChannels == 1)
{
using AliasedVecs = VecT[vecs_per_thread];
WrappedVecsIteratorT d_wrapped_vecs(reinterpret_cast<VecT*>(d_native_samples + block_offset));
// Load using a wrapped vec iterator
BlockLoadVecT{temp_storage.vec_load}.Load(d_wrapped_vecs, reinterpret_cast<AliasedVecs&>(samples));
}
else
{
using AliasedPixels = PixelT[pixels_per_thread];
WrappedPixelIteratorT d_wrapped_pixels(reinterpret_cast<PixelT*>(d_native_samples + block_offset));
// Load using a wrapped pixel iterator
BlockLoadPixelT{temp_storage.pixel_load}.Load(d_wrapped_pixels, reinterpret_cast<AliasedPixels&>(samples));
}
}
template <bool IsFullTile, bool IsAligned>
_CCCL_DEVICE _CCCL_FORCEINLINE void
LoadTile(OffsetT block_offset, int valid_samples, SampleT (&samples)[pixels_per_thread][NumChannels])
{
if constexpr (IsFullTile)
{
if constexpr (IsAligned)
{
LoadFullAlignedTile(block_offset, samples);
}
else
{
// Load using sample iterator
using AliasedSamples = SampleT[samples_per_thread];
BlockLoadSampleT{temp_storage.sample_load}.Load(
d_wrapped_samples + block_offset, reinterpret_cast<AliasedSamples&>(samples));
}
}
else
{
if constexpr (IsAligned)
{
// Load partially-full, aligned tile using the pixel iterator
using AliasedPixels = PixelT[pixels_per_thread];
WrappedPixelIteratorT d_wrapped_pixels((PixelT*) (d_native_samples + block_offset));
int valid_pixels = valid_samples / NumChannels;
// Load using a wrapped pixel iterator
BlockLoadPixelT{temp_storage.pixel_load}.Load(
d_wrapped_pixels, reinterpret_cast<AliasedPixels&>(samples), valid_pixels);
}
else
{
using AliasedSamples = SampleT[samples_per_thread];
BlockLoadSampleT{temp_storage.sample_load}.Load(
d_wrapped_samples + block_offset, reinterpret_cast<AliasedSamples&>(samples), valid_samples);
}
}
}
template <bool IsFullTile, bool IsStriped>
_CCCL_DEVICE _CCCL_FORCEINLINE void MarkValid(bool (&is_valid)[pixels_per_thread], int valid_samples)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int pixel = 0; pixel < pixels_per_thread; ++pixel)
{
if constexpr (IsStriped)
{
is_valid[pixel] = IsFullTile || (((threadIdx.x + threads_per_block * pixel) * NumChannels) < valid_samples);
}
else
{
is_valid[pixel] = IsFullTile || (((threadIdx.x * pixels_per_thread + pixel) * NumChannels) < valid_samples);
}
}
}
//! @brief Consume a tile of data samples
//!
//! @tparam IsAligned
//! Whether the tile offset is aligned (vec-aligned for single-channel, pixel-aligned for multi-channel)
//!
//! @tparam IsFullTile
//! Whether the tile is full
template <bool IsAligned, bool IsFullTile>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeTile(OffsetT block_offset, int valid_samples)
{
SampleT samples[pixels_per_thread][NumChannels];
bool is_valid[pixels_per_thread];
LoadTile<IsFullTile, IsAligned>(block_offset, valid_samples, samples);
MarkValid<IsFullTile, AgentHistogramPolicyT::LOAD_ALGORITHM == BLOCK_LOAD_STRIPED>(is_valid, valid_samples);
if (prefer_smem)
{
AccumulatePixels(samples, is_valid, temp_storage.histograms, ::cuda::std::bool_constant<is_rle_compress>{});
}
else
{
AccumulatePixels(samples, is_valid, d_privatized_histograms, ::cuda::std::bool_constant<is_rle_compress>{});
}
}
//! @brief Consume row tiles. Specialized for work-stealing from queue
//!
//! @param num_row_pixels
//! The number of multi-channel pixels per row in the region of interest
//!
//! @param num_rows
//! The number of rows in the region of interest
//!
//! @param row_stride_samples
//! The number of samples between starts of consecutive rows in the region of interest
//!
//! @param tiles_per_row
//! Number of image tiles per row
template <bool IsAligned>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeTiles(
OffsetT num_row_pixels,
OffsetT num_rows,
OffsetT row_stride_samples,
int tiles_per_row,
GridQueue<int> tile_queue,
::cuda::std::true_type is_work_stealing)
{
int num_tiles = num_rows * tiles_per_row;
int tile_idx = static_cast<int>((blockIdx.y * gridDim.x) + blockIdx.x);
OffsetT num_even_share_tiles = gridDim.x * gridDim.y;
while (tile_idx < num_tiles)
{
int row = tile_idx / tiles_per_row;
int col = tile_idx - (row * tiles_per_row);
OffsetT row_offset = row * row_stride_samples;
OffsetT col_offset = (col * tile_samples);
OffsetT tile_offset = row_offset + col_offset;
if (col == tiles_per_row - 1)
{
// Consume a partially-full tile at the end of the row
OffsetT num_remaining = (num_row_pixels * NumChannels) - col_offset;
ConsumeTile<IsAligned, false>(tile_offset, num_remaining);
}
else
{
// Consume full tile
ConsumeTile<IsAligned, true>(tile_offset, tile_samples);
}
__syncthreads();
// Get next tile
if (threadIdx.x == 0)
{
temp_storage.tile_idx = tile_queue.Drain(1) + num_even_share_tiles;
}
__syncthreads();
tile_idx = temp_storage.tile_idx;
}
}
//! @brief Consume row tiles. Specialized for even-share (striped across thread blocks)
//!
//! @param num_row_pixels
//! The number of multi-channel pixels per row in the region of interest
//!
//! @param num_rows
//! The number of rows in the region of interest
//!
//! @param row_stride_samples
//! The number of samples between starts of consecutive rows in the region of interest
template <bool IsAligned>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeTiles(
OffsetT num_row_pixels, OffsetT num_rows, OffsetT row_stride_samples, int, GridQueue<int>, ::cuda::std::false_type)
{
for (int row = static_cast<int>(blockIdx.y); row < num_rows; row += static_cast<int>(gridDim.y))
{
OffsetT row_begin = row * row_stride_samples;
OffsetT row_end = row_begin + (num_row_pixels * NumChannels);
OffsetT tile_offset = row_begin + (blockIdx.x * tile_samples);
while (tile_offset < row_end)
{
OffsetT num_remaining = row_end - tile_offset;
if (num_remaining < tile_samples)
{
// Consume partial tile
ConsumeTile<IsAligned, false>(tile_offset, num_remaining);
break;
}
// Consume full tile
ConsumeTile<IsAligned, true>(tile_offset, tile_samples);
tile_offset += gridDim.x * tile_samples;
}
}
}
//---------------------------------------------------------------------
// Parameter extraction
//---------------------------------------------------------------------
//! @brief Constructor
//!
//! @param temp_storage
//! Reference to temp_storage
//!
//! @param d_samples
//! Input data to reduce
//!
//! @param num_output_bins
//! The number bins per final output histogram
//!
//! @param num_privatized_bins
//! The number bins per privatized histogram
//!
//! @param d_output_histograms
//! Reference to final output histograms
//!
//! @param d_privatized_histograms
//! Reference to privatized histograms
//!
//! @param output_decode_op
//! The transform operator for determining output bin-ids from privatized counter indices, one for each channel
//!
//! @param privatized_decode_op
//! The transform operator for determining privatized counter indices from samples, one for each channel
_CCCL_DEVICE _CCCL_FORCEINLINE AgentHistogram(
TempStorage& temp_storage,
SampleIteratorT d_samples,
const int* num_output_bins,
const int* num_privatized_bins,
CounterT** d_output_histograms,
CounterT** d_privatized_histograms,
const OutputDecodeOpT* output_decode_op,
const PrivatizedDecodeOpT* privatized_decode_op)
: temp_storage(temp_storage.Alias())
, d_wrapped_samples(d_samples)
, d_native_samples(NativePointer(d_wrapped_samples))
, num_output_bins(num_output_bins)
, num_privatized_bins(num_privatized_bins)
, d_output_histograms(d_output_histograms)
, output_decode_op(output_decode_op)
, privatized_decode_op(privatized_decode_op)
, prefer_smem((mem_preference == SMEM) ? true : // prefer smem privatized histograms
(mem_preference == GMEM) ? false
: // prefer gmem privatized histograms
blockIdx.x & 1) // prefer blended privatized histograms
{
const int blockId = static_cast<int>((blockIdx.y * gridDim.x) + blockIdx.x);
// TODO(bgruber): d_privatized_histograms seems only used when !prefer_smem, can we skip it if prefer_smem?
// Initialize the locations of this block's privatized histograms
for (int ch = 0; ch < NumActiveChannels; ++ch)
{
const auto offset = static_cast<::cuda::std::int64_t>(blockId) * num_privatized_bins[ch];
this->d_privatized_histograms[ch] = d_privatized_histograms[ch] + offset;
}
}
//! @brief Consume image
//!
//! @param num_row_pixels
//! The number of multi-channel pixels per row in the region of interest
//!
//! @param num_rows
//! The number of rows in the region of interest
//!
//! @param row_stride_samples
//! The number of samples between starts of consecutive rows in the region of interest
//!
//! @param tiles_per_row
//! Number of image tiles per row
//!
//! @param tile_queue
//! Queue descriptor for assigning tiles of work to thread blocks
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeTiles(
OffsetT num_row_pixels, OffsetT num_rows, OffsetT row_stride_samples, int tiles_per_row, GridQueue<int> tile_queue)
{
// Check whether all row starting offsets are vec-aligned (in single-channel) or pixel-aligned (in multi-channel)
constexpr int vec_mask = alignof(VecT) - 1;
constexpr int pixel_mask = alignof(PixelT) - 1;
const size_t row_bytes = sizeof(SampleT) * row_stride_samples;
const bool vec_aligned_rows =
(NumChannels == 1) && (samples_per_thread % vec_size == 0) && // Single channel
((size_t(d_native_samples) & vec_mask) == 0) && // ptr is quad-aligned
((num_rows == 1) || ((row_bytes & vec_mask) == 0)); // number of row-samples is a multiple of the alignment of the
// quad
const bool pixel_aligned_rows =
(NumChannels > 1) && // Multi channel
((size_t(d_native_samples) & pixel_mask) == 0) && // ptr is pixel-aligned
((row_bytes & pixel_mask) == 0); // number of row-samples is a multiple of the alignment of the pixel
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
// Whether rows are aligned and can be vectorized
if ((d_native_samples != nullptr) && (vec_aligned_rows || pixel_aligned_rows))
{
ConsumeTiles<true>(
num_row_pixels, num_rows, row_stride_samples, tiles_per_row, tile_queue, bool_constant_v<is_work_stealing>);
}
else
{
ConsumeTiles<false>(
num_row_pixels, num_rows, row_stride_samples, tiles_per_row, tile_queue, bool_constant_v<is_work_stealing>);
}
_CCCL_PDL_TRIGGER_NEXT_LAUNCH(); // omitting makes no difference in cub.bench.histogram.even.base
}
//! Initialize privatized bin counters. Specialized for privatized shared-memory counters
_CCCL_DEVICE _CCCL_FORCEINLINE void InitBinCounters()
{
if (prefer_smem)
{
ZeroBinCounters(temp_storage.histograms);
}
else
{
ZeroBinCounters(d_privatized_histograms);
}
}
//! Store privatized histogram to device-accessible memory. Specialized for privatized shared-memory counters
_CCCL_DEVICE _CCCL_FORCEINLINE void StoreOutput()
{
if (prefer_smem)
{
StoreOutput(temp_storage.histograms);
}
else
{
StoreOutput(d_privatized_histograms);
}
}
};
} // namespace detail::histogram
CUB_NAMESPACE_END

View File

@@ -0,0 +1,335 @@
// SPDX-FileCopyrightText: Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/agent_merge_sort.cuh>
#include <cub/block/block_load_to_shared.cuh>
#include <cub/block/block_merge_sort.cuh>
#include <cub/block/block_store.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_namespace.cuh>
#include <cub/util_type.cuh>
#include <thrust/type_traits/is_contiguous_iterator.h>
#include <thrust/type_traits/is_trivially_relocatable.h>
#include <thrust/type_traits/unwrap_contiguous_iterator.h>
#include <cuda/std/__algorithm/min.h>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/is_same.h>
#include <cuda/std/span>
CUB_NAMESPACE_BEGIN
namespace detail::merge
{
// TODO(bgruber): can we unify this one with AgentMerge in agent_merge_sort.cuh?
// TODO(bgruber): pass a merge_policy by value instead of individual template parameters in C++20
template <int ThreadsPerBlock,
int ItemsPerThread,
CacheLoadModifier LoadModifier,
BlockStoreAlgorithm StoreAlgorithm,
bool UseBl2ShForKeys,
bool UseBl2ShForItems,
bool Unroll,
typename KeysIt1,
typename ItemsIt1,
typename KeysIt2,
typename ItemsIt2,
typename KeysOutputIt,
typename ItemsOutputIt,
typename Offset,
typename CompareOp>
struct agent_t
{
static constexpr int threads_per_block = ThreadsPerBlock; // also used for kernel launch bounds and dispatch logic
static constexpr int items_per_tile = ItemsPerThread * ThreadsPerBlock; // also used by dispatch logic
// key and value type are taken from the first input sequence (consistent with old Thrust behavior)
using key_type = it_value_t<KeysIt1>;
using item_type = it_value_t<ItemsIt1>;
using block_load_to_shared = BlockLoadToShared<ThreadsPerBlock>;
using block_store_keys = BlockStore<key_type, ThreadsPerBlock, ItemsPerThread, StoreAlgorithm>;
using block_store_items = BlockStore<item_type, ThreadsPerBlock, ItemsPerThread, StoreAlgorithm>;
static constexpr int bl2sh_minimum_align = cub::detail::LoadToSharedBufferAlignBytes<char>();
template <typename ValueT>
struct alignas(cub::detail::LoadToSharedBufferAlignBytes<ValueT>()) buffer_t
{
// Need extra bytes of padding for TMA because this static buffer has to hold the two dynamically sized buffers.
static constexpr int bytes_needed = cub::detail::LoadToSharedBufferSizeBytes<ValueT>(items_per_tile + 1ULL)
+ (alignof(ValueT) < bl2sh_minimum_align ? 2 * bl2sh_minimum_align : 0);
char c_array[bytes_needed];
};
struct temp_storages_without_bl2sh
{
using keys_smem = ::cuda::std::conditional_t<UseBl2ShForKeys, buffer_t<key_type>, key_type[items_per_tile + 1]>;
using items_smem = ::cuda::std::conditional_t<UseBl2ShForItems, buffer_t<item_type>, item_type[items_per_tile + 1]>;
union
{
typename block_store_keys::TempStorage store_keys;
typename block_store_items::TempStorage store_items;
keys_smem keys_shared;
items_smem items_shared;
};
};
// inherit from data storage, so it's positioned at the start of the shared memory
struct temp_storages_with_bl2sh : temp_storages_without_bl2sh
{
typename block_load_to_shared::TempStorage load2sh;
};
using temp_storages = ::cuda::std::
conditional_t<UseBl2ShForKeys || UseBl2ShForItems, temp_storages_with_bl2sh, temp_storages_without_bl2sh>;
using TempStorage = Uninitialized<temp_storages>;
// Per thread data
temp_storages& storage;
KeysIt1 keys1_in;
ItemsIt1 items1_in;
Offset keys1_count;
KeysIt2 keys2_in;
ItemsIt2 items2_in;
Offset keys2_count;
KeysOutputIt keys_out;
ItemsOutputIt items_out;
CompareOp compare_op;
Offset* key1_beg_offsets;
template <bool IsFullTile>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile(Offset tile_idx, Offset tile_base, int num_remaining)
{
const Offset diag0 = items_per_tile * tile_idx;
Offset diag1 = diag0 + items_per_tile;
if constexpr (IsFullTile)
{
_CCCL_ASSERT(diag1 <= keys1_count + keys2_count, "");
}
else
{
diag1 = keys1_count + keys2_count;
}
// compute bounding box for keys1 & keys2
const Offset keys1_beg = key1_beg_offsets[tile_idx + 0];
const Offset keys1_end = key1_beg_offsets[tile_idx + 1];
const Offset keys2_beg = diag0 - keys1_beg;
const Offset keys2_end = diag1 - keys1_end;
// number of keys per tile
const int keys1_count_tile = static_cast<int>(keys1_end - keys1_beg);
const int keys2_count_tile = static_cast<int>(keys2_end - keys2_beg);
if constexpr (IsFullTile) // NOLINT(bugprone-branch-clone)
{
_CCCL_ASSERT(keys1_count_tile + keys2_count_tile == items_per_tile, "");
}
else
{
_CCCL_ASSERT(keys1_count_tile + keys2_count_tile == num_remaining, "");
}
[[maybe_unused]] auto load2sh = [&] {
if constexpr (UseBl2ShForKeys || UseBl2ShForItems)
{
return block_load_to_shared{storage.load2sh};
}
else
{
return NullType{};
}
}();
key_type keys_loc[ItemsPerThread];
key_type* keys1_shared;
key_type* keys2_shared;
int keys2_offset;
if constexpr (UseBl2ShForKeys)
{
::cuda::std::span keys1_src{THRUST_NS_QUALIFIER::unwrap_contiguous_iterator(keys1_in + keys1_beg),
static_cast<::cuda::std::size_t>(keys1_count_tile)};
::cuda::std::span keys2_src{THRUST_NS_QUALIFIER::unwrap_contiguous_iterator(keys2_in + keys2_beg),
static_cast<::cuda::std::size_t>(keys2_count_tile)};
::cuda::std::span keys_buffers{storage.keys_shared.c_array};
auto keys1_buffer = keys_buffers.first(cub::detail::LoadToSharedBufferSizeBytes<key_type>(keys1_count_tile));
auto keys2_buffer = keys_buffers.last(cub::detail::LoadToSharedBufferSizeBytes<key_type>(keys2_count_tile));
_CCCL_ASSERT(keys1_buffer.end() <= keys2_buffer.begin(),
"Keys buffer needs to be appropriately sized (internal)");
keys1_shared = data(load2sh.CopyAsync(keys1_buffer, keys1_src));
keys2_shared = data(load2sh.CopyAsync(keys2_buffer, keys2_src));
auto token = load2sh.Commit();
// Needed for using keys1_shared as one big buffer including both ranges in SerialMerge
keys2_offset = static_cast<int>(keys2_shared - keys1_shared);
load2sh.Wait(::cuda::std::move(token));
}
else
{
auto keys1_in_cm = try_make_cache_modified_iterator<LoadModifier>(keys1_in);
auto keys2_in_cm = try_make_cache_modified_iterator<LoadModifier>(keys2_in);
merge_sort::gmem_to_reg<ThreadsPerBlock, IsFullTile>(
keys_loc, keys1_in_cm + keys1_beg, keys2_in_cm + keys2_beg, keys1_count_tile, keys2_count_tile);
keys1_shared = &storage.keys_shared[0];
// Needed for using keys1_shared as one big buffer including both ranges in SerialMerge
keys2_offset = keys1_count_tile;
keys2_shared = keys1_shared + keys2_offset;
merge_sort::reg_to_shared<ThreadsPerBlock>(keys1_shared, keys_loc);
__syncthreads();
}
// Now find the merge path for each of the threads.
// We can use int type here, because the number of items in shared memory is limited.
int diag0_thread = ItemsPerThread * static_cast<int>(threadIdx.x);
if constexpr (IsFullTile)
{
_CCCL_ASSERT(num_remaining == items_per_tile, "");
_CCCL_ASSERT(diag0_thread < num_remaining, "");
}
else
{ // for partial tiles, clamp the thread diagonal to the valid items
diag0_thread = (::cuda::std::min) (diag0_thread, num_remaining);
}
const int keys1_beg_thread =
MergePath(keys1_shared, keys2_shared, keys1_count_tile, keys2_count_tile, diag0_thread, compare_op);
const int keys2_beg_thread = diag0_thread - keys1_beg_thread;
const int keys1_count_thread = keys1_count_tile - keys1_beg_thread;
const int keys2_count_thread = keys2_count_tile - keys2_beg_thread;
// perform serial merge
int indices[ItemsPerThread];
cub::detail::serial_merge<Unroll>(
keys1_shared,
keys1_beg_thread,
keys2_offset + keys2_beg_thread,
keys1_count_thread,
keys2_count_thread,
keys_loc,
indices,
compare_op);
// write keys
__syncthreads(); // sync after reading from SMEM before so block store can use SMEM again
if constexpr (IsFullTile)
{
block_store_keys{storage.store_keys}.Store(keys_out + tile_base, keys_loc);
}
else
{
block_store_keys{storage.store_keys}.Store(keys_out + tile_base, keys_loc, num_remaining);
}
// if items are provided, merge them
static constexpr bool have_items = !::cuda::std::is_same_v<item_type, NullType>;
if constexpr (have_items)
{
// Both of these are only needed when either keys or items or both use BlockLoadToShared introducing padding (that
// can differ between the keys and items)
[[maybe_unused]] const auto translate_indices = [&](int items2_offset) -> void {
const int diff = items2_offset - keys2_offset;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < ItemsPerThread; ++i)
{
if (indices[i] >= keys2_offset)
{
indices[i] += diff;
}
}
};
// WAR for MSVC erroring ("declared but never referenced") despite [[maybe_unused]]
(void) translate_indices;
item_type items_loc[ItemsPerThread];
item_type* items1_shared;
if constexpr (UseBl2ShForItems)
{
::cuda::std::span items1_src{THRUST_NS_QUALIFIER::unwrap_contiguous_iterator(items1_in + keys1_beg),
static_cast<::cuda::std::size_t>(keys1_count_tile)};
::cuda::std::span items2_src{THRUST_NS_QUALIFIER::unwrap_contiguous_iterator(items2_in + keys2_beg),
static_cast<::cuda::std::size_t>(keys2_count_tile)};
::cuda::std::span items_buffers{storage.items_shared.c_array};
auto items1_buffer = items_buffers.first(cub::detail::LoadToSharedBufferSizeBytes<item_type>(keys1_count_tile));
auto items2_buffer = items_buffers.last(cub::detail::LoadToSharedBufferSizeBytes<item_type>(keys2_count_tile));
_CCCL_ASSERT(items1_buffer.end() <= items2_buffer.begin(),
"Items buffer needs to be appropriately sized (internal)");
// block_store_keys above uses shared memory, so make sure all threads are done before we write
__syncthreads();
items1_shared = data(load2sh.CopyAsync(items1_buffer, items1_src));
item_type* items2_shared = data(load2sh.CopyAsync(items2_buffer, items2_src));
auto token = load2sh.Commit();
const int items2_offset = static_cast<int>(items2_shared - items1_shared);
translate_indices(items2_offset);
load2sh.Wait(::cuda::std::move(token));
}
else
{
{
auto items1_in_cm = try_make_cache_modified_iterator<LoadModifier>(items1_in);
auto items2_in_cm = try_make_cache_modified_iterator<LoadModifier>(items2_in);
merge_sort::gmem_to_reg<ThreadsPerBlock, IsFullTile>(
items_loc, items1_in_cm + keys1_beg, items2_in_cm + keys2_beg, keys1_count_tile, keys2_count_tile);
__syncthreads(); // block_store_keys above uses SMEM, so make sure all threads are done before we write to it
items1_shared = &storage.items_shared[0];
if constexpr (UseBl2ShForKeys)
{
const int items2_offset = keys1_count_tile;
translate_indices(items2_offset);
}
merge_sort::reg_to_shared<ThreadsPerBlock>(items1_shared, items_loc);
__syncthreads();
}
}
// gather items from shared mem
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < ItemsPerThread; ++i)
{
items_loc[i] = items1_shared[indices[i]];
}
__syncthreads();
// write from reg to gmem
if constexpr (IsFullTile)
{
block_store_items{storage.store_items}.Store(items_out + tile_base, items_loc);
}
else
{
block_store_items{storage.store_items}.Store(items_out + tile_base, items_loc, num_remaining);
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void operator()()
{
const Offset tile_idx = blockIdx.x;
const Offset tile_base = tile_idx * items_per_tile;
const int items_in_tile =
static_cast<int>((::cuda::std::min) (static_cast<Offset>(items_per_tile), keys1_count + keys2_count - tile_base));
if (items_in_tile == items_per_tile)
{
consume_tile</* IsFullTile = */ true>(tile_idx, tile_base, items_per_tile);
}
else
{
consume_tile</* IsFullTile = */ false>(tile_idx, tile_base, items_in_tile);
}
}
};
} // namespace detail::merge
CUB_NAMESPACE_END

View File

@@ -0,0 +1,692 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/block_merge_sort.cuh>
#include <cub/block/block_store.cuh>
#include <cub/device/dispatch/tuning/tuning_merge_sort.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_namespace.cuh>
#include <cub/util_type.cuh>
#include <cuda/std/__algorithm/min.h>
CUB_NAMESPACE_BEGIN
namespace detail::merge_sort
{
template <typename PolicyGetter,
typename KeyInputIteratorT,
typename ValueInputIteratorT,
typename KeyIteratorT,
typename ValueIteratorT,
typename OffsetT,
typename CompareOpT,
typename KeyT,
typename ValueT>
struct AgentBlockSort
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
static constexpr bool KEYS_ONLY = ::cuda::std::is_same_v<ValueT, NullType>;
static constexpr MergeSortPolicy policy = PolicyGetter{}();
static constexpr int BLOCK_THREADS = policy.threads_per_block;
static constexpr int ITEMS_PER_THREAD = policy.items_per_thread;
static constexpr int ITEMS_PER_TILE = BLOCK_THREADS * ITEMS_PER_THREAD;
using BlockMergeSortT = BlockMergeSort<KeyT, BLOCK_THREADS, ITEMS_PER_THREAD, ValueT, 1, 1, policy.unroll>;
using KeysLoadIt = try_make_cache_modified_iterator_t<policy.load_modifier, KeyInputIteratorT>;
using ItemsLoadIt = try_make_cache_modified_iterator_t<policy.load_modifier, ValueInputIteratorT>;
using BlockLoadKeys = BlockLoad<it_value_t<KeysLoadIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.load_algorithm>;
using BlockLoadItems = BlockLoad<it_value_t<ItemsLoadIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.load_algorithm>;
using BlockStoreKeysIt =
BlockStore<it_value_t<KeyIteratorT>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreItemsIt =
BlockStore<it_value_t<ValueIteratorT>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreKeysRaw = BlockStore<KeyT, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreItemsRaw = BlockStore<ValueT, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
union _TempStorage
{
typename BlockLoadKeys::TempStorage load_keys;
typename BlockLoadItems::TempStorage load_items;
typename BlockStoreKeysIt::TempStorage store_keys_it;
typename BlockStoreItemsIt::TempStorage store_items_it;
typename BlockStoreKeysRaw::TempStorage store_keys_raw;
typename BlockStoreItemsRaw::TempStorage store_items_raw;
typename BlockMergeSortT::TempStorage block_merge;
};
/// Alias wrapper allowing storage to be unioned
using TempStorage = Uninitialized<_TempStorage>;
//---------------------------------------------------------------------
// Per thread data
//---------------------------------------------------------------------
bool ping;
_TempStorage& storage;
KeysLoadIt keys_in;
ItemsLoadIt items_in;
OffsetT keys_count;
KeyIteratorT keys_out_it;
ValueIteratorT items_out_it;
KeyT* keys_out_raw;
ValueT* items_out_raw;
CompareOpT compare_op;
_CCCL_DEVICE _CCCL_FORCEINLINE AgentBlockSort(
bool ping_,
TempStorage& storage_,
KeysLoadIt keys_in_,
ItemsLoadIt items_in_,
OffsetT keys_count_,
KeyIteratorT keys_out_it_,
ValueIteratorT items_out_it_,
KeyT* keys_out_raw_,
ValueT* items_out_raw_,
CompareOpT compare_op_)
: ping(ping_)
, storage(storage_.Alias())
, keys_in(keys_in_)
, items_in(items_in_)
, keys_count(keys_count_)
, keys_out_it(keys_out_it_)
, items_out_it(items_out_it_)
, keys_out_raw(keys_out_raw_)
, items_out_raw(items_out_raw_)
, compare_op(compare_op_)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
const auto tile_idx = static_cast<OffsetT>(blockIdx.x);
const auto num_tiles = static_cast<OffsetT>(gridDim.x);
const auto tile_base = tile_idx * ITEMS_PER_TILE;
const int items_in_tile = (::cuda::std::min) (static_cast<int>(keys_count - tile_base), int{ITEMS_PER_TILE});
if (tile_idx < num_tiles - 1)
{
consume_tile<false>(tile_base, ITEMS_PER_TILE);
}
else
{
consume_tile<true>(tile_base, items_in_tile);
}
}
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile(OffsetT tile_base, int num_remaining)
{
ValueT items_local[ITEMS_PER_THREAD];
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
if constexpr (!KEYS_ONLY)
{
if constexpr (IS_LAST_TILE)
{
BlockLoadItems(storage.load_items)
.Load(items_in + tile_base, items_local, num_remaining, *(items_in + tile_base));
}
else
{
BlockLoadItems(storage.load_items).Load(items_in + tile_base, items_local);
}
__syncthreads();
}
KeyT keys_local[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
BlockLoadKeys(storage.load_keys).Load(keys_in + tile_base, keys_local, num_remaining, *(keys_in + tile_base));
}
else
{
BlockLoadKeys(storage.load_keys).Load(keys_in + tile_base, keys_local);
}
__syncthreads();
_CCCL_PDL_TRIGGER_NEXT_LAUNCH();
if constexpr (IS_LAST_TILE)
{
BlockMergeSortT(storage.block_merge).Sort(keys_local, items_local, compare_op, num_remaining, keys_local[0]);
}
else
{
BlockMergeSortT(storage.block_merge).Sort(keys_local, items_local, compare_op);
}
__syncthreads();
if (ping)
{
if constexpr (IS_LAST_TILE)
{
BlockStoreKeysIt(storage.store_keys_it).Store(keys_out_it + tile_base, keys_local, num_remaining);
}
else
{
BlockStoreKeysIt(storage.store_keys_it).Store(keys_out_it + tile_base, keys_local);
}
if constexpr (!KEYS_ONLY)
{
__syncthreads();
if constexpr (IS_LAST_TILE)
{
BlockStoreItemsIt(storage.store_items_it).Store(items_out_it + tile_base, items_local, num_remaining);
}
else
{
BlockStoreItemsIt(storage.store_items_it).Store(items_out_it + tile_base, items_local);
}
}
}
else
{
if constexpr (IS_LAST_TILE)
{
BlockStoreKeysRaw(storage.store_keys_raw).Store(keys_out_raw + tile_base, keys_local, num_remaining);
}
else
{
BlockStoreKeysRaw(storage.store_keys_raw).Store(keys_out_raw + tile_base, keys_local);
}
if constexpr (!KEYS_ONLY)
{
__syncthreads();
if constexpr (IS_LAST_TILE)
{
BlockStoreItemsRaw(storage.store_items_raw).Store(items_out_raw + tile_base, items_local, num_remaining);
}
else
{
BlockStoreItemsRaw(storage.store_items_raw).Store(items_out_raw + tile_base, items_local);
}
}
}
}
};
/**
* \brief This agent is responsible for partitioning a merge path into equal segments
*
* There are two sorted arrays to be merged into one array. If the first array
* is partitioned between parallel workers by slicing it into ranges of equal
* size, there could be a significant workload imbalance. The imbalance is
* caused by the fact that the distribution of elements from the second array
* is unknown beforehand. Instead, the MergePath is partitioned between workers.
* This approach guarantees an equal amount of work being assigned to each worker.
*
* This approach is outlined in the paper:
* Odeh et al, "Merge Path - Parallel Merging Made Simple"
* doi:10.1109/IPDPSW.2012.202
*/
template <typename KeyIteratorT, typename OffsetT, typename CompareOpT, typename KeyT>
struct AgentPartition
{
bool ping;
KeyIteratorT keys_ping;
KeyT* keys_pong;
OffsetT keys_count;
OffsetT partition_idx;
OffsetT* merge_partitions;
CompareOpT compare_op;
OffsetT target_merged_tiles_number;
int items_per_tile;
OffsetT num_partitions;
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
const OffsetT merged_tiles_number = target_merged_tiles_number / 2;
// target_merged_tiles_number is a power of two.
const OffsetT mask = target_merged_tiles_number - 1;
// The first tile number in the tiles group being merged, equal to:
// target_merged_tiles_number * (partition_idx / target_merged_tiles_number)
const OffsetT list = ~mask & partition_idx;
const OffsetT start = items_per_tile * list;
const OffsetT size = items_per_tile * merged_tiles_number;
// Tile number within the tile group being merged, equal to:
// partition_idx / target_merged_tiles_number
const OffsetT local_tile_idx = mask & partition_idx;
const OffsetT keys1_beg = (::cuda::std::min) (keys_count, start);
const OffsetT keys1_end = (::cuda::std::min) (keys_count, detail::safe_add_bound_to_max(start, size));
const OffsetT keys2_beg = keys1_end;
const OffsetT keys2_end = (::cuda::std::min) (keys_count, detail::safe_add_bound_to_max(keys2_beg, size));
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
// The last partition (which is one-past-the-last-tile) is only to mark the end of keys1_end for the merge stage
if (partition_idx + 1 == num_partitions)
{
merge_partitions[partition_idx] = keys1_end;
}
else
{
const OffsetT partition_at = (::cuda::std::min) (keys2_end - keys1_beg, items_per_tile * local_tile_idx);
OffsetT partition_diag =
ping
? MergePath(keys_ping + keys1_beg,
keys_ping + keys2_beg,
keys1_end - keys1_beg,
keys2_end - keys2_beg,
partition_at,
compare_op)
: MergePath(keys_pong + keys1_beg,
keys_pong + keys2_beg,
keys1_end - keys1_beg,
keys2_end - keys2_beg,
partition_at,
compare_op);
merge_partitions[partition_idx] = keys1_beg + partition_diag;
}
// TODO(bgruber): looking at SASS triggering the next launch here just generates a lot of noise and the PRE-EXIT
// just ends of right before EXIT anyway. So let's omit it.
// _CCCL_PDL_TRIGGER_NEXT_LAUNCH();
}
};
/**
* \brief Concatenates up to ITEMS_PER_THREAD elements from input{1,2} into output array
*
* Reads data in a coalesced fashion [BLOCK_THREADS * item + tid] and
* stores the result in output[item].
*/
template <int BLOCK_THREADS, bool IS_FULL_TILE, int ITEMS_PER_THREAD, class T, class It1, class It2>
_CCCL_DEVICE _CCCL_FORCEINLINE void
gmem_to_reg(T (&output)[ITEMS_PER_THREAD], It1 input1, It2 input2, int count1, int count2)
{
if constexpr (IS_FULL_TILE)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ITEMS_PER_THREAD; ++item)
{
const int idx = BLOCK_THREADS * item + threadIdx.x;
// It1 and It2 could have different value types. Convert after load.
output[item] = (idx < count1) ? static_cast<T>(input1[idx]) : static_cast<T>(input2[idx - count1]);
}
}
else
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ITEMS_PER_THREAD; ++item)
{
const int idx = BLOCK_THREADS * item + threadIdx.x;
if (idx < count1 + count2)
{
output[item] = (idx < count1) ? static_cast<T>(input1[idx]) : static_cast<T>(input2[idx - count1]);
}
}
}
}
/// \brief Stores data in a coalesced fashion in[item] -> out[BLOCK_THREADS * item + tid]
template <int BLOCK_THREADS, int ITEMS_PER_THREAD, class T, class It>
_CCCL_DEVICE _CCCL_FORCEINLINE void reg_to_shared(It output, T (&input)[ITEMS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ITEMS_PER_THREAD; ++item)
{
const int idx = BLOCK_THREADS * item + threadIdx.x;
output[idx] = input[item];
}
}
/// \brief The agent is responsible for merging N consecutive sorted arrays into N/2 sorted arrays.
template <typename PolicyGetter, // TODO(bgruber): pass policy as NTTP in C++20
typename KeyIteratorT,
typename ValueIteratorT,
typename OffsetT,
typename CompareOpT,
typename KeyT,
typename ValueT>
struct AgentMerge
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
static constexpr bool KEYS_ONLY = ::cuda::std::is_same_v<ValueT, NullType>;
static constexpr MergeSortPolicy policy = PolicyGetter{}();
static constexpr int BLOCK_THREADS = policy.threads_per_block;
static constexpr int ITEMS_PER_THREAD = policy.items_per_thread;
static constexpr int ITEMS_PER_TILE = BLOCK_THREADS * ITEMS_PER_THREAD;
using KeysLoadPingIt = try_make_cache_modified_iterator_t<policy.load_modifier, KeyIteratorT>;
using ItemsLoadPingIt = try_make_cache_modified_iterator_t<policy.load_modifier, ValueIteratorT>;
using KeysLoadPongIt = try_make_cache_modified_iterator_t<policy.load_modifier, KeyT*>;
using ItemsLoadPongIt = try_make_cache_modified_iterator_t<policy.load_modifier, ValueT*>;
using KeysOutputPongIt = KeyIteratorT;
using ItemsOutputPongIt = ValueIteratorT;
using KeysOutputPingIt = KeyT*;
using ItemsOutputPingIt = ValueT*;
using BlockStoreKeysPong =
BlockStore<it_value_t<KeysOutputPongIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreItemsPong =
BlockStore<it_value_t<ItemsOutputPongIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreKeysPing =
BlockStore<it_value_t<KeysOutputPingIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
using BlockStoreItemsPing =
BlockStore<it_value_t<ItemsOutputPingIt>, BLOCK_THREADS, ITEMS_PER_THREAD, policy.store_algorithm>;
/// Parameterized BlockReduce primitive
union _TempStorage
{
typename BlockStoreKeysPing::TempStorage store_keys_ping;
typename BlockStoreItemsPing::TempStorage store_items_ping;
typename BlockStoreKeysPong::TempStorage store_keys_pong;
typename BlockStoreItemsPong::TempStorage store_items_pong;
KeyT keys_shared[ITEMS_PER_TILE + 1];
ValueT items_shared[ITEMS_PER_TILE + 1];
};
/// Alias wrapper allowing storage to be unioned
using TempStorage = Uninitialized<_TempStorage>;
//---------------------------------------------------------------------
// Per thread data
//---------------------------------------------------------------------
bool ping;
_TempStorage& storage;
KeysLoadPingIt keys_in_ping;
ItemsLoadPingIt items_in_ping;
KeysLoadPongIt keys_in_pong;
ItemsLoadPongIt items_in_pong;
OffsetT keys_count;
KeysOutputPongIt keys_out_pong;
ItemsOutputPongIt items_out_pong;
KeysOutputPingIt keys_out_ping;
ItemsOutputPingIt items_out_ping;
CompareOpT compare_op;
OffsetT* merge_partitions;
OffsetT target_merged_tiles_number;
//---------------------------------------------------------------------
// Utility functions
//---------------------------------------------------------------------
template <bool IS_FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void consume_tile(int tid, OffsetT tile_idx, OffsetT tile_base, int count)
{
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
const OffsetT partition_beg = merge_partitions[tile_idx + 0];
const OffsetT partition_end = merge_partitions[tile_idx + 1];
// target_merged_tiles_number is a power of two.
const OffsetT merged_tiles_number = target_merged_tiles_number / 2;
const OffsetT mask = target_merged_tiles_number - 1;
// The first tile number in the tiles group being merged, equal to:
// target_merged_tiles_number * (tile_idx / target_merged_tiles_number)
const OffsetT list = ~mask & tile_idx;
const OffsetT start = ITEMS_PER_TILE * list;
const OffsetT size = ITEMS_PER_TILE * merged_tiles_number;
const OffsetT diag = ITEMS_PER_TILE * tile_idx - start;
const OffsetT keys1_beg = partition_beg - start;
OffsetT keys1_end = partition_end - start;
const OffsetT keys_end_dist_from_start = keys_count - start;
const OffsetT max_keys2 = (keys_end_dist_from_start > size) ? (keys_end_dist_from_start - size) : 0;
// We have the following invariants:
// diag >= keys1_beg, because diag is the distance of the total merge path so far (keys1 + keys2)
// diag+ITEMS_PER_TILE >= keys1_end, because diag+ITEMS_PER_TILE is the distance of the merge path for the next tile
// and keys1_end is key1's component of that path
const OffsetT keys2_beg = (::cuda::std::min) (max_keys2, diag - keys1_beg);
OffsetT keys2_end =
(::cuda::std::min) (max_keys2,
detail::safe_add_bound_to_max(diag, static_cast<OffsetT>(ITEMS_PER_TILE)) - keys1_end);
// Check if it's the last tile in the tile group being merged
if (mask == (mask & tile_idx))
{
keys1_end = (::cuda::std::min) (keys_count - start, size);
keys2_end = (::cuda::std::min) (max_keys2, size);
}
// number of keys per tile
const int num_keys1 = static_cast<int>(keys1_end - keys1_beg);
const int num_keys2 = static_cast<int>(keys2_end - keys2_beg);
// load keys1 & keys2
KeyT keys_local[ITEMS_PER_THREAD];
if (ping)
{
gmem_to_reg<BLOCK_THREADS, IS_FULL_TILE>(
keys_local, keys_in_ping + start + keys1_beg, keys_in_ping + start + size + keys2_beg, num_keys1, num_keys2);
}
else
{
gmem_to_reg<BLOCK_THREADS, IS_FULL_TILE>(
keys_local, keys_in_pong + start + keys1_beg, keys_in_pong + start + size + keys2_beg, num_keys1, num_keys2);
}
reg_to_shared<BLOCK_THREADS>(&storage.keys_shared[0], keys_local);
// preload items into registers already
//
[[maybe_unused]] ValueT items_local[ITEMS_PER_THREAD];
if constexpr (!KEYS_ONLY)
{
if (ping)
{
gmem_to_reg<BLOCK_THREADS, IS_FULL_TILE>(
items_local,
items_in_ping + start + keys1_beg,
items_in_ping + start + size + keys2_beg,
num_keys1,
num_keys2);
}
else
{
gmem_to_reg<BLOCK_THREADS, IS_FULL_TILE>(
items_local,
items_in_pong + start + keys1_beg,
items_in_pong + start + size + keys2_beg,
num_keys1,
num_keys2);
}
}
__syncthreads();
_CCCL_PDL_TRIGGER_NEXT_LAUNCH();
// use binary search in shared memory
// to find merge path for each of thread
// we can use int type here, because the number of
// items in shared memory is limited
//
const int diag0_local = (::cuda::std::min) (num_keys1 + num_keys2, ITEMS_PER_THREAD * tid);
const int keys1_beg_local = MergePath(
&storage.keys_shared[0], &storage.keys_shared[num_keys1], num_keys1, num_keys2, diag0_local, compare_op);
const int keys1_end_local = num_keys1;
const int keys2_beg_local = diag0_local - keys1_beg_local;
const int keys2_end_local = num_keys2;
const int num_keys1_local = keys1_end_local - keys1_beg_local;
const int num_keys2_local = keys2_end_local - keys2_beg_local;
// perform serial merge
//
int indices[ITEMS_PER_THREAD];
detail::serial_merge<policy.unroll>(
&storage.keys_shared[0],
keys1_beg_local,
keys2_beg_local + num_keys1,
num_keys1_local,
num_keys2_local,
keys_local,
indices,
compare_op);
__syncthreads();
// write keys
if (ping)
{
if constexpr (IS_FULL_TILE)
{
BlockStoreKeysPing(storage.store_keys_ping).Store(keys_out_ping + tile_base, keys_local);
}
else
{
BlockStoreKeysPing(storage.store_keys_ping).Store(keys_out_ping + tile_base, keys_local, num_keys1 + num_keys2);
}
}
else
{
if constexpr (IS_FULL_TILE)
{
BlockStoreKeysPong(storage.store_keys_pong).Store(keys_out_pong + tile_base, keys_local);
}
else
{
BlockStoreKeysPong(storage.store_keys_pong).Store(keys_out_pong + tile_base, keys_local, num_keys1 + num_keys2);
}
}
// if items are provided, merge them
if constexpr (!KEYS_ONLY)
{
__syncthreads();
reg_to_shared<BLOCK_THREADS>(&storage.items_shared[0], items_local);
__syncthreads();
// gather items from shared mem
//
_CCCL_PRAGMA_UNROLL_FULL()
for (int item = 0; item < ITEMS_PER_THREAD; ++item)
{
items_local[item] = storage.items_shared[indices[item]];
}
__syncthreads();
// write from reg to gmem
//
if (ping)
{
if constexpr (IS_FULL_TILE)
{
BlockStoreItemsPing(storage.store_items_ping).Store(items_out_ping + tile_base, items_local);
}
else
{
BlockStoreItemsPing(storage.store_items_ping).Store(items_out_ping + tile_base, items_local, count);
}
}
else
{
if constexpr (IS_FULL_TILE)
{
BlockStoreItemsPong(storage.store_items_pong).Store(items_out_pong + tile_base, items_local);
}
else
{
BlockStoreItemsPong(storage.store_items_pong).Store(items_out_pong + tile_base, items_local, count);
}
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE AgentMerge(
bool ping_,
TempStorage& storage_,
KeysLoadPingIt keys_in_ping_,
ItemsLoadPingIt items_in_ping_,
KeysLoadPongIt keys_in_pong_,
ItemsLoadPongIt items_in_pong_,
OffsetT keys_count_,
KeysOutputPingIt keys_out_ping_,
ItemsOutputPingIt items_out_ping_,
KeysOutputPongIt keys_out_pong_,
ItemsOutputPongIt items_out_pong_,
CompareOpT compare_op_,
OffsetT* merge_partitions_,
OffsetT target_merged_tiles_number_)
: ping(ping_)
, storage(storage_.Alias())
, keys_in_ping(keys_in_ping_)
, items_in_ping(items_in_ping_)
, keys_in_pong(keys_in_pong_)
, items_in_pong(items_in_pong_)
, keys_count(keys_count_)
, keys_out_pong(keys_out_pong_)
, items_out_pong(items_out_pong_)
, keys_out_ping(keys_out_ping_)
, items_out_ping(items_out_ping_)
, compare_op(compare_op_)
, merge_partitions(merge_partitions_)
, target_merged_tiles_number(target_merged_tiles_number_)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
const int tile_idx = static_cast<int>(blockIdx.x);
const int num_tiles = static_cast<int>(gridDim.x);
const OffsetT tile_base = OffsetT(tile_idx) * ITEMS_PER_TILE;
const int tid = static_cast<int>(threadIdx.x);
const int items_in_tile =
static_cast<int>((::cuda::std::min) (static_cast<OffsetT>(ITEMS_PER_TILE), keys_count - tile_base));
if (tile_idx < num_tiles - 1)
{
consume_tile<true>(tid, tile_idx, tile_base, ITEMS_PER_TILE);
}
else
{
consume_tile<false>(tid, tile_idx, tile_base, items_in_tile);
}
}
};
} // namespace detail::merge_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,758 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2018, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* \file
* AgentRadixSortDownsweep implements a stateful abstraction of CUDA thread
* blocks for participating in device-wide radix sort downsweep .
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_exchange.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_radix_rank.cuh>
#include <cub/block/block_store.cuh>
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/thread/thread_load.cuh>
#include <cub/util_device.cuh>
#include <cub/util_type.cuh>
#include <cuda/__warp/warp_shuffle.h>
#include <cuda/std/cstdint>
CUB_NAMESPACE_BEGIN
namespace detail
{
/**
* @brief Parameterizable tuning policy type for AgentRadixSortDownsweep
*
* @tparam NominalThreadsPerBlock4B
* Threads per thread block
*
* @tparam NominalItemsPerThread4B
* Items per thread (per tile of input)
*
* @tparam ComputeT
* Dominant compute type
*
* @tparam LoadAlgorithm
* The BlockLoad algorithm to use
*
* @tparam LoadModifier
* Cache load modifier for reading keys (and values)
*
* @tparam RankAlgorithm
* The radix ranking algorithm to use
*
* @tparam ScanAlgorithm
* The block scan algorithm to use
*
* @tparam RadixBits
* The number of radix bits, i.e., log2(bins)
*/
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
RadixRankAlgorithm RankAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
struct agent_radix_sort_downsweep_policy : ScalingType
{
/// The number of radix bits, i.e., log2(bins)
static constexpr int RADIX_BITS = RadixBits;
/// The BlockLoad algorithm to use
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
/// Cache load modifier for reading keys (and values)
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
/// The radix ranking algorithm to use
static constexpr RadixRankAlgorithm RANK_ALGORITHM = RankAlgorithm;
/// The BlockScan algorithm to use
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
};
} // namespace detail
/******************************************************************************
* Tuning policy types
******************************************************************************/
//! Deprecated [Since 3.5]
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
RadixRankAlgorithm RankAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
using AgentRadixSortDownsweepPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceRadixSort") = detail::agent_radix_sort_downsweep_policy<
NominalThreadsPerBlock4B,
NominalItemsPerThread4B,
ComputeT,
LoadAlgorithm,
LoadModifier,
RankAlgorithm,
ScanAlgorithm,
RadixBits,
ScalingType>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::radix_sort
{
/**
* @brief AgentRadixSortDownsweep implements a stateful abstraction of CUDA thread blocks for participating in
* device-wide radix sort downsweep .
*
* @tparam AgentRadixSortDownsweepPolicy
* Parameterized AgentRadixSortDownsweepPolicy tuning policy type
*
* @tparam IS_DESCENDING
* Whether or not the sorted-order is high-to-low
*
* @tparam KeyT
* KeyT type
*
* @tparam ValueT
* ValueT type
*
* @tparam OffsetT
* Signed integer type for global offsets
*/
template <typename AgentRadixSortDownsweepPolicy,
bool IS_DESCENDING,
typename KeyT,
typename ValueT,
typename OffsetT,
typename DecomposerT = identity_decomposer_t>
struct AgentRadixSortDownsweep
{
//---------------------------------------------------------------------
// Type definitions and constants
//---------------------------------------------------------------------
using traits = radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
using bit_ordered_conversion = typename traits::bit_ordered_conversion_policy;
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = AgentRadixSortDownsweepPolicy::LOAD_ALGORITHM;
static constexpr CacheLoadModifier LOAD_MODIFIER = AgentRadixSortDownsweepPolicy::LOAD_MODIFIER;
static constexpr RadixRankAlgorithm RANK_ALGORITHM = AgentRadixSortDownsweepPolicy::RANK_ALGORITHM;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = AgentRadixSortDownsweepPolicy::SCAN_ALGORITHM;
static constexpr int BLOCK_THREADS = AgentRadixSortDownsweepPolicy::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = AgentRadixSortDownsweepPolicy::ITEMS_PER_THREAD;
static constexpr int RADIX_BITS = AgentRadixSortDownsweepPolicy::RADIX_BITS;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;
static constexpr bool KEYS_ONLY = ::cuda::std::is_same_v<ValueT, NullType>;
static constexpr bool LOAD_WARP_STRIPED =
RANK_ALGORITHM == RADIX_RANK_MATCH || RANK_ALGORITHM == RADIX_RANK_MATCH_EARLY_COUNTS_ANY
|| RANK_ALGORITHM == RADIX_RANK_MATCH_EARLY_COUNTS_ATOMIC_OR;
// Input iterator wrapper type (for applying cache modifier)s
using KeysItr = CacheModifiedInputIterator<LOAD_MODIFIER, bit_ordered_type, OffsetT>;
using ValuesItr = CacheModifiedInputIterator<LOAD_MODIFIER, ValueT, OffsetT>;
// Radix ranking type to use
using BlockRadixRankT = block_radix_rank_t<RANK_ALGORITHM, BLOCK_THREADS, RADIX_BITS, IS_DESCENDING, SCAN_ALGORITHM>;
// Digit extractor type
using fundamental_digit_extractor_t = BFEDigitExtractor<KeyT>;
using digit_extractor_t = typename traits::template digit_extractor_t<fundamental_digit_extractor_t, DecomposerT>;
/// Number of bin-starting offsets tracked per thread
static constexpr int BINS_TRACKED_PER_THREAD = BlockRadixRankT::BINS_TRACKED_PER_THREAD;
// BlockLoad type (keys)
using BlockLoadKeysT = BlockLoad<bit_ordered_type, BLOCK_THREADS, ITEMS_PER_THREAD, LOAD_ALGORITHM>;
// BlockLoad type (values)
using BlockLoadValuesT = BlockLoad<ValueT, BLOCK_THREADS, ITEMS_PER_THREAD, LOAD_ALGORITHM>;
// Value exchange array type
using ValueExchangeT = ValueT[TILE_ITEMS];
/**
* Shared memory storage layout
*/
union __align__(16) _TempStorage
{
typename BlockLoadKeysT::TempStorage load_keys;
typename BlockLoadValuesT::TempStorage load_values;
typename BlockRadixRankT::TempStorage radix_rank;
struct KeysAndOffsets
{
bit_ordered_type exchange_keys[TILE_ITEMS];
OffsetT relative_bin_offsets[RADIX_DIGITS];
} keys_and_offsets;
Uninitialized<ValueExchangeT> exchange_values;
OffsetT exclusive_digit_prefix[RADIX_DIGITS];
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Thread fields
//---------------------------------------------------------------------
// Shared storage for this CTA
_TempStorage& temp_storage;
// Input and output device pointers
KeysItr d_keys_in;
ValuesItr d_values_in;
bit_ordered_type* d_keys_out;
ValueT* d_values_out;
// The global scatter base offset for each digit (valid in the first RADIX_DIGITS threads)
OffsetT bin_offset[BINS_TRACKED_PER_THREAD];
uint32_t current_bit;
uint32_t num_bits;
// Whether to short-circuit
int short_circuit;
DecomposerT decomposer;
//---------------------------------------------------------------------
// Utility methods
//---------------------------------------------------------------------
_CCCL_DEVICE _CCCL_FORCEINLINE digit_extractor_t digit_extractor()
{
return traits::template digit_extractor<fundamental_digit_extractor_t>(current_bit, num_bits, decomposer);
}
/**
* Scatter ranked keys through shared memory, then to device-accessible memory
*/
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterKeys(
bit_ordered_type (&twiddled_keys)[ITEMS_PER_THREAD],
OffsetT (&relative_bin_offsets)[ITEMS_PER_THREAD],
int (&ranks)[ITEMS_PER_THREAD],
OffsetT valid_items)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
temp_storage.keys_and_offsets.exchange_keys[ranks[ITEM]] = twiddled_keys[ITEM];
}
__syncthreads();
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
bit_ordered_type key = temp_storage.keys_and_offsets.exchange_keys[threadIdx.x + (ITEM * BLOCK_THREADS)];
uint32_t digit = digit_extractor().Digit(key);
relative_bin_offsets[ITEM] = temp_storage.keys_and_offsets.relative_bin_offsets[digit];
key = bit_ordered_conversion::from_bit_ordered(decomposer, key);
if (FULL_TILE
|| (static_cast<OffsetT>(threadIdx.x + (ITEM * BLOCK_THREADS)) // NOLINT(bugprone-misplaced-widening-cast)
< valid_items))
{
d_keys_out[relative_bin_offsets[ITEM] + threadIdx.x + (ITEM * BLOCK_THREADS)] = key;
}
}
}
/**
* Scatter ranked values through shared memory, then to device-accessible memory
*/
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterValues(
ValueT (&values)[ITEMS_PER_THREAD],
OffsetT (&relative_bin_offsets)[ITEMS_PER_THREAD],
int (&ranks)[ITEMS_PER_THREAD],
OffsetT valid_items)
{
__syncthreads();
ValueExchangeT& exchange_values = temp_storage.exchange_values.Alias();
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
exchange_values[ranks[ITEM]] = values[ITEM];
}
__syncthreads();
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
ValueT value = exchange_values[threadIdx.x + (ITEM * BLOCK_THREADS)];
if (FULL_TILE
|| (static_cast<OffsetT>(threadIdx.x + (ITEM * BLOCK_THREADS)) // NOLINT(bugprone-misplaced-widening-cast)
< valid_items))
{
d_values_out[relative_bin_offsets[ITEM] + threadIdx.x + (ITEM * BLOCK_THREADS)] = value;
}
}
}
/**
* Load a tile of keys (specialized for full tile, block load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadKeys(
bit_ordered_type (&keys)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
bit_ordered_type oob_item,
::cuda::std::true_type is_full_tile,
::cuda::std::false_type warp_striped)
{
BlockLoadKeysT(temp_storage.load_keys).Load(d_keys_in + block_offset, keys);
__syncthreads();
}
/**
* Load a tile of keys (specialized for partial tile, block load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadKeys(
bit_ordered_type (&keys)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
bit_ordered_type oob_item,
::cuda::std::false_type is_full_tile,
::cuda::std::false_type warp_striped)
{
// Register pressure work-around: moving valid_items through shfl prevents compiler
// from reusing guards/addressing from prior guarded loads
valid_items = ::cuda::device::warp_shuffle_idx(valid_items, 0);
BlockLoadKeysT(temp_storage.load_keys).Load(d_keys_in + block_offset, keys, valid_items, oob_item);
__syncthreads();
}
/**
* Load a tile of keys (specialized for full tile, warp-striped load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadKeys(
bit_ordered_type (&keys)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
bit_ordered_type oob_item,
::cuda::std::true_type is_full_tile,
::cuda::std::true_type warp_striped)
{
LoadDirectWarpStriped(threadIdx.x, d_keys_in + block_offset, keys);
}
/**
* Load a tile of keys (specialized for partial tile, warp-striped load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadKeys(
bit_ordered_type (&keys)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
bit_ordered_type oob_item,
::cuda::std::false_type is_full_tile,
::cuda::std::true_type warp_striped)
{
// Register pressure work-around: moving valid_items through shfl prevents compiler
// from reusing guards/addressing from prior guarded loads
valid_items = ::cuda::device::warp_shuffle_idx(valid_items, 0);
LoadDirectWarpStriped(threadIdx.x, d_keys_in + block_offset, keys, valid_items, oob_item);
}
/**
* Load a tile of values (specialized for full tile, block load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadValues(
ValueT (&values)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
::cuda::std::true_type is_full_tile,
::cuda::std::false_type warp_striped)
{
BlockLoadValuesT(temp_storage.load_values).Load(d_values_in + block_offset, values);
__syncthreads();
}
/**
* Load a tile of values (specialized for partial tile, block load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadValues(
ValueT (&values)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
::cuda::std::false_type is_full_tile,
::cuda::std::false_type warp_striped)
{
// Register pressure work-around: moving valid_items through shfl prevents compiler
// from reusing guards/addressing from prior guarded loads
valid_items = ::cuda::device::warp_shuffle_idx(valid_items, 0);
BlockLoadValuesT(temp_storage.load_values).Load(d_values_in + block_offset, values, valid_items);
__syncthreads();
}
/**
* Load a tile of items (specialized for full tile, warp-striped load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadValues(
ValueT (&values)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
::cuda::std::true_type is_full_tile,
::cuda::std::true_type warp_striped)
{
LoadDirectWarpStriped(threadIdx.x, d_values_in + block_offset, values);
}
/**
* Load a tile of items (specialized for partial tile, warp-striped load)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadValues(
ValueT (&values)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
::cuda::std::false_type is_full_tile,
::cuda::std::true_type warp_striped)
{
// Register pressure work-around: moving valid_items through shfl prevents compiler
// from reusing guards/addressing from prior guarded loads
valid_items = ::cuda::device::warp_shuffle_idx(valid_items, 0);
LoadDirectWarpStriped(threadIdx.x, d_values_in + block_offset, values, valid_items);
}
/**
* Truck along associated values
*/
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void GatherScatterValues(
OffsetT (&relative_bin_offsets)[ITEMS_PER_THREAD],
int (&ranks)[ITEMS_PER_THREAD],
OffsetT block_offset,
OffsetT valid_items,
::cuda::std::false_type /*is_keys_only*/)
{
ValueT values[ITEMS_PER_THREAD];
__syncthreads();
LoadValues(values, block_offset, valid_items, bool_constant_v<FULL_TILE>, bool_constant_v<LOAD_WARP_STRIPED>);
ScatterValues<FULL_TILE>(values, relative_bin_offsets, ranks, valid_items);
}
/**
* Truck along associated values (specialized for key-only sorting)
*/
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void GatherScatterValues(
OffsetT (& /*relative_bin_offsets*/)[ITEMS_PER_THREAD],
int (& /*ranks*/)[ITEMS_PER_THREAD],
OffsetT /*block_offset*/,
OffsetT /*valid_items*/,
::cuda::std::true_type /*is_keys_only*/)
{}
/**
* Process tile
*/
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessTile(OffsetT block_offset, OffsetT valid_items = TILE_ITEMS)
{
bit_ordered_type keys[ITEMS_PER_THREAD];
int ranks[ITEMS_PER_THREAD];
OffsetT relative_bin_offsets[ITEMS_PER_THREAD];
// Assign default (min/max) value to all keys
bit_ordered_type default_key =
IS_DESCENDING ? traits::min_raw_binary_key(decomposer) : traits::max_raw_binary_key(decomposer);
// Load tile of keys
LoadKeys(
keys, block_offset, valid_items, default_key, bool_constant_v<FULL_TILE>, bool_constant_v<LOAD_WARP_STRIPED>);
_CCCL_PRAGMA_UNROLL_FULL()
for (int KEY = 0; KEY < ITEMS_PER_THREAD; KEY++)
{
keys[KEY] = bit_ordered_conversion::to_bit_ordered(decomposer, keys[KEY]);
}
// Rank the twiddled keys
int exclusive_digit_prefix[BINS_TRACKED_PER_THREAD];
BlockRadixRankT(temp_storage.radix_rank).RankKeys(keys, ranks, digit_extractor(), exclusive_digit_prefix);
__syncthreads();
// Share exclusive digit prefix
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
// Store exclusive prefix
temp_storage.exclusive_digit_prefix[bin_idx] = exclusive_digit_prefix[track];
}
}
__syncthreads();
// Get inclusive digit prefix
int inclusive_digit_prefix[BINS_TRACKED_PER_THREAD];
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
if (IS_DESCENDING)
{
// Get inclusive digit prefix from exclusive prefix (higher bins come first)
inclusive_digit_prefix[track] =
(bin_idx == 0) ? (BLOCK_THREADS * ITEMS_PER_THREAD) : temp_storage.exclusive_digit_prefix[bin_idx - 1];
}
else
{
// Get inclusive digit prefix from exclusive prefix (lower bins come first)
inclusive_digit_prefix[track] =
(bin_idx == RADIX_DIGITS - 1)
? (BLOCK_THREADS * ITEMS_PER_THREAD)
: temp_storage.exclusive_digit_prefix[bin_idx + 1];
}
}
}
__syncthreads();
// Update global scatter base offsets for each digit
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
bin_offset[track] -= exclusive_digit_prefix[track];
temp_storage.keys_and_offsets.relative_bin_offsets[bin_idx] = bin_offset[track];
bin_offset[track] += inclusive_digit_prefix[track];
}
}
__syncthreads();
// Scatter keys
ScatterKeys<FULL_TILE>(keys, relative_bin_offsets, ranks, valid_items);
// Gather/scatter values
GatherScatterValues<FULL_TILE>(relative_bin_offsets, ranks, block_offset, valid_items, bool_constant_v<KEYS_ONLY>);
}
//---------------------------------------------------------------------
// Copy shortcut
//---------------------------------------------------------------------
/**
* Copy tiles within the range of input
*/
template <typename InputIteratorT, typename T>
_CCCL_DEVICE _CCCL_FORCEINLINE void Copy(InputIteratorT d_in, T* d_out, OffsetT block_offset, OffsetT block_end)
{
// Simply copy the input
while (block_end - block_offset >= TILE_ITEMS)
{
T items[ITEMS_PER_THREAD];
LoadDirectStriped<BLOCK_THREADS>(threadIdx.x, d_in + block_offset, items);
__syncthreads();
StoreDirectStriped<BLOCK_THREADS>(threadIdx.x, d_out + block_offset, items);
block_offset += TILE_ITEMS;
}
// Clean up last partial tile with guarded-I/O
if (block_offset < block_end)
{
OffsetT valid_items = block_end - block_offset;
T items[ITEMS_PER_THREAD];
LoadDirectStriped<BLOCK_THREADS>(threadIdx.x, d_in + block_offset, items, valid_items);
__syncthreads();
StoreDirectStriped<BLOCK_THREADS>(threadIdx.x, d_out + block_offset, items, valid_items);
}
}
/**
* Copy tiles within the range of input (specialized for NullType)
*/
template <typename InputIteratorT>
_CCCL_DEVICE _CCCL_FORCEINLINE void
Copy(InputIteratorT /*d_in*/, NullType* /*d_out*/, OffsetT /*block_offset*/, OffsetT /*block_end*/)
{}
//---------------------------------------------------------------------
// Interface
//---------------------------------------------------------------------
/**
* Constructor
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AgentRadixSortDownsweep(
TempStorage& temp_storage,
OffsetT (&bin_offset)[BINS_TRACKED_PER_THREAD],
OffsetT num_items,
const KeyT* d_keys_in,
KeyT* d_keys_out,
const ValueT* d_values_in,
ValueT* d_values_out,
int current_bit,
int num_bits,
DecomposerT decomposer = {})
: temp_storage(temp_storage.Alias())
, d_keys_in(reinterpret_cast<const bit_ordered_type*>(d_keys_in))
, d_values_in(d_values_in)
, d_keys_out(reinterpret_cast<bit_ordered_type*>(d_keys_out))
, d_values_out(d_values_out)
, current_bit(current_bit)
, num_bits(num_bits)
, short_circuit(1)
, decomposer(decomposer)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
this->bin_offset[track] = bin_offset[track];
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
// Short circuit if the histogram has only bin counts of only zeros or problem-size
short_circuit = short_circuit && ((bin_offset[track] == 0) || (bin_offset[track] == num_items));
}
}
short_circuit = __syncthreads_and(short_circuit);
}
/**
* Constructor
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AgentRadixSortDownsweep(
TempStorage& temp_storage,
OffsetT num_items,
OffsetT* d_spine,
const KeyT* d_keys_in,
KeyT* d_keys_out,
const ValueT* d_values_in,
ValueT* d_values_out,
int current_bit,
int num_bits,
DecomposerT decomposer = {})
: temp_storage(temp_storage.Alias())
, d_keys_in(reinterpret_cast<const bit_ordered_type*>(d_keys_in))
, d_values_in(d_values_in)
, d_keys_out(reinterpret_cast<bit_ordered_type*>(d_keys_out))
, d_values_out(d_values_out)
, current_bit(current_bit)
, num_bits(num_bits)
, short_circuit(1)
, decomposer(decomposer)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
// Load digit bin offsets (each of the first RADIX_DIGITS threads will load an offset for that digit)
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
if (IS_DESCENDING)
{
bin_idx = RADIX_DIGITS - bin_idx - 1;
}
// Short circuit if the first block's histogram has only bin counts of only zeros or problem-size
OffsetT first_block_bin_offset = d_spine[gridDim.x * bin_idx];
short_circuit = short_circuit && ((first_block_bin_offset == 0) || (first_block_bin_offset == num_items));
// Load my block's bin offset for my bin
bin_offset[track] = d_spine[(gridDim.x * bin_idx) + blockIdx.x];
}
}
short_circuit = __syncthreads_and(short_circuit);
}
/**
* Distribute keys from a segment of input tiles.
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessRegion(OffsetT block_offset, OffsetT block_end)
{
if (short_circuit)
{
// Copy keys
Copy(d_keys_in, d_keys_out, block_offset, block_end);
// Copy values
Copy(d_values_in, d_values_out, block_offset, block_end);
}
else
{
// Process full tiles of tile_items
_CCCL_PRAGMA_NOUNROLL()
while (block_end - block_offset >= TILE_ITEMS)
{
ProcessTile<true>(block_offset);
block_offset += TILE_ITEMS;
__syncthreads();
}
// Clean up last partial tile with guarded-I/O
if (block_offset < block_end)
{
ProcessTile<false>(block_offset, block_end - block_offset);
}
}
}
};
} // namespace detail::radix_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,288 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2020, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* \file
* agent_radix_sort_histogram.cuh implements a stateful abstraction of CUDA
* thread blocks for participating in the device histogram kernel used for
* one-sweep radix sorting.
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/thread/thread_reduce.cuh>
#include <cub/util_math.cuh>
#include <cub/util_type.cuh>
#include <cuda/__cmath/ceil_div.h>
#include <cuda/__ptx/instructions/get_sreg.h>
#include <cuda/std/__algorithm/max.h>
#include <cuda/std/__algorithm/min.h>
#include <cuda/std/__functional/operations.h>
#include <cuda/std/__type_traits/is_void.h>
CUB_NAMESPACE_BEGIN
namespace detail
{
//! @param ComputeT If void, use NOMINAL_4B_NUM_PARTS directly for NUM_PARTS. Otherwise, perform scaling.
template <int ThreadsPerBlock, int ItemsPerThread, int NOMINAL_4B_NUM_PARTS, typename ComputeT, int RadixBits>
struct agent_radix_sort_histogram_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
// need to discard sizeof(ComputeType) in case it's void
template <typename ComputeType = ComputeT>
_CCCL_HOST_DEVICE_API static constexpr int num_parts_helper()
{
if constexpr (::cuda::std::is_void_v<ComputeT>)
{
return NOMINAL_4B_NUM_PARTS;
}
else
{
return ::cuda::std::max(1, NOMINAL_4B_NUM_PARTS * 4 / ::cuda::std::max(int{sizeof(ComputeType)}, 4));
}
}
/** NUM_PARTS is the number of private histograms (parts) each histogram is split
* into. Each warp lane is assigned to a specific part based on the lane
* ID. However, lanes with the same ID in different warp use the same private
* histogram. This arrangement helps reduce the degree of conflicts in atomic
* operations. */
static constexpr int NUM_PARTS = num_parts_helper<ComputeT>();
static constexpr int RADIX_BITS = RadixBits;
};
template <int ThreadsPerBlock, int RadixBits>
struct agent_radix_sort_exclusive_sum_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int RADIX_BITS = RadixBits;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock, int ItemsPerThread, int NOMINAL_4B_NUM_PARTS, typename ComputeT, int RadixBits>
using AgentRadixSortHistogramPolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceRadixSort") =
detail::agent_radix_sort_histogram_policy<ThreadsPerBlock, ItemsPerThread, NOMINAL_4B_NUM_PARTS, ComputeT, RadixBits>;
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock, int RadixBits>
using AgentRadixSortExclusiveSumPolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceRadixSort") =
detail::agent_radix_sort_exclusive_sum_policy<ThreadsPerBlock, RadixBits>;
namespace detail::radix_sort
{
template <typename AgentRadixSortHistogramPolicy,
bool IS_DESCENDING,
typename KeyT,
typename OffsetT,
typename DecomposerT = identity_decomposer_t>
struct AgentRadixSortHistogram
{
// constants
static constexpr int ITEMS_PER_THREAD = AgentRadixSortHistogramPolicy::ITEMS_PER_THREAD;
static constexpr int BLOCK_THREADS = AgentRadixSortHistogramPolicy::BLOCK_THREADS;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
static constexpr int RADIX_BITS = AgentRadixSortHistogramPolicy::RADIX_BITS;
static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;
static constexpr int MAX_NUM_PASSES = (sizeof(KeyT) * 8 + RADIX_BITS - 1) / RADIX_BITS;
static constexpr int NUM_PARTS = AgentRadixSortHistogramPolicy::NUM_PARTS;
using traits = radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
using bit_ordered_conversion = typename traits::bit_ordered_conversion_policy;
using Twiddle = RadixSortTwiddle<IS_DESCENDING, KeyT>;
using ShmemCounterT = uint32_t;
using ShmemAtomicCounterT = ShmemCounterT;
using fundamental_digit_extractor_t = ShiftDigitExtractor<KeyT>;
using digit_extractor_t = typename traits::template digit_extractor_t<fundamental_digit_extractor_t, DecomposerT>;
struct _TempStorage
{
ShmemAtomicCounterT bins[MAX_NUM_PASSES][RADIX_DIGITS][NUM_PARTS];
};
struct TempStorage : Uninitialized<_TempStorage>
{};
// thread fields
// shared memory storage
_TempStorage& s;
// bins for the histogram
OffsetT* d_bins_out;
// data to compute the histogram
const bit_ordered_type* d_keys_in;
// number of data items
OffsetT num_items;
// begin and end bits for sorting
int begin_bit, end_bit;
// number of sorting passes
int num_passes; // NOLINT(modernize-use-default-member-init)
DecomposerT decomposer;
_CCCL_DEVICE _CCCL_FORCEINLINE AgentRadixSortHistogram(
TempStorage& temp_storage,
OffsetT* d_bins_out,
const KeyT* d_keys_in,
OffsetT num_items,
int begin_bit,
int end_bit,
DecomposerT decomposer = {})
: s(temp_storage.Alias())
, d_bins_out(d_bins_out)
, d_keys_in(reinterpret_cast<const bit_ordered_type*>(d_keys_in))
, num_items(num_items)
, begin_bit(begin_bit)
, end_bit(end_bit)
, num_passes((end_bit - begin_bit + RADIX_BITS - 1) / RADIX_BITS)
, decomposer(decomposer)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void Init()
{
// Initialize bins to 0.
_CCCL_PRAGMA_UNROLL_FULL()
for (int bin = static_cast<int>(threadIdx.x); bin < RADIX_DIGITS; bin += BLOCK_THREADS)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int pass = 0; pass < num_passes; ++pass)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int part = 0; part < NUM_PARTS; ++part)
{
s.bins[pass][bin][part] = 0;
}
}
}
__syncthreads();
}
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadTileKeys(OffsetT tile_offset, bit_ordered_type (&keys)[ITEMS_PER_THREAD])
{
// tile_offset < num_items always, hence the line below works
bool full_tile = num_items - tile_offset >= TILE_ITEMS;
if (full_tile)
{
LoadDirectStriped<BLOCK_THREADS>(threadIdx.x, d_keys_in + tile_offset, keys);
}
else
{
LoadDirectStriped<BLOCK_THREADS>(
threadIdx.x, d_keys_in + tile_offset, keys, num_items - tile_offset, Twiddle::DefaultKey(decomposer));
}
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
keys[u] = Twiddle::In(keys[u], decomposer);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
AccumulateSharedHistograms(OffsetT tile_offset, bit_ordered_type (&keys)[ITEMS_PER_THREAD])
{
int part = ::cuda::ptx::get_sreg_laneid() % NUM_PARTS;
_CCCL_PRAGMA_UNROLL_FULL()
for (int current_bit = begin_bit, pass = 0; current_bit < end_bit; current_bit += RADIX_BITS, ++pass)
{
const int num_bits = ::cuda::std::min(+RADIX_BITS, end_bit - current_bit);
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
uint32_t bin = digit_extractor(current_bit, num_bits).Digit(keys[u]);
// Using cuda::atomic<> results in lower performance on GP100,
// so atomicAdd() is used instead.
atomicAdd(&s.bins[pass][bin][part], 1);
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void AccumulateGlobalHistograms()
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int bin = static_cast<int>(threadIdx.x); bin < RADIX_DIGITS; bin += BLOCK_THREADS)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int pass = 0; pass < num_passes; ++pass)
{
OffsetT count = cub::ThreadReduce(s.bins[pass][bin], ::cuda::std::plus<>{});
if (count > 0)
{
// Using cuda::atomic<> here would also require using it in
// other kernels. However, other kernels of onesweep sorting
// (ExclusiveSum, Onesweep) don't need atomic
// access. Therefore, atomicAdd() is used, until
// cuda::atomic_ref<> becomes available.
atomicAdd(&d_bins_out[pass * RADIX_DIGITS + bin], count);
}
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
_CCCL_PDL_TRIGGER_NEXT_LAUNCH();
// Within a portion, avoid overflowing (u)int32 counters.
// Between portions, accumulate results in global memory.
constexpr OffsetT MAX_PORTION_SIZE = 1 << 30;
OffsetT num_portions = ::cuda::ceil_div(num_items, MAX_PORTION_SIZE);
for (OffsetT portion = 0; portion < num_portions; ++portion)
{
// Reset the counters.
Init();
// Process the tiles.
OffsetT portion_offset = portion * MAX_PORTION_SIZE;
OffsetT portion_size = ::cuda::std::min(MAX_PORTION_SIZE, num_items - portion_offset);
for (OffsetT offset = static_cast<OffsetT>(blockIdx.x) * TILE_ITEMS; offset < portion_size;
offset += OffsetT{TILE_ITEMS} * gridDim.x)
{
OffsetT tile_offset = portion_offset + offset;
bit_ordered_type keys[ITEMS_PER_THREAD];
LoadTileKeys(tile_offset, keys);
AccumulateSharedHistograms(tile_offset, keys);
}
__syncthreads();
// Accumulate the result in global memory.
// Wait for global histogram init
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
AccumulateGlobalHistograms();
__syncthreads();
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE digit_extractor_t digit_extractor(int current_bit, int num_bits)
{
return traits::template digit_extractor<fundamental_digit_extractor_t>(current_bit, num_bits, decomposer);
}
};
} // namespace detail::radix_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,742 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2020, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* \file
* agent_radix_sort_onesweep.cuh implements a stateful abstraction of CUDA
* thread blocks for participating in the device one-sweep radix sort kernel.
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_radix_rank.cuh>
#include <cub/block/block_store.cuh>
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/util_ptx.cuh>
#include <cub/util_type.cuh>
#include <cuda/__ptx/instructions/get_sreg.h>
#include <cuda/std/__concepts/same_as.h>
#include <cuda/std/__fwd/format.h>
#include <cuda/std/__host_stdlib/ostream>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/integral_constant.h>
#include <cuda/std/__type_traits/is_same.h>
CUB_NAMESPACE_BEGIN
/** \brief cub::RadixSortStoreAlgorithm enumerates different algorithms to write
* partitioned elements (keys, values) stored in shared memory into global
* memory. Currently applies only to writing 4B keys in full tiles; in all other cases,
* RADIX_SORT_STORE_DIRECT is used.
*/
enum RadixSortStoreAlgorithm
{
/** \brief Elements are statically distributed among block threads, which write them
* into the appropriate partition in global memory. This results in fewer instructions
* and more writes in flight at a given moment, but may generate more transactions. */
RADIX_SORT_STORE_DIRECT,
/** \brief Elements are distributed among warps in a block distribution. Each warp
* goes through its elements and tries to write them while minimizing the number of
* memory transactions. This results in fewer memory transactions, but more
* instructions and less writes in flight at a given moment. */
RADIX_SORT_STORE_ALIGNED
};
#if _CCCL_HOSTED()
namespace detail
{
[[nodiscard]] _CCCL_HOST_DEVICE_API constexpr const char* to_string(RadixSortStoreAlgorithm algo) noexcept
{
switch (algo)
{
case RADIX_SORT_STORE_DIRECT:
return "RADIX_SORT_STORE_DIRECT";
case RADIX_SORT_STORE_ALIGNED:
return "RADIX_SORT_STORE_ALIGNED";
}
return "<unknown RadixSortStoreAlgorithm>";
}
} // namespace detail
#endif // _CCCL_HOSTED()
#if _CCCL_HOSTED() && !defined(_CCCL_DOXYGEN_INVOKED)
inline ::std::ostream& operator<<(::std::ostream& os, RadixSortStoreAlgorithm algo)
{
return os << CUB_NS_QUALIFIER::detail::to_string(algo);
}
#endif // _CCCL_HOSTED() && !_CCCL_DOXYGEN_INVOKED
CUB_NAMESPACE_END
#if __cpp_lib_format >= 201907L && !defined(_CCCL_DOXYGEN_INVOKED)
template <::cuda::std::same_as<char> CharT>
struct std::formatter<CUB_NS_QUALIFIER::RadixSortStoreAlgorithm, CharT> : formatter<const CharT*, CharT>
{
template <class FmtCtx>
auto format(const CUB_NS_QUALIFIER::RadixSortStoreAlgorithm& algo, FmtCtx& ctx) const
{
return formatter<const CharT*, CharT>::format(CUB_NS_QUALIFIER::detail::to_string(algo), ctx);
}
};
#endif // __cpp_lib_format >= 201907L && !defined(_CCCL_DOXYGEN_INVOKED)
CUB_NAMESPACE_BEGIN
namespace detail
{
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
/** Number of private histograms to use in the ranker;
ignored if the ranking algorithm is not one of RADIX_RANK_MATCH_EARLY_COUNTS_* */
int RankNumParts,
/** Ranking algorithm used in the onesweep kernel. Only algorithms that
support warp-strided key arrangement and count callbacks are supported. */
RadixRankAlgorithm RankAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
RadixSortStoreAlgorithm StoreAlgorithm,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
struct agent_radix_sort_onesweep_policy : ScalingType
{
static constexpr int RANK_NUM_PARTS = RankNumParts;
static constexpr int RADIX_BITS = RadixBits;
static constexpr RadixRankAlgorithm RANK_ALGORITHM = RankAlgorithm;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
static constexpr RadixSortStoreAlgorithm STORE_ALGORITHM = StoreAlgorithm;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
int RankNumParts,
RadixRankAlgorithm RankAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
RadixSortStoreAlgorithm StoreAlgorithm,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
using AgentRadixSortOnesweepPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceRadixSort") = detail::agent_radix_sort_onesweep_policy<
NominalThreadsPerBlock4B,
NominalItemsPerThread4B,
ComputeT,
RankNumParts,
RankAlgorithm,
ScanAlgorithm,
StoreAlgorithm,
RadixBits,
ScalingType>;
namespace detail::radix_sort
{
template <typename AgentRadixSortOnesweepPolicy,
bool IS_DESCENDING,
typename KeyT,
typename ValueT,
typename OffsetT,
typename PortionOffsetT,
typename DecomposerT = identity_decomposer_t>
struct AgentRadixSortOnesweep
{
// constants
static constexpr int ITEMS_PER_THREAD = AgentRadixSortOnesweepPolicy::ITEMS_PER_THREAD;
static constexpr bool KEYS_ONLY = ::cuda::std::is_same_v<ValueT, NullType>;
static constexpr int BLOCK_THREADS = AgentRadixSortOnesweepPolicy::BLOCK_THREADS;
static constexpr int RANK_NUM_PARTS = AgentRadixSortOnesweepPolicy::RANK_NUM_PARTS;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
static constexpr int RADIX_BITS = AgentRadixSortOnesweepPolicy::RADIX_BITS;
static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;
static constexpr int BINS_PER_THREAD = (RADIX_DIGITS + BLOCK_THREADS - 1) / BLOCK_THREADS;
static constexpr bool FULL_BINS = BINS_PER_THREAD * BLOCK_THREADS == RADIX_DIGITS;
static constexpr int WARP_THREADS = warp_threads;
static constexpr int BLOCK_WARPS = BLOCK_THREADS / WARP_THREADS;
static constexpr int WARP_MASK = ~0;
static constexpr int LOOKBACK_PARTIAL_MASK = 1 << (PortionOffsetT(sizeof(PortionOffsetT)) * 8 - 2);
static constexpr int LOOKBACK_GLOBAL_MASK = 1 << (PortionOffsetT(sizeof(PortionOffsetT)) * 8 - 1);
static constexpr int LOOKBACK_KIND_MASK = LOOKBACK_PARTIAL_MASK | LOOKBACK_GLOBAL_MASK;
static constexpr int LOOKBACK_VALUE_MASK = ~LOOKBACK_KIND_MASK;
using traits = radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
using bit_ordered_conversion = typename traits::bit_ordered_conversion_policy;
using fundamental_digit_extractor_t = ShiftDigitExtractor<KeyT>;
using digit_extractor_t = typename traits::template digit_extractor_t<fundamental_digit_extractor_t, DecomposerT>;
using AtomicOffsetT = PortionOffsetT;
static constexpr RadixRankAlgorithm RANK_ALGORITHM = AgentRadixSortOnesweepPolicy::RANK_ALGORITHM;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = AgentRadixSortOnesweepPolicy::SCAN_ALGORITHM;
static constexpr RadixSortStoreAlgorithm STORE_ALGORITHM =
sizeof(bit_ordered_type) == sizeof(uint32_t)
? AgentRadixSortOnesweepPolicy::STORE_ALGORITHM
: RADIX_SORT_STORE_DIRECT;
using Twiddle = RadixSortTwiddle<IS_DESCENDING, KeyT>;
static_assert(RANK_ALGORITHM == RADIX_RANK_MATCH || RANK_ALGORITHM == RADIX_RANK_MATCH_EARLY_COUNTS_ANY
|| RANK_ALGORITHM == RADIX_RANK_MATCH_EARLY_COUNTS_ATOMIC_OR,
"for onesweep agent, the ranking algorithm must warp-strided key arrangement");
using BlockRadixRankT = ::cuda::std::_If<
RANK_ALGORITHM == RADIX_RANK_MATCH_EARLY_COUNTS_ATOMIC_OR,
BlockRadixRankMatchEarlyCounts<BLOCK_THREADS, RADIX_BITS, false, SCAN_ALGORITHM, WARP_MATCH_ATOMIC_OR, RANK_NUM_PARTS>,
::cuda::std::_If<
RANK_ALGORITHM == RADIX_RANK_MATCH,
BlockRadixRankMatch<BLOCK_THREADS, RADIX_BITS, false, SCAN_ALGORITHM>,
BlockRadixRankMatchEarlyCounts<BLOCK_THREADS, RADIX_BITS, false, SCAN_ALGORITHM, WARP_MATCH_ANY, RANK_NUM_PARTS>>>;
// temporary storage
struct TempStorage_
{
union
{
bit_ordered_type keys_out[TILE_ITEMS];
ValueT values_out[TILE_ITEMS];
typename BlockRadixRankT::TempStorage rank_temp_storage;
};
union
{
OffsetT global_offsets[RADIX_DIGITS];
PortionOffsetT block_idx;
};
};
using TempStorage = Uninitialized<TempStorage_>;
// thread variables
TempStorage_& s;
// kernel parameters
AtomicOffsetT* d_lookback;
AtomicOffsetT* d_ctrs;
OffsetT* d_bins_out;
const OffsetT* d_bins_in;
bit_ordered_type* d_keys_out;
const bit_ordered_type* d_keys_in;
ValueT* d_values_out;
const ValueT* d_values_in;
PortionOffsetT num_items;
int current_bit;
int num_bits;
// other thread variables
int warp;
int lane;
DecomposerT decomposer;
PortionOffsetT block_idx;
bool full_block;
_CCCL_DEVICE _CCCL_FORCEINLINE digit_extractor_t digit_extractor()
{
return traits::template digit_extractor<fundamental_digit_extractor_t>(current_bit, num_bits, decomposer);
}
// helper methods
_CCCL_DEVICE _CCCL_FORCEINLINE uint32_t Digit(bit_ordered_type key)
{
return digit_extractor().Digit(key);
}
_CCCL_DEVICE _CCCL_FORCEINLINE int ThreadBin(int u)
{
return threadIdx.x * BINS_PER_THREAD + u;
}
_CCCL_DEVICE _CCCL_FORCEINLINE void LookbackPartial(int (&bins)[BINS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
int bin = ThreadBin(u);
if (FULL_BINS || bin < RADIX_DIGITS)
{
// write the local sum into the bin
AtomicOffsetT& loc = d_lookback[block_idx * RADIX_DIGITS + bin];
PortionOffsetT value = bins[u] | LOOKBACK_PARTIAL_MASK;
ThreadStore<STORE_VOLATILE>(&loc, value);
}
}
}
struct CountsCallback
{
using AgentT =
AgentRadixSortOnesweep<AgentRadixSortOnesweepPolicy, IS_DESCENDING, KeyT, ValueT, OffsetT, PortionOffsetT, DecomposerT>;
AgentT& agent;
int (&bins)[BINS_PER_THREAD];
bit_ordered_type (&keys)[ITEMS_PER_THREAD];
static constexpr bool EMPTY = false;
_CCCL_DEVICE _CCCL_FORCEINLINE
CountsCallback(AgentT& agent, int (&bins)[BINS_PER_THREAD], bit_ordered_type (&keys)[ITEMS_PER_THREAD])
: agent(agent)
, bins(bins)
, keys(keys)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void operator()(int (&other_bins)[BINS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
bins[u] = other_bins[u];
}
// Wait for lookback init
_CCCL_PDL_GRID_DEPENDENCY_SYNC();
agent.LookbackPartial(bins);
agent.TryShortCircuit(keys, bins);
}
};
_CCCL_DEVICE _CCCL_FORCEINLINE void LookbackGlobal(int (&bins)[BINS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
int bin = ThreadBin(u);
if (FULL_BINS || bin < RADIX_DIGITS)
{
PortionOffsetT inc_sum = bins[u];
int want_mask = ~0;
// backtrack as long as necessary
for (PortionOffsetT block_jdx = block_idx - 1; block_jdx >= 0; --block_jdx)
{
// wait for some value to appear
PortionOffsetT value_j = 0;
AtomicOffsetT& loc_j = d_lookback[block_jdx * RADIX_DIGITS + bin];
do
{
__threadfence_block(); // prevent hoisting loads from loop
value_j = ThreadLoad<LOAD_VOLATILE>(&loc_j);
} while (value_j == 0);
inc_sum += value_j & LOOKBACK_VALUE_MASK;
want_mask = __ballot_sync(want_mask, (value_j & LOOKBACK_GLOBAL_MASK) == 0);
if (value_j & LOOKBACK_GLOBAL_MASK)
{
break;
}
}
AtomicOffsetT& loc_i = d_lookback[block_idx * RADIX_DIGITS + bin];
PortionOffsetT value_i = inc_sum | LOOKBACK_GLOBAL_MASK;
ThreadStore<STORE_VOLATILE>(&loc_i, value_i);
s.global_offsets[bin] += inc_sum - bins[u];
}
}
_CCCL_PDL_TRIGGER_NEXT_LAUNCH();
}
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadKeys(OffsetT tile_offset, bit_ordered_type (&keys)[ITEMS_PER_THREAD])
{
if (full_block)
{
LoadDirectWarpStriped(threadIdx.x, d_keys_in + tile_offset, keys);
}
else
{
LoadDirectWarpStriped(
threadIdx.x, d_keys_in + tile_offset, keys, num_items - tile_offset, Twiddle::DefaultKey(decomposer));
}
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
keys[u] = Twiddle::In(keys[u], decomposer);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadValues(OffsetT tile_offset, ValueT (&values)[ITEMS_PER_THREAD])
{
if (full_block)
{
LoadDirectWarpStriped(threadIdx.x, d_values_in + tile_offset, values);
}
else
{
int tile_items = num_items - tile_offset;
LoadDirectWarpStriped(threadIdx.x, d_values_in + tile_offset, values, tile_items);
}
}
/** Checks whether "short-circuiting" is possible. Short-circuiting happens
* if all TILE_ITEMS keys fall into the same bin, i.e. have the same digit
* value (note that it only happens for full tiles). If short-circuiting is
* performed, the part of the ranking algorithm after the CountsCallback, as
* well as the rest of the sorting (e.g. scattering keys and values to
* shared and global memory) are skipped; updates related to decoupled
* look-back are still performed. Instead, the keys assigned to the current
* thread block are written cooperatively into a contiguous location in
* d_keys_out corresponding to their digit. The values (if also sorting
* values) assigned to the current thread block are similarly copied from
* d_values_in to d_values_out. */
_CCCL_DEVICE _CCCL_FORCEINLINE void
TryShortCircuit(bit_ordered_type (&keys)[ITEMS_PER_THREAD], int (&bins)[BINS_PER_THREAD])
{
// check if any bin can be short-circuited
bool short_circuit = false;
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
if (FULL_BINS || ThreadBin(u) < RADIX_DIGITS)
{
short_circuit = short_circuit || bins[u] == TILE_ITEMS;
}
}
short_circuit = __syncthreads_or(short_circuit);
if (!short_circuit)
{
return;
}
ShortCircuitCopy(keys, bins);
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
ShortCircuitCopy(bit_ordered_type (&keys)[ITEMS_PER_THREAD], int (&bins)[BINS_PER_THREAD])
{
// short-circuit handling; note that global look-back is still required
// compute offsets
uint32_t common_bin = Digit(keys[0]);
int offsets[BINS_PER_THREAD];
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
int bin = ThreadBin(u);
offsets[u] = bin > common_bin ? TILE_ITEMS : 0;
}
// global lookback
LoadBinsToOffsetsGlobal(offsets);
LookbackGlobal(bins);
UpdateBinsGlobal(bins, offsets);
__syncthreads();
// scatter the keys
OffsetT global_offset = s.global_offsets[common_bin];
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
keys[u] = Twiddle::Out(keys[u], decomposer);
}
if (full_block)
{
StoreDirectWarpStriped(threadIdx.x, d_keys_out + global_offset, keys);
}
else
{
int tile_items = num_items - block_idx * TILE_ITEMS;
StoreDirectWarpStriped(threadIdx.x, d_keys_out + global_offset, keys, tile_items);
}
if (!KEYS_ONLY)
{
// gather and scatter the values
ValueT values[ITEMS_PER_THREAD];
LoadValues(block_idx * TILE_ITEMS, values); // NOLINT(bugprone-misplaced-widening-cast)
if (full_block)
{
StoreDirectWarpStriped(threadIdx.x, d_values_out + global_offset, values);
}
else
{
int tile_items = num_items - block_idx * TILE_ITEMS;
StoreDirectWarpStriped(threadIdx.x, d_values_out + global_offset, values, tile_items);
}
}
// exit early
ThreadExit();
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScatterKeysShared(bit_ordered_type (&keys)[ITEMS_PER_THREAD], int (&ranks)[ITEMS_PER_THREAD])
{
// write to shared memory
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
s.keys_out[ranks[u]] = keys[u];
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScatterValuesShared(ValueT (&values)[ITEMS_PER_THREAD], int (&ranks)[ITEMS_PER_THREAD])
{
// write to shared memory
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
s.values_out[ranks[u]] = values[u];
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void LoadBinsToOffsetsGlobal(int (&offsets)[BINS_PER_THREAD])
{
// global offset - global part
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
int bin = ThreadBin(u);
if (FULL_BINS || bin < RADIX_DIGITS)
{
s.global_offsets[bin] = d_bins_in[bin] - offsets[u];
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void UpdateBinsGlobal(int (&bins)[BINS_PER_THREAD], int (&offsets)[BINS_PER_THREAD])
{
bool last_block = (block_idx + 1) * TILE_ITEMS >= num_items;
if (d_bins_out != nullptr && last_block)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < BINS_PER_THREAD; ++u)
{
int bin = ThreadBin(u);
if (FULL_BINS || bin < RADIX_DIGITS)
{
d_bins_out[bin] = s.global_offsets[bin] + offsets[u] + bins[u];
}
}
}
}
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterKeysGlobalDirect()
{
int tile_items = FULL_TILE ? TILE_ITEMS : num_items - block_idx * TILE_ITEMS;
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
int idx = threadIdx.x + u * BLOCK_THREADS;
bit_ordered_type key = s.keys_out[idx];
OffsetT global_idx = idx + s.global_offsets[Digit(key)];
if (FULL_TILE || idx < tile_items)
{
d_keys_out[global_idx] = Twiddle::Out(key, decomposer);
}
__syncwarp(WARP_MASK);
}
}
template <bool FULL_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterValuesGlobalDirect(int (&digits)[ITEMS_PER_THREAD])
{
int tile_items = FULL_TILE ? TILE_ITEMS : num_items - block_idx * TILE_ITEMS;
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
int idx = threadIdx.x + u * BLOCK_THREADS;
ValueT value = s.values_out[idx];
OffsetT global_idx = idx + s.global_offsets[digits[u]];
if (FULL_TILE || idx < tile_items)
{
d_values_out[global_idx] = value;
}
__syncwarp(WARP_MASK);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterKeysGlobalAligned()
{
// this only works with full tiles
constexpr int ITEMS_PER_WARP = TILE_ITEMS / BLOCK_WARPS;
constexpr int ALIGN = 8;
constexpr auto CACHE_MODIFIER = STORE_CG;
int warp_start = warp * ITEMS_PER_WARP;
int warp_end = (warp + 1) * ITEMS_PER_WARP;
int warp_offset = warp_start;
while (warp_offset < warp_end - WARP_THREADS)
{
int idx = warp_offset + lane;
bit_ordered_type key = s.keys_out[idx];
bit_ordered_type key_out = Twiddle::Out(key, decomposer);
OffsetT global_idx = idx + s.global_offsets[Digit(key)];
int last_lane = WARP_THREADS - 1;
int num_writes = WARP_THREADS;
if (lane == last_lane)
{
num_writes -= int(global_idx + 1) % ALIGN;
}
num_writes = __shfl_sync(WARP_MASK, num_writes, last_lane);
if (lane < num_writes)
{
ThreadStore<CACHE_MODIFIER>(&d_keys_out[global_idx], key_out);
}
warp_offset += num_writes;
}
{
int num_writes = warp_end - warp_offset;
if (lane < num_writes)
{
int idx = warp_offset + lane;
bit_ordered_type key = s.keys_out[idx];
OffsetT global_idx = idx + s.global_offsets[Digit(key)];
ThreadStore<CACHE_MODIFIER>(&d_keys_out[global_idx], Twiddle::Out(key, decomposer));
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterKeysGlobal()
{
// write block data to global memory
if (full_block)
{
if constexpr (STORE_ALGORITHM == RADIX_SORT_STORE_ALIGNED)
{
ScatterKeysGlobalAligned();
}
else
{
ScatterKeysGlobalDirect<true>();
}
}
else
{
ScatterKeysGlobalDirect<false>();
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterValuesGlobal(int (&digits)[ITEMS_PER_THREAD])
{
// write block data to global memory
if (full_block)
{
ScatterValuesGlobalDirect<true>(digits);
}
else
{
ScatterValuesGlobalDirect<false>(digits);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void ComputeKeyDigits(int (&digits)[ITEMS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int u = 0; u < ITEMS_PER_THREAD; ++u)
{
int idx = threadIdx.x + u * BLOCK_THREADS;
digits[u] = Digit(s.keys_out[idx]);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
GatherScatterValues(int (&ranks)[ITEMS_PER_THREAD], ::cuda::std::false_type keys_only)
{
// compute digits corresponding to the keys
int digits[ITEMS_PER_THREAD];
ComputeKeyDigits(digits);
// load values
ValueT values[ITEMS_PER_THREAD];
LoadValues(block_idx * TILE_ITEMS, values); // NOLINT(bugprone-misplaced-widening-cast)
// scatter values
__syncthreads();
ScatterValuesShared(values, ranks);
__syncthreads();
ScatterValuesGlobal(digits);
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
GatherScatterValues(int (&ranks)[ITEMS_PER_THREAD], ::cuda::std::true_type keys_only)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void Process()
{
// load keys
// if warp1 < warp2, all elements of warp1 occur before those of warp2
// in the source array
bit_ordered_type keys[ITEMS_PER_THREAD];
LoadKeys(block_idx * TILE_ITEMS, keys); // NOLINT(bugprone-misplaced-widening-cast)
// rank keys
int ranks[ITEMS_PER_THREAD];
int exclusive_digit_prefix[BINS_PER_THREAD];
int bins[BINS_PER_THREAD];
BlockRadixRankT(s.rank_temp_storage)
.RankKeys(keys, ranks, digit_extractor(), exclusive_digit_prefix, CountsCallback(*this, bins, keys));
// scatter keys in shared memory
__syncthreads();
ScatterKeysShared(keys, ranks);
// compute global offsets
LoadBinsToOffsetsGlobal(exclusive_digit_prefix);
LookbackGlobal(bins);
UpdateBinsGlobal(bins, exclusive_digit_prefix);
// scatter keys in global memory
__syncthreads();
ScatterKeysGlobal();
// scatter values if necessary
GatherScatterValues(ranks, bool_constant_v<KEYS_ONLY>);
}
_CCCL_DEVICE _CCCL_FORCEINLINE //
AgentRadixSortOnesweep(
TempStorage& temp_storage,
AtomicOffsetT* d_lookback,
AtomicOffsetT* d_ctrs,
OffsetT* d_bins_out,
const OffsetT* d_bins_in,
KeyT* d_keys_out,
const KeyT* d_keys_in,
ValueT* d_values_out,
const ValueT* d_values_in,
PortionOffsetT num_items,
int current_bit,
int num_bits,
DecomposerT decomposer = {})
: s(temp_storage.Alias())
, d_lookback(d_lookback)
, d_ctrs(d_ctrs)
, d_bins_out(d_bins_out)
, d_bins_in(d_bins_in)
, d_keys_out(reinterpret_cast<bit_ordered_type*>(d_keys_out))
, d_keys_in(reinterpret_cast<const bit_ordered_type*>(d_keys_in))
, d_values_out(d_values_out)
, d_values_in(d_values_in)
, num_items(num_items)
, current_bit(current_bit)
, num_bits(num_bits)
, warp(static_cast<int>(threadIdx.x / WARP_THREADS))
, lane(static_cast<int>(::cuda::ptx::get_sreg_laneid()))
, decomposer(decomposer)
{
// initialization
if (threadIdx.x == 0)
{
s.block_idx = atomicAdd(d_ctrs, 1);
}
__syncthreads();
block_idx = s.block_idx;
full_block = (block_idx + 1) * TILE_ITEMS <= num_items;
}
};
} // namespace detail::radix_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,517 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2018, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* \file
* AgentRadixSortUpsweep implements a stateful abstraction of CUDA thread blocks for participating in device-wide radix
* sort upsweep .
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/thread/thread_load.cuh>
#include <cub/thread/thread_reduce.cuh>
#include <cub/util_device.cuh>
#include <cub/util_type.cuh>
#include <cub/warp/warp_reduce.cuh>
#include <cuda/__ptx/instructions/get_sreg.h>
#include <cuda/__utility/static_for.h>
#include <cuda/std/__algorithm/max.h>
#include <cuda/std/__algorithm/min.h>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
/**
* @brief Parameterizable tuning policy type for AgentRadixSortUpsweep
*
* @tparam NominalThreadsPerBlock4B
* Threads per thread block
*
* @tparam NominalItemsPerThread4B
* Items per thread (per tile of input)
*
* @tparam ComputeT
* Dominant compute type
*
* @tparam LoadModifier
* Cache load modifier for reading keys
*
* @tparam RadixBits
* The number of radix bits, i.e., log2(bins)
*/
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
CacheLoadModifier LoadModifier,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
struct agent_radix_sort_upsweep_policy : ScalingType
{
/// The number of radix bits, i.e., log2(bins)
static constexpr int RADIX_BITS = RadixBits;
/// Cache load modifier for reading keys
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
CacheLoadModifier LoadModifier,
int RadixBits,
typename ScalingType = detail::RegBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
using AgentRadixSortUpsweepPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceRadixSort") = detail::agent_radix_sort_upsweep_policy<
NominalThreadsPerBlock4B,
NominalItemsPerThread4B,
ComputeT,
LoadModifier,
RadixBits,
ScalingType>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::radix_sort
{
/**
* @brief AgentRadixSortUpsweep implements a stateful abstraction of CUDA thread blocks for
* participating in device-wide radix sort upsweep .
*
* @tparam AgentRadixSortUpsweepPolicy
* Parameterized AgentRadixSortUpsweepPolicy tuning policy type
*
* @tparam KeyT
* KeyT type
*
* @tparam DecomposerT = identity_decomposer_t
* Signed integer type for global offsets
*/
template <typename AgentRadixSortUpsweepPolicy,
typename KeyT,
typename OffsetT,
typename DecomposerT = identity_decomposer_t>
struct AgentRadixSortUpsweep
{
//---------------------------------------------------------------------
// Type definitions and constants
//---------------------------------------------------------------------
using traits = radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
using bit_ordered_conversion = typename traits::bit_ordered_conversion_policy;
// Integer type for digit counters (to be packed into words of PackedCounters)
using DigitCounter = unsigned char;
// Integer type for packing DigitCounters into columns of shared memory banks
using PackedCounter = unsigned int;
static constexpr CacheLoadModifier LOAD_MODIFIER = AgentRadixSortUpsweepPolicy::LOAD_MODIFIER;
static constexpr int RADIX_BITS = AgentRadixSortUpsweepPolicy::RADIX_BITS;
static constexpr int BLOCK_THREADS = AgentRadixSortUpsweepPolicy::BLOCK_THREADS;
static constexpr int KEYS_PER_THREAD = AgentRadixSortUpsweepPolicy::ITEMS_PER_THREAD;
static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;
static constexpr int LOG_WARP_THREADS = log2_warp_threads;
static constexpr int WARP_THREADS = 1 << LOG_WARP_THREADS;
static constexpr int WARPS = (BLOCK_THREADS + WARP_THREADS - 1) / WARP_THREADS;
static constexpr int TILE_ITEMS = BLOCK_THREADS * KEYS_PER_THREAD;
static constexpr int BYTES_PER_COUNTER = sizeof(DigitCounter);
static constexpr int OG_BYTES_PER_COUNTER = Log2<BYTES_PER_COUNTER>::VALUE;
static constexpr int PACKING_RATIO = sizeof(PackedCounter) / sizeof(DigitCounter);
static constexpr int LOG_PACKING_RATIO = Log2<PACKING_RATIO>::VALUE;
static constexpr int LOG_COUNTER_LANES = ::cuda::std::max(0, int(RADIX_BITS) - int(LOG_PACKING_RATIO));
static constexpr int COUNTER_LANES = 1 << LOG_COUNTER_LANES;
// To prevent counter overflow, we must periodically unpack and aggregate the
// digit counters back into registers. Each counter lane is assigned to a
// warp for aggregation.
static constexpr int LANES_PER_WARP = ::cuda::std::max(1, (COUNTER_LANES + WARPS - 1) / WARPS);
// Unroll tiles in batches without risk of counter overflow
static constexpr int UNROLL_COUNT = ::cuda::std::min(64, 255 / KEYS_PER_THREAD);
static constexpr int UNROLLED_ELEMENTS = UNROLL_COUNT * TILE_ITEMS;
// Input iterator wrapper type (for applying cache modifier)s
using KeysItr = CacheModifiedInputIterator<LOAD_MODIFIER, bit_ordered_type, OffsetT>;
// Digit extractor type
using fundamental_digit_extractor_t = BFEDigitExtractor<KeyT>;
using digit_extractor_t = typename traits::template digit_extractor_t<fundamental_digit_extractor_t, DecomposerT>;
/**
* Shared memory storage layout
*/
union __align__(16) _TempStorage
{
DigitCounter thread_counters[COUNTER_LANES][BLOCK_THREADS][PACKING_RATIO];
PackedCounter packed_thread_counters[COUNTER_LANES][BLOCK_THREADS];
OffsetT block_counters[WARP_THREADS][RADIX_DIGITS];
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Thread fields (aggregate state bundle)
//---------------------------------------------------------------------
// Shared storage for this CTA
_TempStorage& temp_storage;
// Thread-local counters for periodically aggregating composite-counter lanes
OffsetT local_counts[LANES_PER_WARP][PACKING_RATIO];
// Input and output device pointers
KeysItr d_keys_in;
// Target bits
int current_bit;
int num_bits;
DecomposerT decomposer;
//---------------------------------------------------------------------
// Utility methods
//---------------------------------------------------------------------
_CCCL_DEVICE _CCCL_FORCEINLINE digit_extractor_t digit_extractor()
{
return traits::template digit_extractor<fundamental_digit_extractor_t>(current_bit, num_bits, decomposer);
}
/**
* Decode a key and increment corresponding smem digit counter
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void Bucket(bit_ordered_type key)
{
// Perform transform op
bit_ordered_type converted_key = bit_ordered_conversion::to_bit_ordered(decomposer, key);
// Extract current digit bits
uint32_t digit = digit_extractor().Digit(converted_key);
// Get sub-counter offset
uint32_t sub_counter = digit & (PACKING_RATIO - 1);
// Get row offset
uint32_t row_offset = digit >> LOG_PACKING_RATIO;
_CCCL_ASSERT(row_offset < COUNTER_LANES, "");
// Increment counter
temp_storage.thread_counters[row_offset][threadIdx.x][sub_counter]++;
}
/**
* Reset composite counters
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ResetDigitCounters()
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int LANE = 0; LANE < COUNTER_LANES; LANE++)
{
temp_storage.packed_thread_counters[LANE][threadIdx.x] = 0;
}
}
/**
* Reset the unpacked counters in each thread
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ResetUnpackedCounters()
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int LANE = 0; LANE < LANES_PER_WARP; LANE++)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int UNPACKED_COUNTER = 0; UNPACKED_COUNTER < PACKING_RATIO; UNPACKED_COUNTER++)
{
local_counts[LANE][UNPACKED_COUNTER] = 0;
}
}
}
/**
* Extracts and aggregates the digit counters for each counter lane
* owned by this warp
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void UnpackDigitCounts()
{
unsigned int warp_id = threadIdx.x >> LOG_WARP_THREADS;
unsigned int warp_tid = ::cuda::ptx::get_sreg_laneid();
_CCCL_PRAGMA_UNROLL_FULL()
for (int LANE = 0; LANE < LANES_PER_WARP; LANE++)
{
const int counter_lane = (LANE * WARPS) + warp_id;
if (counter_lane < COUNTER_LANES)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int PACKED_COUNTER = 0; PACKED_COUNTER < BLOCK_THREADS; PACKED_COUNTER += WARP_THREADS)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int UNPACKED_COUNTER = 0; UNPACKED_COUNTER < PACKING_RATIO; UNPACKED_COUNTER++)
{
OffsetT counter = temp_storage.thread_counters[counter_lane][warp_tid + PACKED_COUNTER][UNPACKED_COUNTER];
local_counts[LANE][UNPACKED_COUNTER] += counter;
}
}
}
}
}
/**
* Processes a single, full tile
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessFullTile(OffsetT block_offset)
{
// Tile of keys
bit_ordered_type keys[KEYS_PER_THREAD];
LoadDirectStriped<BLOCK_THREADS>(threadIdx.x, d_keys_in + block_offset, keys);
// Prevent hoisting
__syncthreads();
// Bucket tile of keys
cuda::static_for<KEYS_PER_THREAD>([&](auto ic) {
Bucket(keys[ic]);
});
}
/**
* Processes a single load (may have some threads masked off)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessPartialTile(OffsetT block_offset, const OffsetT& block_end)
{
// Process partial tile if necessary using single loads
for (OffsetT offset = threadIdx.x; offset < block_end - block_offset; offset += BLOCK_THREADS)
{
// Load and bucket key
bit_ordered_type key = d_keys_in[block_offset + offset];
Bucket(key);
}
}
//---------------------------------------------------------------------
// Interface
//---------------------------------------------------------------------
/**
* Constructor
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AgentRadixSortUpsweep(
TempStorage& temp_storage, const KeyT* d_keys_in, int current_bit, int num_bits, DecomposerT decomposer = {})
: temp_storage(temp_storage.Alias())
, d_keys_in(reinterpret_cast<const bit_ordered_type*>(d_keys_in))
, current_bit(current_bit)
, num_bits(num_bits)
, decomposer(decomposer)
{}
/**
* Compute radix digit histograms from a segment of input tiles.
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessRegion(OffsetT block_offset, const OffsetT& block_end)
{
// Reset digit counters in smem and unpacked counters in registers
ResetDigitCounters();
ResetUnpackedCounters();
// Unroll batches of full tiles
while (block_end - block_offset >= UNROLLED_ELEMENTS)
{
for (int i = 0; i < UNROLL_COUNT; ++i)
{
ProcessFullTile(block_offset);
block_offset += TILE_ITEMS;
}
__syncthreads();
// Aggregate back into local_count registers to prevent overflow
UnpackDigitCounts();
__syncthreads();
// Reset composite counters in lanes
ResetDigitCounters();
}
// Unroll single full tiles
while (block_end - block_offset >= TILE_ITEMS)
{
ProcessFullTile(block_offset);
block_offset += TILE_ITEMS;
}
// Process partial tile if necessary
ProcessPartialTile(block_offset, block_end);
__syncthreads();
// Aggregate back into local_count registers
UnpackDigitCounts();
}
/**
* Extract counts (saving them to the external array)
*/
template <bool IS_DESCENDING>
_CCCL_DEVICE _CCCL_FORCEINLINE void ExtractCounts(OffsetT* counters, int bin_stride = 1, int bin_offset = 0)
{
unsigned int warp_id = threadIdx.x >> LOG_WARP_THREADS;
unsigned int warp_tid = ::cuda::ptx::get_sreg_laneid();
// Place unpacked digit counters in shared memory
_CCCL_PRAGMA_UNROLL_FULL()
for (int LANE = 0; LANE < LANES_PER_WARP; LANE++)
{
int counter_lane = (LANE * WARPS) + warp_id;
if (counter_lane < COUNTER_LANES)
{
int digit_row = counter_lane << LOG_PACKING_RATIO;
_CCCL_PRAGMA_UNROLL_FULL()
for (int UNPACKED_COUNTER = 0; UNPACKED_COUNTER < PACKING_RATIO; UNPACKED_COUNTER++)
{
int bin_idx = digit_row + UNPACKED_COUNTER;
temp_storage.block_counters[warp_tid][bin_idx] = local_counts[LANE][UNPACKED_COUNTER];
}
}
}
__syncthreads();
// Rake-reduce bin_count reductions
// Whole blocks
_CCCL_PRAGMA_UNROLL_FULL()
for (int BIN_BASE = RADIX_DIGITS % BLOCK_THREADS; (BIN_BASE + BLOCK_THREADS) <= RADIX_DIGITS;
BIN_BASE += BLOCK_THREADS)
{
int bin_idx = static_cast<int>(BIN_BASE + threadIdx.x);
OffsetT bin_count = 0;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < WARP_THREADS; ++i)
{
bin_count += temp_storage.block_counters[i][bin_idx];
}
if (IS_DESCENDING)
{
bin_idx = RADIX_DIGITS - bin_idx - 1;
}
counters[(bin_stride * bin_idx) + bin_offset] = bin_count;
}
// Remainder
if ((RADIX_DIGITS % BLOCK_THREADS != 0) && (threadIdx.x < RADIX_DIGITS))
{
int bin_idx = static_cast<int>(threadIdx.x);
OffsetT bin_count = 0;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < WARP_THREADS; ++i)
{
bin_count += temp_storage.block_counters[i][bin_idx];
}
if (IS_DESCENDING)
{
bin_idx = RADIX_DIGITS - bin_idx - 1;
}
counters[(bin_stride * bin_idx) + bin_offset] = bin_count;
}
}
/**
* @brief Extract counts
*
* @param[out] bin_count
* The exclusive prefix sum for the digits
* [(threadIdx.x * BINS_TRACKED_PER_THREAD) ... (threadIdx.x * BINS_TRACKED_PER_THREAD) + BINS_TRACKED_PER_THREAD -
* 1]
*/
template <int BINS_TRACKED_PER_THREAD>
_CCCL_DEVICE _CCCL_FORCEINLINE void ExtractCounts(OffsetT (&bin_count)[BINS_TRACKED_PER_THREAD])
{
unsigned int warp_id = threadIdx.x >> LOG_WARP_THREADS;
unsigned int warp_tid = ::cuda::ptx::get_sreg_laneid();
// Place unpacked digit counters in shared memory
_CCCL_PRAGMA_UNROLL_FULL()
for (int LANE = 0; LANE < LANES_PER_WARP; LANE++)
{
int counter_lane = (LANE * WARPS) + warp_id;
if (counter_lane < COUNTER_LANES)
{
int digit_row = counter_lane << LOG_PACKING_RATIO;
_CCCL_PRAGMA_UNROLL_FULL()
for (int UNPACKED_COUNTER = 0; UNPACKED_COUNTER < PACKING_RATIO; UNPACKED_COUNTER++)
{
int bin_idx = digit_row + UNPACKED_COUNTER;
temp_storage.block_counters[warp_tid][bin_idx] = local_counts[LANE][UNPACKED_COUNTER];
}
}
}
__syncthreads();
// Rake-reduce bin_count reductions
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
bin_count[track] = 0;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < WARP_THREADS; ++i)
{
bin_count[track] += temp_storage.block_counters[i][bin_idx];
}
}
}
}
};
} // namespace detail::radix_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,611 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2022, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
//! @file
//! cub::AgentReduce implements a stateful abstraction of CUDA thread blocks for participating in device-wide reduction.
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/block_reduce.cuh>
#include <cub/detail/type_traits.cuh>
#include <cub/grid/grid_even_share.cuh>
#include <cub/grid/grid_mapping.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_device.cuh>
#include <cub/util_type.cuh>
#include <thrust/type_traits/is_trivially_relocatable.h>
#include <cuda/std/__algorithm/min.h>
#include <cuda/std/__functional/identity.h>
#include <cuda/std/__functional/operations.h>
#include <cuda/std/__memory/is_sufficiently_aligned.h>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/is_pointer.h>
#include <cuda/std/__type_traits/is_same.h>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
// TODO(bgruber): drop in CCCL 4.0
/**
* Parameterizable tuning policy type for AgentReduce
* @tparam NominalThreadsPerBlock4B Threads per thread block
* @tparam NominalItemsPerThread4B Items per thread (per tile of input)
* @tparam ComputeT Dominant compute type
* @tparam VectorLoadLength Number of items per vectorized load
* @tparam BlockAlgorithm Cooperative block-wide reduction algorithm to use
* @tparam LoadModifier Cache load modifier for reading input elements
*/
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
int VectorLoadLength,
BlockReduceAlgorithm BlockAlgorithm,
CacheLoadModifier LoadModifier,
typename ScalingType = MemBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
struct agent_reduce_policy : ScalingType
{
/// Number of items per vectorized load
static constexpr int VECTOR_LOAD_LENGTH = VectorLoadLength;
/// Cooperative block-wide reduction algorithm to use
static constexpr BlockReduceAlgorithm BLOCK_ALGORITHM = BlockAlgorithm;
/// Cache load modifier for reading input elements
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
int VectorLoadLength,
BlockReduceAlgorithm BlockAlgorithm,
CacheLoadModifier LoadModifier,
typename ScalingType = detail::MemBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>>
using AgentReducePolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceReduce") = detail::agent_reduce_policy<
NominalThreadsPerBlock4B,
NominalItemsPerThread4B,
ComputeT,
VectorLoadLength,
BlockAlgorithm,
LoadModifier,
ScalingType>;
namespace detail
{
template <int ThreadsPerBlock,
int WarpThreads,
int NominalItemsPerThread4B,
typename ComputeT,
int VectorLoadLength,
CacheLoadModifier LoadModifier>
struct agent_warp_reduce_policy
{
/// Number of threads per warp
static constexpr int WARP_THREADS = WarpThreads;
/// Number of items per vectorized load
static constexpr int VECTOR_LOAD_LENGTH = VectorLoadLength;
/// Number of threads per block
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
/// Number of items per thread. When `ComputeT` is `void`, the nominal value is used as-is (no scaling),
/// allowing to pass actual items_per_thread to opt out of the legacy 4B scaling.
static constexpr int ITEMS_PER_THREAD =
::cuda::std::conditional_t<::cuda::std::is_same_v<ComputeT, void>,
NoScaling<0, NominalItemsPerThread4B>,
MemBoundScaling<0, NominalItemsPerThread4B, ComputeT>>::ITEMS_PER_THREAD;
/// Cache load modifier for reading input elements
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
/// Number of items per tile
constexpr static int ITEMS_PER_TILE = ITEMS_PER_THREAD * WARP_THREADS;
/// Number of segments per block
constexpr static int SEGMENTS_PER_BLOCK = BLOCK_THREADS / WARP_THREADS;
static_assert((BLOCK_THREADS % WARP_THREADS) == 0, "Block should be multiple of warp");
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int WarpThreads,
int NominalItemsPerThread4B,
typename ComputeT,
int VectorLoadLength,
CacheLoadModifier LoadModifier>
using AgentWarpReducePolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceSegmentedReduce") = detail::
agent_warp_reduce_policy<ThreadsPerBlock, WarpThreads, NominalItemsPerThread4B, ComputeT, VectorLoadLength, LoadModifier>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::reduce
{
/**
* @brief AgentReduceImpl implements a stateful abstraction of CUDA thread blocks
* and warps, for participating in device-wide reduction .
*
* Each thread reduces only the values it loads. If `FIRST_TILE`, this partial
* reduction is stored into `thread_aggregate`. Otherwise it is accumulated
* into `thread_aggregate`.
*
* @tparam AgentReducePolicy
* Parameterized AgentReducePolicy tuning policy type
*
* @tparam InputIteratorT
* Random-access iterator type for input
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam ReductionOp
* Binary reduction operator type having member
* `auto operator()(T &&a, U &&b)`
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*
* @tparam TransformOp
* Unary operator type having member `auto operator()(T &&a)`
*
* @tparam CollectiveReduceT
* Block or Warp reduction type
*
* @tparam NumThreads
* Number of threads participating in the collective reduction
*
* @tparam IsWarpReduction
* Whether or not this is a warp reduction
*/
template <typename AgentReducePolicy,
typename InputIteratorT,
typename OffsetT,
typename ReductionOp,
typename AccumT,
typename TransformOp,
typename CollectiveReduceT,
int NumThreads,
bool IsWarpReduction = false>
struct AgentReduceImpl
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
/// The input value type
using InputT = it_value_t<InputIteratorT>;
/// Vector type of InputT for data movement
using VectorT = typename CubVector<InputT, AgentReducePolicy::VECTOR_LOAD_LENGTH>::Type;
/// Input iterator wrapper type (for applying cache modifier)
// Wrap the native input pointer with CacheModifiedInputIterator
// or directly use the supplied input iterator type
using WrappedInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<InputIteratorT>,
CacheModifiedInputIterator<AgentReducePolicy::LOAD_MODIFIER, InputT, OffsetT>,
InputIteratorT>;
/// Constants
static constexpr int ITEMS_PER_THREAD = AgentReducePolicy::ITEMS_PER_THREAD;
static constexpr int TILE_ITEMS = NumThreads * ITEMS_PER_THREAD;
static constexpr int vec_size = ::cuda::std::min(ITEMS_PER_THREAD, AgentReducePolicy::VECTOR_LOAD_LENGTH);
// Can vectorize according to the policy if the input iterator is a native
// pointer to a primitive type
// TODO(bgruber): we should not check for `is_pointer_v` but `contiguous_iterator` and unwrap it
static constexpr bool ATTEMPT_VECTORIZATION =
(vec_size > 1) && (ITEMS_PER_THREAD % vec_size == 0)
&& (::cuda::std::is_pointer_v<InputIteratorT>)
// TODO(bgruber): remove the check for is_primitive<ValueT> in CCCL 4.0
&&(is_primitive<InputT>::value || THRUST_NS_QUALIFIER::is_trivially_relocatable_v<InputT>)
// vectorizing large types leads to regressions again, see https://github.com/NVIDIA/cccl/issues/9761
// TODO(bgruber): this should be decided by tuning
&&sizeof(InputT)
<= 8;
static constexpr CacheLoadModifier LOAD_MODIFIER = AgentReducePolicy::LOAD_MODIFIER;
/// Shared memory type required by this thread block
struct _TempStorage
{
typename CollectiveReduceT::TempStorage reduce;
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
_TempStorage& temp_storage; ///< Reference to temp_storage
InputIteratorT d_in; ///< Input data to reduce
WrappedInputIteratorT d_wrapped_in; ///< Wrapped input data to reduce
ReductionOp reduction_op; ///< Binary reduction operator
TransformOp transform_op; ///< Transform operator
unsigned int lane_id; ///< Local thread index inside a Warp or Block
//---------------------------------------------------------------------
// Utility
//---------------------------------------------------------------------
// Whether the input is aligned with the vector type
template <typename Iterator, bool AttemptVectorization = ATTEMPT_VECTORIZATION>
[[nodiscard]] _CCCL_DEVICE_API static bool IsAligned(Iterator d_in) noexcept
{
if constexpr (AttemptVectorization)
{
return ::cuda::std::is_sufficiently_aligned<alignof(VectorT)>(d_in);
}
else
{
return false;
}
}
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
/**
* @brief Constructor
* @param temp_storage Reference to temp_storage
* @param d_in Input data to reduce
* @param reduction_op Binary reduction operator
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AgentReduceImpl(
TempStorage& temp_storage, InputIteratorT d_in, ReductionOp reduction_op, TransformOp transform_op, int lane_id)
: temp_storage(temp_storage.Alias())
, d_in(d_in)
, d_wrapped_in(d_in)
, reduction_op(reduction_op)
, transform_op(transform_op)
, lane_id(lane_id)
{}
//---------------------------------------------------------------------
// Tile consumption
//---------------------------------------------------------------------
/**
* @brief Consume a full tile of input
* @tparam IsFirstTile Whether this is a full tile
* @param block_offset The offset the tile to consume
* @param input_is_vector_aligned Whether we can vectorize loads
*/
template <int IsFirstTile, bool CanVectorize>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeFullTile(AccumT& thread_aggregate, OffsetT block_offset)
{
if constexpr (CanVectorize)
{
// Fabricate a vectorized input iterator
InputT* d_in_unqualified = const_cast<InputT*>(d_in) + block_offset + (lane_id * vec_size);
CacheModifiedInputIterator<AgentReducePolicy::LOAD_MODIFIER, VectorT, OffsetT> d_vec_in(
reinterpret_cast<VectorT*>(d_in_unqualified));
// Load items as vector items
InputT input_items[ITEMS_PER_THREAD];
VectorT* vec_items = reinterpret_cast<VectorT*>(input_items);
// Alias items as an array of VectorT and load it in striped fashion
static constexpr int words = ITEMS_PER_THREAD / vec_size;
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < words; ++i)
{
vec_items[i] = d_vec_in[NumThreads * i];
}
// Convert from input type to output type
AccumT items[ITEMS_PER_THREAD];
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = 0; i < ITEMS_PER_THREAD; ++i)
{
items[i] = transform_op(input_items[i]);
}
// Reduce items within each thread stripe
thread_aggregate =
IsFirstTile ? cub::ThreadReduce(items, reduction_op) : cub::ThreadReduce(items, reduction_op, thread_aggregate);
}
else
{
// Scalar path: load items in striped fashion and reduce items within each thread stripe
AccumT items[ITEMS_PER_THREAD];
load_transform_direct_striped<NumThreads>(lane_id, d_wrapped_in + block_offset, items, transform_op);
thread_aggregate =
IsFirstTile ? cub::ThreadReduce(items, reduction_op) : cub::ThreadReduce(items, reduction_op, thread_aggregate);
}
}
/**
* Consume a partial tile of input
* @tparam IsFirstTile Whether or not this is a full tile
* @param block_offset The offset the tile to consume
* @param valid_items The number of valid items in the tile
*/
template <int IsFirstTile>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumePartialTile(AccumT& thread_aggregate, OffsetT block_offset, int valid_items)
{
// Partial tile
int thread_offset = lane_id;
// Read first item
if (IsFirstTile && (thread_offset < valid_items))
{
thread_aggregate =
transform_op(d_wrapped_in[block_offset + thread_offset]); // NOLINT(bugprone-misplaced-widening-cast)
thread_offset += NumThreads;
}
// Continue reading items (block-striped)
while (thread_offset < valid_items)
{
InputT item(d_wrapped_in[block_offset + thread_offset]); // NOLINT(bugprone-misplaced-widening-cast)
thread_aggregate = reduction_op(thread_aggregate, transform_op(item));
thread_offset += NumThreads;
}
}
//---------------------------------------------------------------
// Consume a contiguous segment of tiles
//---------------------------------------------------------------------
/**
* @brief Reduce a contiguous segment of input tiles
* @param even_share GridEvenShare descriptor
*/
template <bool CanVectorize>
_CCCL_DEVICE _CCCL_FORCEINLINE AccumT ConsumeRange(GridEvenShare<OffsetT>& even_share)
{
AccumT thread_aggregate{};
if (even_share.block_end - even_share.block_offset < TILE_ITEMS)
{
// First tile isn't full (not all threads have valid items)
int valid_items = even_share.block_end - even_share.block_offset;
ConsumePartialTile<true>(thread_aggregate, even_share.block_offset, valid_items);
// For Warp Reduction, we need to explicitly handle the valid_items,
// whereas for Block Reduction it is implicitly handled
if constexpr (IsWarpReduction)
{
valid_items = (NumThreads <= valid_items) ? NumThreads : valid_items;
}
return CollectiveReduceT(temp_storage.reduce).Reduce(thread_aggregate, reduction_op, valid_items);
}
// Extracting this into a function saves 8% of generated kernel size by allowing to reuse
// the block reduction below. This also workaround hang in nvcc.
ConsumeFullTileRange<CanVectorize>(thread_aggregate, even_share);
// Compute block-wide reduction (all threads have valid items)
return CollectiveReduceT(temp_storage.reduce).Reduce(thread_aggregate, reduction_op);
}
/**
* @brief Reduce a contiguous segment of input tiles
* @param[in] block_offset Threadblock begin offset (inclusive)
* @param[in] block_end Threadblock end offset (exclusive)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AccumT ConsumeRange(OffsetT block_offset, OffsetT block_end)
{
GridEvenShare<OffsetT> even_share;
even_share.template BlockInit<TILE_ITEMS>(block_offset, block_end);
return IsAligned(d_in + block_offset)
? ConsumeRange<ATTEMPT_VECTORIZATION>(even_share)
: ConsumeRange<false>(even_share);
}
/**
* Reduce a contiguous segment of input tiles
* @param[in] even_share GridEvenShare descriptor
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AccumT ConsumeTiles(GridEvenShare<OffsetT>& even_share)
{
// Initialize GRID_MAPPING_STRIP_MINE even-share descriptor for this thread block
even_share.template BlockInit<TILE_ITEMS, GRID_MAPPING_STRIP_MINE>();
return IsAligned(d_in) ? ConsumeRange<ATTEMPT_VECTORIZATION>(even_share) : ConsumeRange<false>(even_share);
}
private:
/**
* @brief Reduce a contiguous segment of input tiles with more than `TILE_ITEMS` elements
* @param even_share GridEvenShare descriptor
* @param input_is_vector_aligned Whether we can vectorize loads
*/
template <bool CanVectorize>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeFullTileRange(AccumT& thread_aggregate, GridEvenShare<OffsetT>& even_share)
{
// At least one full block
ConsumeFullTile<true, CanVectorize>(thread_aggregate, even_share.block_offset);
if (even_share.block_end - even_share.block_offset < even_share.block_stride)
{
// Exit early to handle offset overflow
return;
}
even_share.block_offset += even_share.block_stride;
// Consume subsequent full tiles of input, at least one full tile was processed, so
// `even_share.block_end >= TILE_ITEMS`
while (even_share.block_offset <= even_share.block_end - TILE_ITEMS)
{
ConsumeFullTile<false, CanVectorize>(thread_aggregate, even_share.block_offset);
if (even_share.block_end - even_share.block_offset < even_share.block_stride)
{
// Exit early to handle offset overflow
return;
}
even_share.block_offset += even_share.block_stride;
}
// Consume a partially-full tile
if (even_share.block_offset < even_share.block_end)
{
int valid_items = even_share.block_end - even_share.block_offset;
ConsumePartialTile<false>(thread_aggregate, even_share.block_offset, valid_items);
}
}
};
/**
* @brief AgentReduce implements a stateful abstraction of CUDA thread blocks
* and warps, for participating in device-wide reduction .
*
* Each thread reduces only the values it loads. If `FIRST_TILE`, this partial
* reduction is stored into `thread_aggregate`. Otherwise it is accumulated
* into `thread_aggregate`.
*
* @tparam AgentReducePolicy
* Parameterized AgentReducePolicy tuning policy type
*
* @tparam InputIteratorT
* Random-access iterator type for input
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam ReductionOp
* Binary reduction operator type having member
* `auto operator()(T &&a, U &&b)`
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*
* @tparam TransformOp
* Unary operator type having member `auto operator()(T &&a)`
*/
template <typename AgentReducePolicy,
typename InputIteratorT,
typename OffsetT,
typename ReductionOp,
typename AccumT,
typename TransformOp = ::cuda::std::identity>
struct AgentReduce
: AgentReduceImpl<AgentReducePolicy,
InputIteratorT,
OffsetT,
ReductionOp,
AccumT,
TransformOp,
BlockReduce<AccumT, AgentReducePolicy::BLOCK_THREADS, AgentReducePolicy::BLOCK_ALGORITHM>,
AgentReducePolicy::BLOCK_THREADS>
{
using base_t =
AgentReduceImpl<AgentReducePolicy,
InputIteratorT,
OffsetT,
ReductionOp,
AccumT,
TransformOp,
BlockReduce<AccumT, AgentReducePolicy::BLOCK_THREADS, AgentReducePolicy::BLOCK_ALGORITHM>,
AgentReducePolicy::BLOCK_THREADS>;
_CCCL_DEVICE _CCCL_FORCEINLINE AgentReduce(
typename base_t::TempStorage& temp_storage,
InputIteratorT d_in,
ReductionOp reduction_op,
TransformOp transform_op = {})
: base_t(temp_storage, d_in, reduction_op, transform_op, threadIdx.x)
{}
};
/**
* @brief AgentWarpReduce implements a stateful abstraction of CUDA warps,
* for participating in device-wide reduction .
*
* Each thread reduces only the values it loads. If `FIRST_TILE`, this partial
* reduction is stored into `thread_aggregate`. Otherwise it is accumulated
* into `thread_aggregate`.
*
* @tparam AgentReducePolicy
* Parameterized AgentReducePolicy tuning policy type
*
* @tparam InputIteratorT
* Random-access iterator type for input
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam ReductionOp
* Binary reduction operator type having member
* `auto operator()(T &&a, U &&b)`
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*
* @tparam TransformOp
* Unary operator type having member `auto operator()(T &&a)`
*/
template <typename AgentReducePolicy,
typename InputIteratorT,
typename OffsetT,
typename ReductionOp,
typename AccumT,
typename TransformOp = ::cuda::std::identity>
struct AgentWarpReduce
: AgentReduceImpl<AgentReducePolicy,
InputIteratorT,
OffsetT,
ReductionOp,
AccumT,
TransformOp,
WarpReduce<AccumT, AgentReducePolicy::WARP_THREADS>,
AgentReducePolicy::WARP_THREADS,
true>
{
using base_t =
AgentReduceImpl<AgentReducePolicy,
InputIteratorT,
OffsetT,
ReductionOp,
AccumT,
TransformOp,
WarpReduce<AccumT, AgentReducePolicy::WARP_THREADS>,
AgentReducePolicy::WARP_THREADS,
true>;
_CCCL_DEVICE _CCCL_FORCEINLINE AgentWarpReduce(
typename base_t::TempStorage& temp_storage,
InputIteratorT d_in,
ReductionOp reduction_op,
TransformOp transform_op = {})
: base_t(temp_storage, d_in, reduction_op, transform_op, threadIdx.x % AgentReducePolicy::WARP_THREADS)
{}
};
} // namespace detail::reduce
CUB_NAMESPACE_END

View File

@@ -0,0 +1,757 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2022, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
//! @file
//! cub::detail::reduce_by_key::AgentReduceByKey implements a stateful abstraction of CUDA thread blocks for
//! participating in device-wide reduce-value-by-key.
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/single_pass_scan_operators.cuh>
#include <cub/block/block_discontinuity.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cuda/__functional/operator_properties.h>
#include <cuda/std/__functional/operations.h>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/is_pointer.h>
#include <cuda/std/__type_traits/is_same.h>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
template <int ThreadsPerBlock,
int ItemsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockScanAlgorithm ScanAlgorithm,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
struct agent_reduce_by_key_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
struct detail
{
using delay_constructor_t = DelayConstructorT;
};
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int ItemsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockScanAlgorithm ScanAlgorithm,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
using AgentReduceByKeyPolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceReduce::ReduceByKey") = detail::
agent_reduce_by_key_policy<ThreadsPerBlock, ItemsPerThread, LoadAlgorithm, LoadModifier, ScanAlgorithm, DelayConstructorT>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::reduce_by_key
{
/**
* @brief AgentReduceByKey implements a stateful abstraction of CUDA thread
* blocks for participating in device-wide reduce-value-by-key
*
* @tparam AgentReduceByKeyPolicyT
* Parameterized AgentReduceByKeyPolicy tuning policy type
*
* @tparam KeysInputIteratorT
* Random-access input iterator type for keys
*
* @tparam UniqueOutputIteratorT
* Random-access output iterator type for keys
*
* @tparam ValuesInputIteratorT
* Random-access input iterator type for values
*
* @tparam AggregatesOutputIteratorT
* Random-access output iterator type for values
*
* @tparam NumRunsOutputIteratorT
* Output iterator type for recording number of items selected
*
* @tparam EqualityOpT
* KeyT equality operator type
*
* @tparam ReductionOpT
* ValueT reduction operator type
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*/
template <typename AgentReduceByKeyPolicyT,
typename KeysInputIteratorT,
typename UniqueOutputIteratorT,
typename ValuesInputIteratorT,
typename AggregatesOutputIteratorT,
typename NumRunsOutputIteratorT,
typename EqualityOpT,
typename ReductionOpT,
typename OffsetT,
typename AccumT,
typename StreamingContextT>
struct AgentReduceByKey
{
// Whether or not this is a streaming invocation (i.e., multiple kernel invocations over partitions of the input)
static constexpr bool is_streaming_invocation = !::cuda::std::is_same_v<StreamingContextT, NullType>;
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
// The input keys type
using KeyInputT = it_value_t<KeysInputIteratorT>;
// The output keys type
using KeyOutputT = non_void_value_t<UniqueOutputIteratorT, KeyInputT>;
// The input values type
using ValueInputT = it_value_t<ValuesInputIteratorT>;
// Tuple type for scanning (pairs accumulated segment-value with
// segment-index)
using OffsetValuePairT = KeyValuePair<OffsetT, AccumT>;
// Tuple type for pairing keys and values
using KeyValuePairT = KeyValuePair<KeyOutputT, AccumT>;
// Tile status descriptor interface type
using ScanTileStateT = ReduceByKeyScanTileState<AccumT, OffsetT>;
// Guarded inequality functor
template <typename _EqualityOpT>
struct GuardedInequalityWrapper
{
/// Wrapped equality operator
_EqualityOpT op;
/// Items remaining
int num_remaining;
/// Constructor
_CCCL_HOST_DEVICE _CCCL_FORCEINLINE GuardedInequalityWrapper(_EqualityOpT op, int num_remaining)
: op(op)
, num_remaining(num_remaining)
{}
/// Boolean inequality operator, returns <tt>(a != b)</tt>
template <typename T>
_CCCL_HOST_DEVICE _CCCL_FORCEINLINE bool operator()(const T& a, const T& b, int idx) const
{
if (idx < num_remaining)
{
return !op(a, b); // In bounds
}
// Return true if first out-of-bounds item, false otherwise
return (idx == num_remaining);
}
};
// Constants
static constexpr int BLOCK_THREADS = AgentReduceByKeyPolicyT::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = AgentReduceByKeyPolicyT::ITEMS_PER_THREAD;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
static constexpr int TWO_PHASE_SCATTER = (ITEMS_PER_THREAD > 1);
// Cache-modified Input iterator wrapper type (for applying cache modifier)
// for keys Wrap the native input pointer with
// CacheModifiedValuesInputIterator or directly use the supplied input
// iterator type
using WrappedKeysInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<KeysInputIteratorT>,
CacheModifiedInputIterator<AgentReduceByKeyPolicyT::LOAD_MODIFIER, KeyInputT, OffsetT>,
KeysInputIteratorT>;
// Cache-modified Input iterator wrapper type (for applying cache modifier)
// for values Wrap the native input pointer with
// CacheModifiedValuesInputIterator or directly use the supplied input
// iterator type
using WrappedValuesInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<ValuesInputIteratorT>,
CacheModifiedInputIterator<AgentReduceByKeyPolicyT::LOAD_MODIFIER, ValueInputT, OffsetT>,
ValuesInputIteratorT>;
// Cache-modified Input iterator wrapper type (for applying cache modifier)
// for fixup values Wrap the native input pointer with
// CacheModifiedValuesInputIterator or directly use the supplied input
// iterator type
using WrappedFixupInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<AggregatesOutputIteratorT>,
CacheModifiedInputIterator<AgentReduceByKeyPolicyT::LOAD_MODIFIER, ValueInputT, OffsetT>,
AggregatesOutputIteratorT>;
// Reduce-value-by-segment scan operator
using ReduceBySegmentOpT = ReduceBySegmentOp<ReductionOpT>;
// Parameterized BlockLoad type for keys
using BlockLoadKeysT =
BlockLoad<KeyOutputT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentReduceByKeyPolicyT::LOAD_ALGORITHM>;
// Parameterized BlockLoad type for values
using BlockLoadValuesT = BlockLoad<AccumT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentReduceByKeyPolicyT::LOAD_ALGORITHM>;
// Parameterized BlockDiscontinuity type for keys
using BlockDiscontinuityKeys = BlockDiscontinuity<KeyOutputT, BLOCK_THREADS>;
// Parameterized BlockScan type
using BlockScanT = BlockScan<OffsetValuePairT, BLOCK_THREADS, AgentReduceByKeyPolicyT::SCAN_ALGORITHM>;
// Callback type for obtaining tile prefix during block scan
using DelayConstructorT = typename AgentReduceByKeyPolicyT::detail::delay_constructor_t;
using TilePrefixCallbackOpT =
TilePrefixCallbackOp<OffsetValuePairT, ReduceBySegmentOpT, ScanTileStateT, DelayConstructorT>;
// Key and value exchange types
using KeyExchangeT = KeyOutputT[TILE_ITEMS + 1];
using ValueExchangeT = AccumT[TILE_ITEMS + 1];
// Shared memory type for this thread block
union _TempStorage
{
struct ScanStorage
{
// Smem needed for tile scanning
typename BlockScanT::TempStorage scan;
// Smem needed for cooperative prefix callback
typename TilePrefixCallbackOpT::TempStorage prefix;
// Smem needed for discontinuity detection
typename BlockDiscontinuityKeys::TempStorage discontinuity;
} scan_storage;
// Smem needed for loading keys
typename BlockLoadKeysT::TempStorage load_keys;
// Smem needed for loading values
typename BlockLoadValuesT::TempStorage load_values;
// Smem needed for compacting key value pairs(allows non POD items in this
// union)
Uninitialized<KeyValuePairT[TILE_ITEMS + 1]> raw_exchange;
};
// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
/// Reference to temp_storage
_TempStorage& temp_storage;
/// Input keys
WrappedKeysInputIteratorT d_keys_in;
/// Unique output keys
UniqueOutputIteratorT d_unique_out;
/// Input values
WrappedValuesInputIteratorT d_values_in;
/// Output value aggregates
AggregatesOutputIteratorT d_aggregates_out;
/// Output pointer for total number of segments identified
NumRunsOutputIteratorT d_num_runs_out;
/// KeyT equality operator
EqualityOpT equality_op;
/// Reduction operator
ReductionOpT reduction_op;
/// Reduce-by-segment scan operator
ReduceBySegmentOpT scan_op;
/// Streaming context providing context about this partition for streaming invocations
StreamingContextT streaming_context;
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
/**
* @param temp_storage
* Reference to temp_storage
*
* @param d_keys_in
* Input keys
*
* @param d_unique_out
* Unique output keys
*
* @param d_values_in
* Input values
*
* @param d_aggregates_out
* Output value aggregates
*
* @param d_num_runs_out
* Output pointer for total number of segments identified
*
* @param equality_op
* KeyT equality operator
*
* @param reduction_op
* ValueT reduction operator
*
* @param streaming_context
* Streaming context providing context about this partition for streaming invocations
*/
template <typename StreamingContext>
_CCCL_DEVICE _CCCL_FORCEINLINE AgentReduceByKey(
TempStorage& temp_storage,
KeysInputIteratorT d_keys_in,
UniqueOutputIteratorT d_unique_out,
ValuesInputIteratorT d_values_in,
AggregatesOutputIteratorT d_aggregates_out,
NumRunsOutputIteratorT d_num_runs_out,
EqualityOpT equality_op,
ReductionOpT reduction_op,
StreamingContext streaming_context)
: temp_storage(temp_storage.Alias())
, d_keys_in(d_keys_in)
, d_unique_out(d_unique_out + streaming_context.num_uniques())
, d_values_in(d_values_in)
, d_aggregates_out(d_aggregates_out + streaming_context.num_uniques())
, d_num_runs_out(d_num_runs_out)
, equality_op(equality_op)
, reduction_op(reduction_op)
, scan_op(reduction_op)
, streaming_context(streaming_context)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE AgentReduceByKey(
TempStorage& temp_storage,
KeysInputIteratorT d_keys_in,
UniqueOutputIteratorT d_unique_out,
ValuesInputIteratorT d_values_in,
AggregatesOutputIteratorT d_aggregates_out,
NumRunsOutputIteratorT d_num_runs_out,
EqualityOpT equality_op,
ReductionOpT reduction_op,
NullType streaming_context)
: temp_storage(temp_storage.Alias())
, d_keys_in(d_keys_in)
, d_unique_out(d_unique_out)
, d_values_in(d_values_in)
, d_aggregates_out(d_aggregates_out)
, d_num_runs_out(d_num_runs_out)
, equality_op(equality_op)
, reduction_op(reduction_op)
, scan_op(reduction_op)
, streaming_context(streaming_context)
{}
//---------------------------------------------------------------------
// Scatter utility methods
//---------------------------------------------------------------------
/**
* Directly scatter flagged items to output offsets
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterDirect(
KeyValuePairT (&scatter_items)[ITEMS_PER_THREAD],
OffsetT (&segment_flags)[ITEMS_PER_THREAD],
OffsetT (&segment_indices)[ITEMS_PER_THREAD])
{
// Scatter flagged keys and values
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
if (segment_flags[ITEM])
{
d_unique_out[segment_indices[ITEM]] = scatter_items[ITEM].key;
d_aggregates_out[segment_indices[ITEM]] = scatter_items[ITEM].value;
}
}
}
/**
* 2-phase scatter flagged items to output offsets
*
* The exclusive scan causes each head flag to be paired with the previous
* value aggregate: the scatter offsets must be decremented for value
* aggregates
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ScatterTwoPhase(
KeyValuePairT (&scatter_items)[ITEMS_PER_THREAD],
OffsetT (&segment_flags)[ITEMS_PER_THREAD],
OffsetT (&segment_indices)[ITEMS_PER_THREAD],
OffsetT num_tile_segments,
OffsetT num_tile_segments_prefix)
{
__syncthreads();
// Compact and scatter pairs
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
if (segment_flags[ITEM])
{
temp_storage.raw_exchange.Alias()[segment_indices[ITEM] - num_tile_segments_prefix] = scatter_items[ITEM];
}
}
__syncthreads();
for (int item = static_cast<int>(threadIdx.x); item < num_tile_segments; item += BLOCK_THREADS)
{
KeyValuePairT pair = temp_storage.raw_exchange.Alias()[item];
d_unique_out[num_tile_segments_prefix + item] = pair.key; // NOLINT(bugprone-misplaced-widening-cast)
d_aggregates_out[num_tile_segments_prefix + item] = pair.value; // NOLINT(bugprone-misplaced-widening-cast)
}
}
/**
* Scatter flagged items
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void Scatter(
KeyValuePairT (&scatter_items)[ITEMS_PER_THREAD],
OffsetT (&segment_flags)[ITEMS_PER_THREAD],
OffsetT (&segment_indices)[ITEMS_PER_THREAD],
OffsetT num_tile_segments,
OffsetT num_tile_segments_prefix)
{
// Do a one-phase scatter if (a) two-phase is disabled or (b) the average
// number of selected items per thread is less than one
if (TWO_PHASE_SCATTER && (num_tile_segments > BLOCK_THREADS))
{
ScatterTwoPhase(scatter_items, segment_flags, segment_indices, num_tile_segments, num_tile_segments_prefix);
}
else
{
ScatterDirect(scatter_items, segment_flags, segment_indices);
}
}
//---------------------------------------------------------------------
// Cooperatively scan a device-wide sequence of tiles with other CTAs
//---------------------------------------------------------------------
/**
* @brief Process a tile of input (dynamic chained scan)
*
* @tparam IS_LAST_TILE
* Whether the current tile is the last tile
*
* @param num_remaining
* Number of global input items remaining (including this tile)
*
* @param tile_idx
* Tile index
*
* @param tile_offset
* Tile offset
*
* @param tile_state
* Global tile state descriptor
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeTile(OffsetT num_remaining, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state)
{
// Tile keys
KeyOutputT keys[ITEMS_PER_THREAD];
// Tile keys shuffled up
KeyOutputT prev_keys[ITEMS_PER_THREAD];
// Tile values
AccumT values[ITEMS_PER_THREAD];
// Segment head flags
OffsetT head_flags[ITEMS_PER_THREAD];
// Segment indices
OffsetT segment_indices[ITEMS_PER_THREAD];
// Zipped values and segment flags|indices
OffsetValuePairT scan_items[ITEMS_PER_THREAD];
// Zipped key value pairs for scattering
KeyValuePairT scatter_items[ITEMS_PER_THREAD];
// Load keys
if (IS_LAST_TILE)
{
BlockLoadKeysT(temp_storage.load_keys).Load(d_keys_in + tile_offset, keys, num_remaining);
}
else
{
BlockLoadKeysT(temp_storage.load_keys).Load(d_keys_in + tile_offset, keys);
}
// Load tile predecessor key in first thread
KeyOutputT tile_predecessor;
if (threadIdx.x == 0)
{
// if (tile_idx == 0)
// first tile gets repeat of first item (thus first item will not
// be flagged as a head)
// else
// Subsequent tiles get last key from previous tile
if constexpr (is_streaming_invocation)
{
tile_predecessor = (tile_idx == 0) ? streaming_context.predecessor_key() : d_keys_in[tile_offset - 1];
}
else
{
tile_predecessor = (tile_idx == 0) ? keys[0] : d_keys_in[tile_offset - 1];
}
}
__syncthreads();
// Load values
if (IS_LAST_TILE)
{
BlockLoadValuesT(temp_storage.load_values).Load(d_values_in + tile_offset, values, num_remaining);
}
else
{
BlockLoadValuesT(temp_storage.load_values).Load(d_values_in + tile_offset, values);
}
__syncthreads();
// Initialize head-flags and shuffle up the previous keys
if (IS_LAST_TILE)
{
// Use custom flag operator to additionally flag the first out-of-bounds item
GuardedInequalityWrapper<EqualityOpT> flag_op(equality_op, num_remaining);
BlockDiscontinuityKeys(temp_storage.scan_storage.discontinuity)
.FlagHeads(head_flags, keys, prev_keys, flag_op, tile_predecessor);
}
else
{
InequalityWrapper<EqualityOpT> flag_op(equality_op);
BlockDiscontinuityKeys(temp_storage.scan_storage.discontinuity)
.FlagHeads(head_flags, keys, prev_keys, flag_op, tile_predecessor);
}
// Reset head-flag on the very first item to make sure we don't start a new run for data where
// (key[0] == key[0]) is false (e.g., when key[0] is NaN)
if constexpr (is_streaming_invocation)
{
if (streaming_context.is_first_partition() && threadIdx.x == 0 && tile_idx == 0)
{
head_flags[0] = 0;
}
}
else
{
if (threadIdx.x == 0 && tile_idx == 0)
{
head_flags[0] = 0;
}
}
// Zip values and head flags
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
scan_items[ITEM].value = values[ITEM];
scan_items[ITEM].key = head_flags[ITEM];
}
// Perform exclusive tile scan
// Inclusive block-wide scan aggregate
OffsetValuePairT block_aggregate;
// Number of segments prior to this tile
OffsetT num_segments_prefix;
// The tile prefix folded with block_aggregate
OffsetValuePairT total_aggregate;
if (tile_idx == 0)
{
// Scan first tile
// First partition does not need to account for preceding partitions
if constexpr (is_streaming_invocation)
{
if (streaming_context.is_first_partition())
{
BlockScanT(temp_storage.scan_storage.scan).ExclusiveScan(scan_items, scan_items, scan_op, block_aggregate);
num_segments_prefix = 0;
total_aggregate = block_aggregate;
}
// Subsequent partitions need to account for preceding partitions
else
{
auto init_value = OffsetValuePairT{0, streaming_context.prefix()};
BlockScanT(temp_storage.scan_storage.scan)
.ExclusiveScan(scan_items, scan_items, init_value, scan_op, block_aggregate);
num_segments_prefix = 0;
// note, block_aggregate does not include the prefix
block_aggregate = scan_op(init_value, block_aggregate);
total_aggregate = block_aggregate;
}
}
else
{
BlockScanT(temp_storage.scan_storage.scan).ExclusiveScan(scan_items, scan_items, scan_op, block_aggregate);
num_segments_prefix = 0;
total_aggregate = block_aggregate;
}
// Update tile status if there are successor tiles
if ((!IS_LAST_TILE) && (threadIdx.x == 0))
{
tile_state.SetInclusive(0, block_aggregate);
}
}
else
{
// Scan non-first tile
TilePrefixCallbackOpT prefix_op(tile_state, temp_storage.scan_storage.prefix, scan_op, tile_idx);
BlockScanT(temp_storage.scan_storage.scan).ExclusiveScan(scan_items, scan_items, scan_op, prefix_op);
block_aggregate = prefix_op.GetBlockAggregate();
num_segments_prefix = prefix_op.GetExclusivePrefix().key;
total_aggregate = prefix_op.GetInclusivePrefix();
}
// Rezip scatter items and segment indices
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
scatter_items[ITEM].key = prev_keys[ITEM];
scatter_items[ITEM].value = scan_items[ITEM].value;
segment_indices[ITEM] = scan_items[ITEM].key;
}
// At this point, each flagged segment head has:
// - The key for the previous segment
// - The reduced value from the previous segment
// - The segment index for the reduced value
// Scatter flagged keys and values
OffsetT num_tile_segments = block_aggregate.key;
Scatter(scatter_items, head_flags, segment_indices, num_tile_segments, num_segments_prefix);
// Last thread in last tile will output final count (and last pair, if necessary)
if ((IS_LAST_TILE) && (threadIdx.x == BLOCK_THREADS - 1))
{
OffsetT num_segments = num_segments_prefix + num_tile_segments;
// If the last tile is a full tile, we need to write out the run ending with the last item
// If this was not a full tile, we already have flagged the head of one-past-the-last-item
if (num_remaining == TILE_ITEMS)
{
if constexpr (is_streaming_invocation)
{
if (streaming_context.is_last_partition())
{
d_unique_out[num_segments] = keys[ITEMS_PER_THREAD - 1];
d_aggregates_out[num_segments] = total_aggregate.value;
num_segments++;
}
else
{
// Write the prefix aggregate of this partition as context for the subsequent partition
streaming_context.write_prefix(total_aggregate.value);
}
}
else
{
d_unique_out[num_segments] = keys[ITEMS_PER_THREAD - 1];
d_aggregates_out[num_segments] = total_aggregate.value;
num_segments++;
}
}
if constexpr (is_streaming_invocation)
{
// Add the number of unique items in this partition to the global aggregate
auto total_uniques = streaming_context.add_num_uniques(num_segments);
// If this is the last partition, write out the number of unique items
if (streaming_context.is_last_partition())
{
*d_num_runs_out = total_uniques;
}
}
else
{
*d_num_runs_out = num_segments;
}
}
}
/**
* @brief Scan tiles of items as part of a dynamic chained scan
*
* @param num_items
* Total number of input items
*
* @param tile_state
* Global tile state descriptor
*
* @param start_tile
* The starting tile for the current grid
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeRange(OffsetT num_items, ScanTileStateT& tile_state, int start_tile)
{
// Blocks are launched in increasing order, so just assign one tile per
// block
// Current tile index
int tile_idx = static_cast<int>(start_tile + blockIdx.x);
// Global offset for the current tile
OffsetT tile_offset = OffsetT(TILE_ITEMS) * tile_idx;
// Remaining items (including this tile)
OffsetT num_remaining = num_items - tile_offset;
if (num_remaining > TILE_ITEMS)
{
// Not last tile
ConsumeTile<false>(num_remaining, tile_idx, tile_offset, tile_state);
}
else if (num_remaining > 0)
{
// Last tile
ConsumeTile<true>(num_remaining, tile_idx, tile_offset, tile_state);
}
}
};
} // namespace detail::reduce_by_key
CUB_NAMESPACE_END

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,567 @@
// SPDX-FileCopyrightText: Copyright (c) 2011, Duane Merrill. All rights reserved.
// SPDX-FileCopyrightText: Copyright (c) 2011-2022, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* @file
* @brief cub::AgentScan implements a stateful abstraction of CUDA thread blocks
* for participating in device-wide prefix scan.
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/single_pass_scan_operators.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/grid/grid_queue.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_device.cuh>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/is_pointer.h>
#include <cuda/std/__type_traits/is_same.h>
CUB_NAMESPACE_BEGIN
namespace detail
{
// TODO(bgruber): remove when C++20 is the minimum, since then we can pass policy values as NTTPs
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockStoreAlgorithm StoreAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
typename ScalingType = detail::MemBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>,
typename DelayConstructorT = detail::default_delay_constructor_t<ComputeT>>
struct agent_scan_policy : ScalingType
{
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr BlockStoreAlgorithm STORE_ALGORITHM = StoreAlgorithm;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
struct detail
{
using delay_constructor_t = DelayConstructorT;
};
};
} // namespace detail
/******************************************************************************
* Tuning policy types
******************************************************************************/
/**
* @brief Parameterizable tuning policy type for AgentScan
*
* @tparam NominalThreadsPerBlock4B
* Threads per thread block
*
* @tparam NominalItemsPerThread4B
* Items per thread (per tile of input)
*
* @tparam ComputeT
* Dominant compute type
*
* @tparam LoadAlgorithm
* The BlockLoad algorithm to use
*
* @tparam LoadModifier
* Cache load modifier for reading input elements
*
* @tparam StoreAlgorithm
* The BlockStore algorithm to use
*
* @tparam ScanAlgorithm
* The BlockScan algorithm to use
*
* @tparam DelayConstructorT
* Implementation detail, do not specify directly, requirements on the
* content of this type are subject to breaking change.
*/
//! Deprecated [Since 3.5]
template <int NominalThreadsPerBlock4B,
int NominalItemsPerThread4B,
typename ComputeT,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockStoreAlgorithm StoreAlgorithm,
BlockScanAlgorithm ScanAlgorithm,
typename ScalingType = detail::MemBoundScaling<NominalThreadsPerBlock4B, NominalItemsPerThread4B, ComputeT>,
typename DelayConstructorT = detail::default_delay_constructor_t<ComputeT>>
using AgentScanPolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceScan") = detail::agent_scan_policy<
NominalThreadsPerBlock4B,
NominalItemsPerThread4B,
ComputeT,
LoadAlgorithm,
LoadModifier,
StoreAlgorithm,
ScanAlgorithm,
ScalingType,
DelayConstructorT>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::scan
{
/**
* @brief AgentScan implements a stateful abstraction of CUDA thread blocks for
* participating in device-wide prefix scan.
* @tparam AgentScanPolicyT
* Parameterized AgentScanPolicyT tuning policy type
*
* @tparam InputIteratorT
* Random-access input iterator type
*
* @tparam OutputIteratorT
* Random-access output iterator type
*
* @tparam ScanOpT
* Scan functor type
*
* @tparam InitValueT
* The init_value element for ScanOpT type (cub::NullType for inclusive scan)
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*/
template <typename AgentScanPolicyT,
typename InputIteratorT,
typename OutputIteratorT,
typename ScanOpT,
typename InitValueT,
typename OffsetT,
typename AccumT,
bool ForceInclusive = false,
bool UsePDL = false,
bool StableReductionOrder = false>
struct AgentScan
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
// The input value type
using InputT = cub::detail::it_value_t<InputIteratorT>;
// Tile status descriptor interface type
using ScanTileStateT = ScanTileState<AccumT>;
// Input iterator wrapper type (for applying cache modifier)
// Wrap the native input pointer with CacheModifiedInputIterator
// or directly use the supplied input iterator type
using WrappedInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<InputIteratorT>,
CacheModifiedInputIterator<AgentScanPolicyT::LOAD_MODIFIER, InputT, OffsetT>,
InputIteratorT>;
// Inclusive scan if no init_value type is provided
static constexpr bool HAS_INIT = !::cuda::std::is_same_v<InitValueT, NullType>;
static constexpr bool IS_INCLUSIVE = ForceInclusive || !HAS_INIT; // We are relying on either initial value not being
// `NullType` or the ForceInclusive tag to be true
// for inclusive scan to get picked up.
static constexpr int BLOCK_THREADS = AgentScanPolicyT::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = AgentScanPolicyT::ITEMS_PER_THREAD;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
// Parameterized BlockLoad type
using BlockLoadT =
BlockLoad<AccumT,
AgentScanPolicyT::BLOCK_THREADS,
AgentScanPolicyT::ITEMS_PER_THREAD,
AgentScanPolicyT::LOAD_ALGORITHM>;
// Parameterized BlockStore type
using BlockStoreT =
BlockStore<AccumT,
AgentScanPolicyT::BLOCK_THREADS,
AgentScanPolicyT::ITEMS_PER_THREAD,
AgentScanPolicyT::STORE_ALGORITHM>;
// Parameterized BlockScan type
using BlockScanT = BlockScan<AccumT, AgentScanPolicyT::BLOCK_THREADS, AgentScanPolicyT::SCAN_ALGORITHM>;
// Callback type for obtaining tile prefix during block scan
using DelayConstructorT = typename AgentScanPolicyT::detail::delay_constructor_t;
using TilePrefixCallbackOpT =
TilePrefixCallbackOp<AccumT, ScanOpT, ScanTileStateT, DelayConstructorT, StableReductionOrder>;
// Stateful BlockScan prefix callback type for managing a running total while
// scanning consecutive tiles
using RunningPrefixCallbackOp = BlockScanRunningPrefixOp<AccumT, ScanOpT>;
// Shared memory type for this thread block
union _TempStorage
{
// Smem needed for tile loading
typename BlockLoadT::TempStorage load;
// Smem needed for tile storing
typename BlockStoreT::TempStorage store;
struct ScanStorage
{
// Smem needed for cooperative prefix callback
typename TilePrefixCallbackOpT::TempStorage prefix;
// Smem needed for tile scanning
typename BlockScanT::TempStorage scan;
} scan_storage;
};
// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
_TempStorage& temp_storage; ///< Reference to temp_storage
WrappedInputIteratorT d_in; ///< Input data
OutputIteratorT d_out; ///< Output data
ScanOpT scan_op; ///< Binary scan operator
InitValueT init_value; ///< The init_value element for ScanOpT
//---------------------------------------------------------------------
// Block scan utility methods
//---------------------------------------------------------------------
template <bool Inclusive = IS_INCLUSIVE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScanFirstTile(AccumT (&items)[ITEMS_PER_THREAD], InitValueT init_value, ScanOpT scan_op, AccumT& block_aggregate)
{
BlockScanT blockScan(temp_storage.scan_storage.scan);
if constexpr (Inclusive)
{
if constexpr (HAS_INIT)
{
blockScan.InclusiveScan(items, items, init_value, scan_op, block_aggregate);
block_aggregate = scan_op(init_value, block_aggregate);
}
else
{
blockScan.InclusiveScan(items, items, scan_op, block_aggregate);
}
}
else
{
blockScan.ExclusiveScan(items, items, init_value, scan_op, block_aggregate);
block_aggregate = scan_op(init_value, block_aggregate);
}
}
template <typename PrefixCallback, bool Inclusive = IS_INCLUSIVE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScanSubsequentTile(AccumT (&items)[ITEMS_PER_THREAD], ScanOpT scan_op, PrefixCallback& prefix_op)
{
BlockScanT blockScan(temp_storage.scan_storage.scan);
if constexpr (Inclusive)
{
blockScan.InclusiveScan(items, items, scan_op, prefix_op);
}
else
{
blockScan.ExclusiveScan(items, items, scan_op, prefix_op);
}
}
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
/**
* @param temp_storage
* Reference to temp_storage
*
* @param d_in
* Input data
*
* @param d_out
* Output data
*
* @param scan_op
* Binary scan operator
*
* @param init_value
* Initial value to seed the exclusive scan
*/
_CCCL_DEVICE _CCCL_FORCEINLINE AgentScan(
TempStorage& temp_storage, InputIteratorT d_in, OutputIteratorT d_out, ScanOpT scan_op, InitValueT init_value)
: temp_storage(temp_storage.Alias())
, d_in(d_in)
, d_out(d_out)
, scan_op(scan_op)
, init_value(init_value)
{}
//---------------------------------------------------------------------
// Cooperatively scan a device-wide sequence of tiles with other CTAs
//---------------------------------------------------------------------
/**
* Process a tile of input (dynamic chained scan)
* @tparam IS_LAST_TILE
* Whether the current tile is the last tile
*
* @param num_remaining
* Number of global input items remaining (including this tile)
*
* @param tile_idx
* Tile index
*
* @param tile_offset
* Tile offset
*
* @param tile_state
* Global tile state descriptor
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeTile(OffsetT num_remaining, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state)
{
// Load items
AccumT items[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last element with the first element because collectives are
// not suffix guarded.
BlockLoadT(temp_storage.load).Load(d_in + tile_offset, items, num_remaining, *(d_in + tile_offset));
}
else
{
BlockLoadT(temp_storage.load).Load(d_in + tile_offset, items);
}
__syncthreads();
// Perform tile scan
if (tile_idx == 0)
{
// Scan first tile
AccumT block_aggregate;
ScanFirstTile(items, init_value, scan_op, block_aggregate);
if ((!IS_LAST_TILE) && (threadIdx.x == 0))
{
tile_state.SetInclusive(0, block_aggregate);
}
}
else
{
// Scan non-first tile
TilePrefixCallbackOpT prefix_op(tile_state, temp_storage.scan_storage.prefix, scan_op, tile_idx);
ScanSubsequentTile(items, scan_op, prefix_op);
}
__syncthreads();
if constexpr (UsePDL)
{
_CCCL_PDL_TRIGGER_NEXT_LAUNCH(); // omitting makes almost no difference in cub.bench.scan.exclusive.sum.base
}
// Store items
if constexpr (IS_LAST_TILE)
{
BlockStoreT(temp_storage.store).Store(d_out + tile_offset, items, num_remaining);
}
else
{
BlockStoreT(temp_storage.store).Store(d_out + tile_offset, items);
}
}
/**
* @brief Scan tiles of items as part of a dynamic chained scan
*
* @param num_items
* Total number of input items
*
* @param tile_state
* Global tile state descriptor
*
* @param start_tile
* The starting tile for the current grid
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeRange(OffsetT num_items, ScanTileStateT& tile_state, int start_tile)
{
// Blocks are launched in increasing order, so just assign one tile per
// block
// Current tile index
int tile_idx = static_cast<int>(start_tile + blockIdx.x);
// Global offset for the current tile
OffsetT tile_offset = OffsetT(TILE_ITEMS) * tile_idx;
// Remaining items (including this tile)
OffsetT num_remaining = num_items - tile_offset;
if (num_remaining > TILE_ITEMS)
{
// Not last tile
ConsumeTile<false>(num_remaining, tile_idx, tile_offset, tile_state);
}
else if (num_remaining > 0)
{
// Last tile
ConsumeTile<true>(num_remaining, tile_idx, tile_offset, tile_state);
}
}
//---------------------------------------------------------------------------
// Scan an sequence of consecutive tiles (independent of other thread blocks)
//---------------------------------------------------------------------------
/**
* @brief Process a tile of input
*
* @param tile_offset
* Tile offset
*
* @param prefix_op
* Running prefix operator
*
* @param valid_items
* Number of valid items in the tile
*/
template <bool IS_FIRST_TILE, bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeTile(OffsetT tile_offset, RunningPrefixCallbackOp& prefix_op, int valid_items = TILE_ITEMS)
{
// Load items
AccumT items[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last element with the first element because collectives are
// not suffix guarded.
BlockLoadT(temp_storage.load).Load(d_in + tile_offset, items, valid_items, *(d_in + tile_offset));
}
else
{
BlockLoadT(temp_storage.load).Load(d_in + tile_offset, items);
}
__syncthreads();
// Block scan
if constexpr (IS_FIRST_TILE)
{
AccumT block_aggregate;
ScanFirstTile(items, init_value, scan_op, block_aggregate);
prefix_op.running_total = block_aggregate;
}
else
{
ScanSubsequentTile(items, scan_op, prefix_op);
}
__syncthreads();
// Store items
if constexpr (IS_LAST_TILE)
{
BlockStoreT(temp_storage.store).Store(d_out + tile_offset, items, valid_items);
}
else
{
BlockStoreT(temp_storage.store).Store(d_out + tile_offset, items);
}
}
/**
* @brief Scan a consecutive share of input tiles
*
* @param[in] range_offset
* Threadblock begin offset (inclusive)
*
* @param[in] range_end
* Threadblock end offset (exclusive)
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeRange(OffsetT range_offset, OffsetT range_end)
{
BlockScanRunningPrefixOp<AccumT, ScanOpT> prefix_op(scan_op);
if (range_offset + TILE_ITEMS <= range_end)
{
// Consume first tile of input (full)
ConsumeTile<true, true>(range_offset, prefix_op);
range_offset += TILE_ITEMS;
// Consume subsequent full tiles of input
while (range_offset + TILE_ITEMS <= range_end)
{
ConsumeTile<false, true>(range_offset, prefix_op);
range_offset += TILE_ITEMS;
}
// Consume a partially-full tile
if (range_offset < range_end)
{
int valid_items = range_end - range_offset;
ConsumeTile<false, false>(range_offset, prefix_op, valid_items);
}
}
else
{
// Consume the first tile of input (partially-full)
int valid_items = range_end - range_offset;
ConsumeTile<true, false>(range_offset, prefix_op, valid_items);
}
}
/**
* @brief Scan a consecutive share of input tiles, seeded with the
* specified prefix value
* @param[in] range_offset
* Threadblock begin offset (inclusive)
*
* @param[in] range_end
* Threadblock end offset (exclusive)
*
* @param[in] prefix
* The prefix to apply to the scan segment
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeRange(OffsetT range_offset, OffsetT range_end, AccumT prefix)
{
BlockScanRunningPrefixOp<AccumT, ScanOpT> prefix_op(prefix, scan_op);
// Consume full tiles of input
while (range_offset + TILE_ITEMS <= range_end)
{
ConsumeTile<true, false>(range_offset, prefix_op);
range_offset += TILE_ITEMS;
}
// Consume a partially-full tile
if (range_offset < range_end)
{
int valid_items = range_end - range_offset;
ConsumeTile<false, false>(range_offset, prefix_op, valid_items);
}
}
};
} // namespace detail::scan
CUB_NAMESPACE_END

View File

@@ -0,0 +1,464 @@
// SPDX-FileCopyrightText: Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* @file
* @brief AgentScanByKey implements a stateful abstraction of CUDA thread blocks
* for participating in device-wide prefix scan by key.
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/single_pass_scan_operators.cuh>
#include <cub/block/block_discontinuity.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_type.cuh>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/enable_if.h>
#include <cuda/std/__type_traits/integral_constant.h>
#include <cuda/std/__type_traits/is_pointer.h>
#include <cuda/std/__type_traits/is_same.h>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
// TODO(bgruber): remove this when C++20 is the minimum, since then we can pass policy values as NTTP
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
BlockLoadAlgorithm LoadAlgorithm = BLOCK_LOAD_DIRECT,
CacheLoadModifier LoadModifier = LOAD_DEFAULT,
BlockScanAlgorithm ScanAlgorithm = BLOCK_SCAN_WARP_SCANS,
BlockStoreAlgorithm StoreAlgorithm = BLOCK_STORE_DIRECT,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
struct agent_scan_by_key_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
static constexpr BlockStoreAlgorithm STORE_ALGORITHM = StoreAlgorithm;
struct detail
{
using delay_constructor_t = DelayConstructorT;
};
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
BlockLoadAlgorithm LoadAlgorithm = BLOCK_LOAD_DIRECT,
CacheLoadModifier LoadModifier = LOAD_DEFAULT,
BlockScanAlgorithm ScanAlgorithm = BLOCK_SCAN_WARP_SCANS,
BlockStoreAlgorithm StoreAlgorithm = BLOCK_STORE_DIRECT,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
using AgentScanByKeyPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceScanByKey") = detail::agent_scan_by_key_policy<
ThreadsPerBlock,
ItemsPerThread,
LoadAlgorithm,
LoadModifier,
ScanAlgorithm,
StoreAlgorithm,
DelayConstructorT>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::scan_by_key
{
/**
* @brief AgentScanByKey implements a stateful abstraction of CUDA thread
* blocks for participating in device-wide prefix scan by key.
*
* @tparam AgentScanByKeyPolicyT
* Parameterized AgentScanPolicyT tuning policy type
*
* @tparam KeysInputIteratorT
* Random-access input iterator type
*
* @tparam ValuesInputIteratorT
* Random-access input iterator type
*
* @tparam ValuesOutputIteratorT
* Random-access output iterator type
*
* @tparam EqualityOp
* Equality functor type
*
* @tparam ScanOpT
* Scan functor type
*
* @tparam InitValueT
* The init_value element for ScanOpT type (cub::NullType for inclusive scan)
*
* @tparam OffsetT
* Signed integer type for global offsets
*
* @tparam AccumT
* The type of intermediate accumulator (according to P2322R6)
*/
template <typename AgentScanByKeyPolicyT,
typename KeysInputIteratorT,
typename ValuesInputIteratorT,
typename ValuesOutputIteratorT,
typename EqualityOp,
typename ScanOpT,
typename InitValueT,
typename OffsetT,
typename AccumT>
struct AgentScanByKey
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
using KeyT = it_value_t<KeysInputIteratorT>;
using InputT = it_value_t<ValuesInputIteratorT>;
using FlagValuePairT = KeyValuePair<int, AccumT>;
using ReduceBySegmentOpT = ScanBySegmentOp<ScanOpT>;
using ScanTileStateT = ReduceByKeyScanTileState<AccumT, int>;
// Constants
// Inclusive scan if no init_value type is provided
static constexpr int IS_INCLUSIVE = ::cuda::std::is_same_v<InitValueT, NullType>;
static constexpr int BLOCK_THREADS = AgentScanByKeyPolicyT::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = AgentScanByKeyPolicyT::ITEMS_PER_THREAD;
static constexpr int ITEMS_PER_TILE = BLOCK_THREADS * ITEMS_PER_THREAD;
using WrappedKeysInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<KeysInputIteratorT>,
CacheModifiedInputIterator<AgentScanByKeyPolicyT::LOAD_MODIFIER, KeyT, OffsetT>,
KeysInputIteratorT>;
using WrappedValuesInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<ValuesInputIteratorT>,
CacheModifiedInputIterator<AgentScanByKeyPolicyT::LOAD_MODIFIER, InputT, OffsetT>,
ValuesInputIteratorT>;
using BlockLoadKeysT = BlockLoad<KeyT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentScanByKeyPolicyT::LOAD_ALGORITHM>;
using BlockLoadValuesT = BlockLoad<AccumT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentScanByKeyPolicyT::LOAD_ALGORITHM>;
using BlockStoreValuesT = BlockStore<AccumT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentScanByKeyPolicyT::STORE_ALGORITHM>;
using BlockDiscontinuityKeysT = BlockDiscontinuity<KeyT, BLOCK_THREADS, 1, 1>;
using DelayConstructorT = typename AgentScanByKeyPolicyT::detail::delay_constructor_t;
using TilePrefixCallbackT =
TilePrefixCallbackOp<FlagValuePairT, ReduceBySegmentOpT, ScanTileStateT, DelayConstructorT>;
using BlockScanT = BlockScan<FlagValuePairT, BLOCK_THREADS, AgentScanByKeyPolicyT::SCAN_ALGORITHM, 1, 1>;
union TempStorage_
{
struct ScanStorage
{
typename BlockScanT::TempStorage scan;
typename TilePrefixCallbackT::TempStorage prefix;
typename BlockDiscontinuityKeysT::TempStorage discontinuity;
} scan_storage;
typename BlockLoadKeysT::TempStorage load_keys;
typename BlockLoadValuesT::TempStorage load_values;
typename BlockStoreValuesT::TempStorage store_values;
};
struct TempStorage : cub::Uninitialized<TempStorage_>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
TempStorage_& storage;
WrappedKeysInputIteratorT d_keys_in;
KeyT* d_keys_prev_in;
WrappedValuesInputIteratorT d_values_in;
ValuesOutputIteratorT d_values_out;
InequalityWrapper<EqualityOp> inequality_op;
ScanOpT scan_op;
ReduceBySegmentOpT pair_scan_op;
InitValueT init_value;
//---------------------------------------------------------------------
// Block scan utility methods (first tile)
//---------------------------------------------------------------------
// Exclusive scan specialization
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScanTile(FlagValuePairT (&scan_items)[ITEMS_PER_THREAD],
FlagValuePairT& tile_aggregate,
::cuda::std::false_type /* is_inclusive */)
{
BlockScanT(storage.scan_storage.scan).ExclusiveScan(scan_items, scan_items, pair_scan_op, tile_aggregate);
}
// Inclusive scan specialization
_CCCL_DEVICE _CCCL_FORCEINLINE void
ScanTile(FlagValuePairT (&scan_items)[ITEMS_PER_THREAD],
FlagValuePairT& tile_aggregate,
::cuda::std::true_type /* is_inclusive */)
{
BlockScanT(storage.scan_storage.scan).InclusiveScan(scan_items, scan_items, pair_scan_op, tile_aggregate);
}
//---------------------------------------------------------------------
// Block scan utility methods (subsequent tiles)
//---------------------------------------------------------------------
// Exclusive scan specialization (with prefix from predecessors)
_CCCL_DEVICE _CCCL_FORCEINLINE void ScanTile(
FlagValuePairT (&scan_items)[ITEMS_PER_THREAD],
FlagValuePairT& tile_aggregate,
TilePrefixCallbackT& prefix_op,
::cuda::std::false_type /* is_inclusive */)
{
BlockScanT(storage.scan_storage.scan).ExclusiveScan(scan_items, scan_items, pair_scan_op, prefix_op);
tile_aggregate = prefix_op.GetBlockAggregate();
}
// Inclusive scan specialization (with prefix from predecessors)
_CCCL_DEVICE _CCCL_FORCEINLINE void ScanTile(
FlagValuePairT (&scan_items)[ITEMS_PER_THREAD],
FlagValuePairT& tile_aggregate,
TilePrefixCallbackT& prefix_op,
::cuda::std::true_type /* is_inclusive */)
{
BlockScanT(storage.scan_storage.scan).InclusiveScan(scan_items, scan_items, pair_scan_op, prefix_op);
tile_aggregate = prefix_op.GetBlockAggregate();
}
//---------------------------------------------------------------------
// Zip utility methods
//---------------------------------------------------------------------
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ZipValuesAndFlags(
OffsetT num_remaining,
AccumT (&values)[ITEMS_PER_THREAD],
OffsetT (&segment_flags)[ITEMS_PER_THREAD],
FlagValuePairT (&scan_items)[ITEMS_PER_THREAD])
{
// Zip values and segment_flags
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
// Set segment_flags for first out-of-bounds item, zero for others
if (IS_LAST_TILE && OffsetT(threadIdx.x * ITEMS_PER_THREAD) + ITEM == num_remaining)
{
segment_flags[ITEM] = 1;
}
scan_items[ITEM].value = values[ITEM];
scan_items[ITEM].key = segment_flags[ITEM];
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void
UnzipValues(AccumT (&values)[ITEMS_PER_THREAD], FlagValuePairT (&scan_items)[ITEMS_PER_THREAD])
{
// Unzip values and segment_flags
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
values[ITEM] = scan_items[ITEM].value;
}
}
template <bool IsNull = ::cuda::std::is_same_v<InitValueT, NullType>, ::cuda::std::enable_if_t<!IsNull, int> = 0>
_CCCL_DEVICE _CCCL_FORCEINLINE void
AddInitToScan(AccumT (&items)[ITEMS_PER_THREAD], OffsetT (&flags)[ITEMS_PER_THREAD])
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
items[ITEM] = flags[ITEM] ? init_value : scan_op(init_value, items[ITEM]);
}
}
template <bool IsNull = ::cuda::std::is_same_v<InitValueT, NullType>, ::cuda::std::enable_if_t<IsNull, int> = 0>
_CCCL_DEVICE _CCCL_FORCEINLINE void
AddInitToScan(AccumT (& /*items*/)[ITEMS_PER_THREAD], OffsetT (& /*flags*/)[ITEMS_PER_THREAD])
{}
//---------------------------------------------------------------------
// Cooperatively scan a device-wide sequence of tiles with other CTAs
//---------------------------------------------------------------------
// Process a tile of input (dynamic chained scan)
//
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeTile(OffsetT /*num_items*/, OffsetT num_remaining, int tile_idx, OffsetT tile_base, ScanTileStateT& tile_state)
{
// Load items
KeyT keys[ITEMS_PER_THREAD];
AccumT values[ITEMS_PER_THREAD];
OffsetT segment_flags[ITEMS_PER_THREAD];
FlagValuePairT scan_items[ITEMS_PER_THREAD];
if (IS_LAST_TILE)
{
// Fill last element with the first element
// because collectives are not suffix guarded
BlockLoadKeysT(storage.load_keys).Load(d_keys_in + tile_base, keys, num_remaining, *(d_keys_in + tile_base));
}
else
{
BlockLoadKeysT(storage.load_keys).Load(d_keys_in + tile_base, keys);
}
__syncthreads();
if (IS_LAST_TILE)
{
// Fill last element with the first element
// because collectives are not suffix guarded
BlockLoadValuesT(storage.load_values)
.Load(d_values_in + tile_base, values, num_remaining, *(d_values_in + tile_base));
}
else
{
BlockLoadValuesT(storage.load_values).Load(d_values_in + tile_base, values);
}
__syncthreads();
// first tile
if (tile_idx == 0)
{
BlockDiscontinuityKeysT(storage.scan_storage.discontinuity).FlagHeads(segment_flags, keys, inequality_op);
// Zip values and segment_flags
ZipValuesAndFlags<IS_LAST_TILE>(num_remaining, values, segment_flags, scan_items);
// Exclusive scan of values and segment_flags
FlagValuePairT tile_aggregate;
ScanTile(scan_items, tile_aggregate, bool_constant_v<IS_INCLUSIVE>);
if (threadIdx.x == 0)
{
if (!IS_LAST_TILE)
{
tile_state.SetInclusive(0, tile_aggregate);
}
scan_items[0].key = 0;
}
}
else
{
KeyT tile_pred_key = (threadIdx.x == 0) ? d_keys_prev_in[tile_idx] : KeyT();
BlockDiscontinuityKeysT(storage.scan_storage.discontinuity)
.FlagHeads(segment_flags, keys, inequality_op, tile_pred_key);
// Zip values and segment_flags
ZipValuesAndFlags<IS_LAST_TILE>(num_remaining, values, segment_flags, scan_items);
FlagValuePairT tile_aggregate;
TilePrefixCallbackT prefix_op(tile_state, storage.scan_storage.prefix, pair_scan_op, tile_idx);
ScanTile(scan_items, tile_aggregate, prefix_op, bool_constant_v<IS_INCLUSIVE>);
}
__syncthreads();
UnzipValues(values, scan_items);
AddInitToScan(values, segment_flags);
// Store items
if (IS_LAST_TILE)
{
BlockStoreValuesT(storage.store_values).Store(d_values_out + tile_base, values, num_remaining);
}
else
{
BlockStoreValuesT(storage.store_values).Store(d_values_out + tile_base, values);
}
}
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
// Dequeue and scan tiles of items as part of a dynamic chained scan
// with Init functor
_CCCL_DEVICE _CCCL_FORCEINLINE AgentScanByKey(
TempStorage& storage,
KeysInputIteratorT d_keys_in,
KeyT* d_keys_prev_in,
ValuesInputIteratorT d_values_in,
ValuesOutputIteratorT d_values_out,
EqualityOp equality_op,
ScanOpT scan_op,
InitValueT init_value)
: storage(storage.Alias())
, d_keys_in(d_keys_in)
, d_keys_prev_in(d_keys_prev_in)
, d_values_in(d_values_in)
, d_values_out(d_values_out)
, inequality_op(equality_op)
, scan_op(scan_op)
, pair_scan_op(scan_op)
, init_value(init_value)
{}
/**
* Scan tiles of items as part of a dynamic chained scan
*
* @param num_items
* Total number of input items
*
* @param tile_state
* Global tile state descriptor
*
* start_tile
* The starting tile for the current grid
*/
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeRange(OffsetT num_items, ScanTileStateT& tile_state, int start_tile)
{
int tile_idx = static_cast<int>(blockIdx.x);
OffsetT tile_base = OffsetT(ITEMS_PER_TILE) * tile_idx;
OffsetT num_remaining = num_items - tile_base;
if (num_remaining > ITEMS_PER_TILE)
{
// Not the last tile (full)
ConsumeTile<false>(num_items, num_remaining, tile_idx, tile_base, tile_state);
}
else if (num_remaining > 0)
{
// The last tile (possibly partially-full)
ConsumeTile<true>(num_items, num_remaining, tile_idx, tile_base, tile_state);
}
}
};
} // namespace detail::scan_by_key
CUB_NAMESPACE_END

View File

@@ -0,0 +1,263 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/agent_radix_sort_downsweep.cuh>
#include <cub/agent/agent_radix_sort_upsweep.cuh>
#include <cub/block/block_radix_sort.cuh>
#include <cub/util_namespace.cuh>
#include <cub/util_type.cuh>
CUB_NAMESPACE_BEGIN
namespace detail::radix_sort
{
/**
* This agent will be implementing the `DeviceSegmentedRadixSort` when the
* https://github.com/NVIDIA/cub/issues/383 is addressed.
*
* @tparam IsDescending
* Whether or not the sorted-order is high-to-low
*
* @tparam SegmentedPolicyT
* Chained tuning policy
*
* @tparam KeyT
* Key type
*
* @tparam ValueT
* Value type
*
* @tparam OffsetT
* Signed integer type for global offsets
*/
template <bool IsDescending,
typename SegmentedPolicyT,
typename KeyT,
typename ValueT,
typename OffsetT,
typename DecomposerT = identity_decomposer_t>
struct AgentSegmentedRadixSort
{
OffsetT num_items;
static constexpr int ITEMS_PER_THREAD = SegmentedPolicyT::ITEMS_PER_THREAD;
static constexpr int BLOCK_THREADS = SegmentedPolicyT::BLOCK_THREADS;
static constexpr int RADIX_BITS = SegmentedPolicyT::RADIX_BITS;
static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;
static constexpr int KEYS_ONLY = ::cuda::std::is_same_v<ValueT, NullType>;
using traits = radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
// Huge segment handlers
using BlockUpsweepT = AgentRadixSortUpsweep<SegmentedPolicyT, KeyT, OffsetT, DecomposerT>;
using DigitScanT = BlockScan<OffsetT, BLOCK_THREADS>;
using BlockDownsweepT = AgentRadixSortDownsweep<SegmentedPolicyT, IsDescending, KeyT, ValueT, OffsetT, DecomposerT>;
/// Number of bin-starting offsets tracked per thread
static constexpr int BINS_TRACKED_PER_THREAD = BlockDownsweepT::BINS_TRACKED_PER_THREAD;
// Small segment handlers
using BlockRadixSortT =
BlockRadixSort<KeyT,
BLOCK_THREADS,
ITEMS_PER_THREAD,
ValueT,
RADIX_BITS,
(SegmentedPolicyT::RANK_ALGORITHM == RADIX_RANK_MEMOIZE),
SegmentedPolicyT::SCAN_ALGORITHM>;
using BlockKeyLoadT = BlockLoad<KeyT, BLOCK_THREADS, ITEMS_PER_THREAD, SegmentedPolicyT::LOAD_ALGORITHM>;
using BlockValueLoadT = BlockLoad<ValueT, BLOCK_THREADS, ITEMS_PER_THREAD, SegmentedPolicyT::LOAD_ALGORITHM>;
union _TempStorage
{
// Huge segment handlers
typename BlockUpsweepT::TempStorage upsweep;
typename BlockDownsweepT::TempStorage downsweep;
struct UnboundBlockSort
{
OffsetT reverse_counts_in[RADIX_DIGITS];
OffsetT reverse_counts_out[RADIX_DIGITS];
typename DigitScanT::TempStorage scan;
} unbound_sort;
// Small segment handlers
typename BlockKeyLoadT::TempStorage keys_load;
typename BlockValueLoadT::TempStorage values_load;
typename BlockRadixSortT::TempStorage sort;
};
using TempStorage = Uninitialized<_TempStorage>;
_TempStorage& temp_storage;
DecomposerT decomposer;
_CCCL_DEVICE _CCCL_FORCEINLINE
AgentSegmentedRadixSort(OffsetT num_items, TempStorage& temp_storage, DecomposerT decomposer = {})
: num_items(num_items)
, temp_storage(temp_storage.Alias())
, decomposer(decomposer)
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessSinglePass(
int begin_bit, int end_bit, const KeyT* d_keys_in, const ValueT* d_values_in, KeyT* d_keys_out, ValueT* d_values_out)
{
KeyT thread_keys[ITEMS_PER_THREAD];
ValueT thread_values[ITEMS_PER_THREAD];
// For FP64 the difference is:
// Lowest() -> -1.79769e+308 = 00...00b -> TwiddleIn -> -0 = 10...00b
// LOWEST -> -nan = 11...11b -> TwiddleIn -> 0 = 00...00b
bit_ordered_type default_key_bits =
IsDescending ? traits::min_raw_binary_key(decomposer) : traits::max_raw_binary_key(decomposer);
KeyT oob_default = reinterpret_cast<KeyT&>(default_key_bits);
if (!KEYS_ONLY)
{
BlockValueLoadT(temp_storage.values_load).Load(d_values_in, thread_values, num_items);
__syncthreads();
}
{
BlockKeyLoadT(temp_storage.keys_load).Load(d_keys_in, thread_keys, num_items, oob_default);
__syncthreads();
}
BlockRadixSortT(temp_storage.sort)
.SortBlockedToStriped(
thread_keys,
thread_values,
begin_bit,
end_bit,
bool_constant_v<IsDescending>,
bool_constant_v<KEYS_ONLY>,
decomposer);
cub::StoreDirectStriped<BLOCK_THREADS>(threadIdx.x, d_keys_out, thread_keys, num_items);
if (!KEYS_ONLY)
{
cub::StoreDirectStriped<BLOCK_THREADS>(threadIdx.x, d_values_out, thread_values, num_items);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessIterative(
int current_bit,
int pass_bits,
const KeyT* d_keys_in,
const ValueT* d_values_in,
KeyT* d_keys_out,
ValueT* d_values_out)
{
// Upsweep
BlockUpsweepT upsweep(temp_storage.upsweep, d_keys_in, current_bit, pass_bits, decomposer);
upsweep.ProcessRegion(OffsetT{}, num_items);
__syncthreads();
// The count of each digit value in this pass (valid in the first RADIX_DIGITS threads)
OffsetT bin_count[BINS_TRACKED_PER_THREAD];
upsweep.ExtractCounts(bin_count);
__syncthreads();
if (IsDescending)
{
// Reverse bin counts
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
temp_storage.unbound_sort.reverse_counts_in[bin_idx] = bin_count[track];
}
}
__syncthreads();
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
bin_count[track] = temp_storage.unbound_sort.reverse_counts_in[RADIX_DIGITS - bin_idx - 1];
}
}
}
// Scan
// The global scatter base offset for each digit value in this pass
// (valid in the first RADIX_DIGITS threads)
OffsetT bin_offset[BINS_TRACKED_PER_THREAD];
DigitScanT(temp_storage.unbound_sort.scan).ExclusiveSum(bin_count, bin_offset);
if (IsDescending)
{
// Reverse bin offsets
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
temp_storage.unbound_sort.reverse_counts_out[threadIdx.x] = bin_offset[track];
}
}
__syncthreads();
_CCCL_PRAGMA_UNROLL_FULL()
for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)
{
int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;
if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))
{
bin_offset[track] = temp_storage.unbound_sort.reverse_counts_out[RADIX_DIGITS - bin_idx - 1];
}
}
}
__syncthreads();
// Downsweep
BlockDownsweepT downsweep(
temp_storage.downsweep,
bin_offset,
num_items,
d_keys_in,
d_keys_out,
d_values_in,
d_values_out,
current_bit,
pass_bits,
decomposer);
downsweep.ProcessRegion(OffsetT{}, num_items);
}
};
} // namespace detail::radix_sort
CUB_NAMESPACE_END

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,349 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_type.cuh>
#include <cub/warp/warp_load.cuh>
#include <cub/warp/warp_merge_sort.cuh>
#include <cub/warp/warp_store.cuh>
#include <nv/target>
CUB_NAMESPACE_BEGIN
namespace detail
{
// TODO(bgruber): drop in CCCL 4.0
template <int ThreadsPerBlock,
int WarpThreadsArg,
int ItemsPerThreadArg,
cub::WarpLoadAlgorithm LoadAlgorithmArg = cub::WARP_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifierArg = cub::LOAD_LDG,
cub::WarpStoreAlgorithm StoreAlgorithmArg = cub::WARP_STORE_DIRECT>
struct agent_sub_warp_merge_sort_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int WARP_THREADS = WarpThreadsArg;
static constexpr int ITEMS_PER_THREAD = ItemsPerThreadArg;
static constexpr int ITEMS_PER_TILE = WARP_THREADS * ITEMS_PER_THREAD;
static constexpr int SEGMENTS_PER_BLOCK = BLOCK_THREADS / WARP_THREADS;
static constexpr cub::WarpLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithmArg;
static constexpr cub::CacheLoadModifier LOAD_MODIFIER = LoadModifierArg;
static constexpr cub::WarpStoreAlgorithm STORE_ALGORITHM = StoreAlgorithmArg;
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int WarpThreadsArg,
int ItemsPerThreadArg,
cub::WarpLoadAlgorithm LoadAlgorithmArg = cub::WARP_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifierArg = cub::LOAD_LDG,
cub::WarpStoreAlgorithm StoreAlgorithmArg = cub::WARP_STORE_DIRECT>
using AgentSubWarpMergeSortPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceSegmentedSort") = detail::agent_sub_warp_merge_sort_policy<
ThreadsPerBlock,
WarpThreadsArg,
ItemsPerThreadArg,
LoadAlgorithmArg,
LoadModifierArg,
StoreAlgorithmArg>;
namespace detail::sub_warp_merge_sort
{
/**
* @brief AgentSubWarpSort implements a sub-warp merge sort.
*
* This agent can work with any power of two number of threads, not exceeding
* 32. The number of threads is defined in the `PolicyT::WARP_THREADS`. Virtual
* warp of `PolicyT::WARP_THREADS` will efficiently load data using
* `PolicyT::LOAD_ALGORITHM`, sort it using `WarpMergeSort`, and store it back
* using `PolicyT::STORE_ALGORITHM`.
*
* @tparam IS_DESCENDING
* Whether or not the sorted-order is high-to-low
*
* @tparam PolicyT
* Chained tuning policy
*
* @tparam KeyT
* Key type
*
* @tparam ValueT
* Value type
*
* @tparam OffsetT
* Signed integer type for global offsets
*/
template <bool IS_DESCENDING, typename PolicyT, typename KeyT, typename ValueT, typename OffsetT>
class AgentSubWarpSort
{
using traits = detail::radix::traits_t<KeyT>;
using bit_ordered_type = typename traits::bit_ordered_type;
struct BinaryOpT
{
template <typename T>
_CCCL_DEVICE bool operator()(T lhs, T rhs) const noexcept
{
if constexpr (IS_DESCENDING)
{
return lhs > rhs;
}
else
{
return lhs < rhs;
}
_CCCL_UNREACHABLE();
}
#if _CCCL_HAS_NVFP16()
_CCCL_DEVICE bool operator()(__half lhs, __half rhs) const noexcept
{
// Need to explicitly cast to float for SM <= 52.
if constexpr (IS_DESCENDING)
{
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_53, (return __hgt(lhs, rhs);), (return __half2float(lhs) > __half2float(rhs);));
}
else
{
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_53, (return __hlt(lhs, rhs);), (return __half2float(lhs) < __half2float(rhs);));
}
_CCCL_UNREACHABLE();
}
#endif // _CCCL_HAS_NVFP16()
#if _CCCL_HAS_NVBF16()
_CCCL_DEVICE bool operator()(__nv_bfloat16 lhs, __nv_bfloat16 rhs) const noexcept
{
// Need to explicitly cast to float for SM < 80.
if constexpr (IS_DESCENDING)
{
NV_IF_ELSE_TARGET(
NV_PROVIDES_SM_80, (return __hgt(lhs, rhs);), (return __bfloat162float(lhs) > __bfloat162float(rhs);));
}
else
{
NV_IF_ELSE_TARGET(
NV_PROVIDES_SM_80, (return __hlt(lhs, rhs);), (return __bfloat162float(lhs) < __bfloat162float(rhs);));
}
_CCCL_UNREACHABLE();
}
#endif // _CCCL_HAS_NVBF16()
};
#if _CCCL_HAS_NVFP16()
_CCCL_DEVICE static bool equal(__half lhs, __half rhs)
{
// Need to explicitly cast to float for SM <= 52.
NV_IF_ELSE_TARGET(NV_PROVIDES_SM_53, (return __heq(lhs, rhs);), (return __half2float(lhs) == __half2float(rhs);));
}
#endif // _CCCL_HAS_NVFP16()
#if _CCCL_HAS_NVBF16()
_CCCL_DEVICE static bool equal(__nv_bfloat16 lhs, __nv_bfloat16 rhs)
{
// Need to explicitly cast to float for SM < 80.
NV_IF_ELSE_TARGET(
NV_PROVIDES_SM_80, (return __heq(lhs, rhs);), (return __bfloat162float(lhs) == __bfloat162float(rhs);));
}
#endif // _CCCL_HAS_NVBF16()
template <typename T>
_CCCL_DEVICE static bool equal(T lhs, T rhs)
{
return lhs == rhs;
}
public:
static constexpr bool KEYS_ONLY = ::cuda::std::is_same_v<ValueT, cub::NullType>;
using WarpMergeSortT = WarpMergeSort<KeyT, PolicyT::ITEMS_PER_THREAD, PolicyT::WARP_THREADS, ValueT>;
using KeysLoadItT = try_make_cache_modified_iterator_t<PolicyT::LOAD_MODIFIER, const KeyT*>;
using ItemsLoadItT = try_make_cache_modified_iterator_t<PolicyT::LOAD_MODIFIER, const ValueT*>;
using WarpLoadKeysT = cub::WarpLoad<KeyT, PolicyT::ITEMS_PER_THREAD, PolicyT::LOAD_ALGORITHM, PolicyT::WARP_THREADS>;
using WarpLoadItemsT =
cub::WarpLoad<ValueT, PolicyT::ITEMS_PER_THREAD, PolicyT::LOAD_ALGORITHM, PolicyT::WARP_THREADS>;
using WarpStoreKeysT =
cub::WarpStore<KeyT, PolicyT::ITEMS_PER_THREAD, PolicyT::STORE_ALGORITHM, PolicyT::WARP_THREADS>;
using WarpStoreItemsT =
cub::WarpStore<ValueT, PolicyT::ITEMS_PER_THREAD, PolicyT::STORE_ALGORITHM, PolicyT::WARP_THREADS>;
union _TempStorage
{
typename WarpLoadKeysT::TempStorage load_keys;
typename WarpLoadItemsT::TempStorage load_items;
typename WarpMergeSortT::TempStorage sort;
typename WarpStoreKeysT::TempStorage store_keys;
typename WarpStoreItemsT::TempStorage store_items;
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
_TempStorage& storage;
_CCCL_DEVICE _CCCL_FORCEINLINE explicit AgentSubWarpSort(TempStorage& temp_storage)
: storage(temp_storage.Alias())
{}
_CCCL_DEVICE _CCCL_FORCEINLINE void ProcessSegment(
int segment_size, KeysLoadItT keys_input, KeyT* keys_output, ItemsLoadItT values_input, ValueT* values_output)
{
WarpMergeSortT warp_merge_sort(storage.sort);
if (segment_size < 3)
{
ShortCircuit(
warp_merge_sort.get_linear_tid(),
segment_size,
keys_input,
keys_output,
values_input,
values_output,
BinaryOpT{});
}
else
{
KeyT keys[PolicyT::ITEMS_PER_THREAD];
ValueT values[PolicyT::ITEMS_PER_THREAD];
KeyT oob_default = [&] {
if constexpr (::cuda::std::is_same_v<bool, KeyT>)
{
// Traits<KeyT>::MAX_KEY for `bool` is 0xFF which is different from `true` and makes
// comparison with oob unreliable.
return !IS_DESCENDING;
}
else
{
// For FP64 the difference is:
// Lowest() -> -1.79769e+308 = 00...00b -> TwiddleIn -> -0 = 10...00b
// LOWEST -> -nan = 11...11b -> TwiddleIn -> 0 = 00...00b
// Segmented sort doesn't support custom types at the moment.
bit_ordered_type default_key_bits = IS_DESCENDING ? traits::min_raw_binary_key(identity_decomposer_t{})
: traits::max_raw_binary_key(identity_decomposer_t{});
return reinterpret_cast<KeyT&>(default_key_bits);
}
}();
WarpLoadKeysT(storage.load_keys).Load(keys_input, keys, segment_size, oob_default);
__syncwarp(warp_merge_sort.get_member_mask());
if (!KEYS_ONLY)
{
WarpLoadItemsT(storage.load_items).Load(values_input, values, segment_size);
__syncwarp(warp_merge_sort.get_member_mask());
}
warp_merge_sort.Sort(keys, values, BinaryOpT{}, segment_size, oob_default);
__syncwarp(warp_merge_sort.get_member_mask());
WarpStoreKeysT(storage.store_keys).Store(keys_output, keys, segment_size);
if (!KEYS_ONLY)
{
__syncwarp(warp_merge_sort.get_member_mask());
WarpStoreItemsT(storage.store_items).Store(values_output, values, segment_size);
}
}
}
private:
/**
* This method implements a shortcut for sorting less than three items.
* Only the first thread of a virtual warp is used for soring.
*/
template <typename CompareOpT>
_CCCL_DEVICE _CCCL_FORCEINLINE void ShortCircuit(
unsigned int linear_tid,
OffsetT segment_size,
KeysLoadItT keys_input,
KeyT* keys_output,
ItemsLoadItT values_input,
ValueT* values_output,
CompareOpT binary_op)
{
if (segment_size == 1)
{
if (linear_tid == 0)
{
if (keys_input.ptr != keys_output)
{
keys_output[0] = keys_input[0];
}
if (!KEYS_ONLY)
{
if (values_input.ptr != values_output)
{
values_output[0] = values_input[0];
}
}
}
}
else if (segment_size == 2)
{
if (linear_tid == 0)
{
KeyT lhs = keys_input[0];
KeyT rhs = keys_input[1];
if (equal(lhs, rhs) || binary_op(lhs, rhs))
{
keys_output[0] = lhs;
keys_output[1] = rhs;
if (!KEYS_ONLY)
{
if (values_output != values_input.ptr)
{
values_output[0] = values_input[0];
values_output[1] = values_input[1];
}
}
}
else
{
keys_output[0] = rhs;
keys_output[1] = lhs;
if (!KEYS_ONLY)
{
// values_output might be an alias for values_input, so
// we have to use registers here
const ValueT lhs_val = values_input[0];
const ValueT rhs_val = values_input[1];
values_output[0] = rhs_val;
values_output[1] = lhs_val;
}
}
}
}
}
};
} // namespace detail::sub_warp_merge_sort
CUB_NAMESPACE_END

View File

@@ -0,0 +1,588 @@
// SPDX-FileCopyrightText: Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/single_pass_scan_operators.cuh>
#include <cub/block/block_discontinuity.cuh>
#include <cub/block/block_exchange.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/iterator/cache_modified_input_iterator.cuh>
#include <cub/util_device.cuh>
#include <cuda/std/__functional/operations.h>
#include <cuda/std/__type_traits/conditional.h>
#include <cuda/std/__type_traits/enable_if.h>
#include <cuda/std/__type_traits/is_pointer.h>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
// TODO(bgruber): remove this when C++20 is the minimum, since then we can pass policy values as NTTP
template <int ThreadsPerBlock,
int ItemsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockScanAlgorithm ScanAlgorithm,
class DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
struct agent_three_way_partition_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
struct detail
{
using delay_constructor_t = DelayConstructorT;
};
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int ItemsPerThread,
BlockLoadAlgorithm LoadAlgorithm,
CacheLoadModifier LoadModifier,
BlockScanAlgorithm ScanAlgorithm,
class DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
using AgentThreeWayPartitionPolicy
CCCL_DEPRECATED_BECAUSE("Use the tuning API for DevicePartition") = detail::agent_three_way_partition_policy<
ThreadsPerBlock,
ItemsPerThread,
LoadAlgorithm,
LoadModifier,
ScanAlgorithm,
DelayConstructorT>;
namespace detail::three_way_partition
{
template <class OffsetT>
struct pair_pack_t
{
OffsetT x, y;
_CCCL_DEVICE pair_pack_t<OffsetT> operator+(const pair_pack_t<OffsetT>& other) const
{
return {x + other.x, y + other.y};
}
};
template <class OffsetT, class = void>
struct accumulator_pack_base_t
{
using pack_t = pair_pack_t<OffsetT>;
_CCCL_DEVICE static pack_t pack(OffsetT f, OffsetT s)
{
return {f, s};
}
_CCCL_DEVICE static OffsetT first(pack_t packed)
{
return packed.x;
}
_CCCL_DEVICE static OffsetT second(pack_t packed)
{
return packed.y;
}
};
template <class OffsetT>
struct accumulator_pack_base_t<OffsetT, ::cuda::std::enable_if_t<sizeof(OffsetT) == 4>>
{
using pack_t = uint64_t;
_CCCL_DEVICE static pack_t pack(OffsetT f, OffsetT s)
{
return (static_cast<pack_t>(f) << 32) | static_cast<pack_t>(s);
}
_CCCL_DEVICE static OffsetT first(pack_t packed)
{
return static_cast<OffsetT>(packed >> 32);
}
_CCCL_DEVICE static OffsetT second(pack_t packed)
{
return static_cast<OffsetT>(packed & 0xFFFFFFFF);
}
};
template <class OffsetT>
struct accumulator_pack_t : accumulator_pack_base_t<OffsetT>
{
using base = accumulator_pack_base_t<OffsetT>;
using typename base::pack_t;
_CCCL_DEVICE static void subtract(pack_t& packed, OffsetT val)
{
packed = base::pack(base::first(packed) - val, base::second(packed) - val);
}
_CCCL_DEVICE static OffsetT sum(pack_t& packed)
{
return base::first(packed) + base::second(packed);
}
_CCCL_DEVICE static pack_t zero()
{
return {};
}
};
/**
* \brief Implements a device-wide three-way partitioning
*
* Splits input data into three parts based on the selection functors. If the
* first functor selects an item, the algorithm places it in the first part.
* Otherwise, if the second functor selects an item, the algorithm places it in
* the second part. If both functors don't select an item, the algorithm places
* it into the unselected part.
*/
template <typename PolicyT,
typename InputIteratorT,
typename FirstOutputIteratorT,
typename SecondOutputIteratorT,
typename UnselectedOutputIteratorT,
typename SelectFirstPartOp,
typename SelectSecondPartOp,
typename OffsetT,
typename StreamingContextT>
struct AgentThreeWayPartition
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
// The input value type
using InputT = it_value_t<InputIteratorT>;
using AccumPackHelperT = accumulator_pack_t<OffsetT>;
using AccumPackT = typename AccumPackHelperT::pack_t;
// Tile status descriptor interface type
using ScanTileStateT = cub::ScanTileState<AccumPackT>;
// Constants
static constexpr int BLOCK_THREADS = PolicyT::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = PolicyT::ITEMS_PER_THREAD;
static constexpr int TILE_ITEMS = BLOCK_THREADS * ITEMS_PER_THREAD;
using WrappedInputIteratorT =
::cuda::std::_If<::cuda::std::is_pointer_v<InputIteratorT>,
cub::CacheModifiedInputIterator<PolicyT::LOAD_MODIFIER, InputT, OffsetT>,
InputIteratorT>;
// Parameterized BlockLoad type for input data
using BlockLoadT = cub::BlockLoad<InputT, BLOCK_THREADS, ITEMS_PER_THREAD, PolicyT::LOAD_ALGORITHM>;
// Parameterized BlockScan type
using BlockScanT = cub::BlockScan<AccumPackT, BLOCK_THREADS, PolicyT::SCAN_ALGORITHM>;
// Callback type for obtaining tile prefix during block scan
using DelayConstructorT = typename PolicyT::detail::delay_constructor_t;
using TilePrefixCallbackOpT =
cub::TilePrefixCallbackOp<AccumPackT, ::cuda::std::plus<>, ScanTileStateT, DelayConstructorT>;
// Item exchange type
using ItemExchangeT = InputT[TILE_ITEMS];
// Shared memory type for this thread block
union _TempStorage
{
struct ScanStorage
{
// Smem needed for tile scanning
typename BlockScanT::TempStorage scan;
// Smem needed for cooperative prefix callback
typename TilePrefixCallbackOpT::TempStorage prefix;
} scan_storage;
// Smem needed for loading items
typename BlockLoadT::TempStorage load_items;
// Smem needed for compacting items (allows non POD items in this union)
cub::Uninitialized<ItemExchangeT> raw_exchange;
};
// Alias wrapper allowing storage to be unioned
struct TempStorage : cub::Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
_TempStorage& temp_storage; ///< Reference to temp_storage
WrappedInputIteratorT d_in; ///< Input items
FirstOutputIteratorT d_first_part_out;
SecondOutputIteratorT d_second_part_out;
UnselectedOutputIteratorT d_unselected_out;
SelectFirstPartOp select_first_part_op;
SelectSecondPartOp select_second_part_op;
OffsetT num_items; ///< Total number of input items
// Note: This is a const reference because we have seen double-digit percentage perf regressions otherwise
const StreamingContextT& streaming_context; ///< Context for the current partition
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
// Constructor
_CCCL_DEVICE _CCCL_FORCEINLINE AgentThreeWayPartition(
TempStorage& temp_storage,
InputIteratorT d_in,
FirstOutputIteratorT d_first_part_out,
SecondOutputIteratorT d_second_part_out,
UnselectedOutputIteratorT d_unselected_out,
SelectFirstPartOp select_first_part_op,
SelectSecondPartOp select_second_part_op,
OffsetT num_items,
const StreamingContextT& streaming_context)
: temp_storage(temp_storage.Alias())
, d_in(d_in)
, d_first_part_out(d_first_part_out)
, d_second_part_out(d_second_part_out)
, d_unselected_out(d_unselected_out)
, select_first_part_op(select_first_part_op)
, select_second_part_op(select_second_part_op)
, num_items(num_items)
, streaming_context(streaming_context)
{}
//---------------------------------------------------------------------
// Utility methods for initializing the selections
//---------------------------------------------------------------------
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void Initialize(
OffsetT num_tile_items, InputT (&items)[ITEMS_PER_THREAD], AccumPackT (&items_selection_flags)[ITEMS_PER_THREAD])
{
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
// Out-of-bounds items are selection_flags
items_selection_flags[ITEM] = AccumPackHelperT::pack(1, 1);
if (!IS_LAST_TILE || (OffsetT(threadIdx.x * ITEMS_PER_THREAD) + ITEM < num_tile_items))
{
OffsetT first_item_selected = select_first_part_op(items[ITEM]);
items_selection_flags[ITEM] =
AccumPackHelperT::pack(first_item_selected, first_item_selected ? 0 : select_second_part_op(items[ITEM]));
}
}
}
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void Scatter(
InputT (&items)[ITEMS_PER_THREAD],
AccumPackT (&items_selection_flags)[ITEMS_PER_THREAD],
AccumPackT (&items_selection_indices)[ITEMS_PER_THREAD],
int num_tile_items,
AccumPackT num_tile_selected,
AccumPackT num_tile_selected_prefix,
OffsetT num_rejected_prefix)
{
__syncthreads();
const OffsetT num_first_selections_prefix = AccumPackHelperT::first(num_tile_selected_prefix);
const OffsetT num_second_selections_prefix = AccumPackHelperT::second(num_tile_selected_prefix);
const int first_item_end = AccumPackHelperT::first(num_tile_selected);
const int second_item_end = first_item_end + AccumPackHelperT::second(num_tile_selected);
// Scatter items to shared memory (rejections first)
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
int item_idx = (threadIdx.x * ITEMS_PER_THREAD) + ITEM;
const OffsetT first_items_selection_indices = AccumPackHelperT::first(items_selection_indices[ITEM]);
const OffsetT second_items_selection_indices = AccumPackHelperT::second(items_selection_indices[ITEM]);
if (!IS_LAST_TILE || (item_idx < num_tile_items))
{
int local_scatter_offset = 0;
if (AccumPackHelperT::first(items_selection_flags[ITEM]))
{
local_scatter_offset = first_items_selection_indices - num_first_selections_prefix;
}
else if (AccumPackHelperT::second(items_selection_flags[ITEM]))
{
local_scatter_offset = first_item_end + second_items_selection_indices - num_second_selections_prefix;
}
else
{
// Medium item
int local_selection_idx = (first_items_selection_indices - num_first_selections_prefix)
+ (second_items_selection_indices - num_second_selections_prefix);
local_scatter_offset = second_item_end + item_idx - local_selection_idx;
}
temp_storage.raw_exchange.Alias()[local_scatter_offset] = items[ITEM];
}
}
__syncthreads();
// Gather items from shared memory and scatter to global
// NOLINTBEGIN(bugprone-misplaced-widening-cast)
auto first_base =
d_first_part_out + (streaming_context.num_previously_selected_first() + num_first_selections_prefix);
auto second_base =
d_second_part_out + (streaming_context.num_previously_selected_second() + num_second_selections_prefix);
auto unselected_base = d_unselected_out + (streaming_context.num_previously_rejected() + num_rejected_prefix);
// NOLINTEND(bugprone-misplaced-widening-cast)
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
int item_idx = (ITEM * BLOCK_THREADS) + threadIdx.x;
if (!IS_LAST_TILE || (item_idx < num_tile_items))
{
InputT item = temp_storage.raw_exchange.Alias()[item_idx];
if (item_idx < first_item_end)
{
first_base[item_idx] = item;
}
else if (item_idx < second_item_end)
{
second_base[item_idx - first_item_end] = item;
}
else
{
int rejection_idx = item_idx - second_item_end;
unselected_base[rejection_idx] = item;
}
}
}
}
//---------------------------------------------------------------------
// Cooperatively scan a device-wide sequence of tiles with other CTAs
//---------------------------------------------------------------------
/**
* Process first tile of input (dynamic chained scan).
* Returns the running count of selections (including this tile)
*
* @param num_tile_items Number of input items comprising this tile
* @param tile_offset Tile offset
* @param first_tile_state Global tile state descriptor
* @param second_tile_state Global tile state descriptor
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeFirstTile(int num_tile_items, OffsetT tile_offset, ScanTileStateT& tile_state, AccumPackT& num_items_selected)
{
InputT items[ITEMS_PER_THREAD];
AccumPackT items_selection_flags[ITEMS_PER_THREAD];
AccumPackT items_selection_indices[ITEMS_PER_THREAD];
// Load items
if constexpr (IS_LAST_TILE)
{
BlockLoadT(temp_storage.load_items)
.Load(d_in + streaming_context.input_offset() + tile_offset, items, num_tile_items);
}
else
{
BlockLoadT(temp_storage.load_items).Load(d_in + streaming_context.input_offset() + tile_offset, items);
}
// Initialize selection_flags
Initialize<IS_LAST_TILE>(num_tile_items, items, items_selection_flags);
__syncthreads();
// Exclusive scan of selection_flags
BlockScanT(temp_storage.scan_storage.scan)
.ExclusiveSum(items_selection_flags, items_selection_indices, num_items_selected);
if (threadIdx.x == 0)
{
// Update tile status if this is not the last tile
if (!IS_LAST_TILE)
{
tile_state.SetInclusive(0, num_items_selected);
}
}
// Discount any out-of-bounds selections
if (IS_LAST_TILE)
{
AccumPackHelperT::subtract(num_items_selected, TILE_ITEMS - num_tile_items);
}
// Scatter flagged items
Scatter<IS_LAST_TILE>(
items,
items_selection_flags,
items_selection_indices,
num_tile_items,
num_items_selected,
// all the prefixes equal to 0 because it's the first tile
AccumPackHelperT::zero(),
0);
}
/**
* Process subsequent tile of input (dynamic chained scan).
* Returns the running count of selections (including this tile)
*
* @param num_tile_items Number of input items comprising this tile
* @param tile_idx Tile index
* @param tile_offset Tile offset
* @param first_tile_state Global tile state descriptor
* @param second_tile_state Global tile state descriptor
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void ConsumeSubsequentTile(
int num_tile_items, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state, AccumPackT& num_items_selected)
{
InputT items[ITEMS_PER_THREAD];
AccumPackT items_selected_flags[ITEMS_PER_THREAD];
AccumPackT items_selected_indices[ITEMS_PER_THREAD];
// Load items
if (IS_LAST_TILE)
{
BlockLoadT(temp_storage.load_items)
.Load(d_in + streaming_context.input_offset() + tile_offset, items, num_tile_items);
}
else
{
BlockLoadT(temp_storage.load_items).Load(d_in + streaming_context.input_offset() + tile_offset, items);
}
// Initialize selection_flags
Initialize<IS_LAST_TILE>(num_tile_items, items, items_selected_flags);
__syncthreads();
// Exclusive scan of values and selection_flags
TilePrefixCallbackOpT prefix_op(tile_state, temp_storage.scan_storage.prefix, ::cuda::std::plus<>{}, tile_idx);
BlockScanT(temp_storage.scan_storage.scan).ExclusiveSum(items_selected_flags, items_selected_indices, prefix_op);
num_items_selected = prefix_op.GetInclusivePrefix();
AccumPackT num_items_in_tile_selected = prefix_op.GetBlockAggregate();
AccumPackT num_items_selected_prefix = prefix_op.GetExclusivePrefix();
__syncthreads();
OffsetT num_rejected_prefix = (tile_idx * TILE_ITEMS) - AccumPackHelperT::sum(num_items_selected_prefix);
// Discount any out-of-bounds selections. There are exactly
// TILE_ITEMS - num_tile_items elements like that because we
// marked them as selected in Initialize method.
if (IS_LAST_TILE)
{
const int num_discount = TILE_ITEMS - num_tile_items;
AccumPackHelperT::subtract(num_items_selected, num_discount);
AccumPackHelperT::subtract(num_items_in_tile_selected, num_discount);
}
// Scatter flagged items
Scatter<IS_LAST_TILE>(
items,
items_selected_flags,
items_selected_indices,
num_tile_items,
num_items_in_tile_selected,
num_items_selected_prefix,
num_rejected_prefix);
}
/**
* Process a tile of input
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeTile(int num_tile_items, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state, AccumPackT& accum)
{
if (tile_idx == 0)
{
ConsumeFirstTile<IS_LAST_TILE>(num_tile_items, tile_offset, tile_state, accum);
}
else
{
ConsumeSubsequentTile<IS_LAST_TILE>(num_tile_items, tile_idx, tile_offset, tile_state, accum);
}
}
/**
* Scan tiles of items as part of a dynamic chained scan
*
* @tparam NumSelectedIteratorT
* Output iterator type for recording number of items selection_flags
*
* @param num_tiles
* Total number of input tiles
*
* @param first_tile_state
* Global tile state descriptor
*
* @param second_tile_state
* Global tile state descriptor
*
* @param d_num_selected_out
* Output total number selection_flags
*/
template <typename NumSelectedIteratorT>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeRange(int num_tiles, ScanTileStateT& tile_state, NumSelectedIteratorT d_num_selected_out)
{
// Blocks are launched in increasing order, so just assign one tile per block
// Current tile index
const int tile_idx = static_cast<int>(blockIdx.x);
// Global offset for the current tile
const OffsetT tile_offset = tile_idx * TILE_ITEMS;
AccumPackT accum;
if (tile_idx < num_tiles - 1)
{
// Not the last tile (full)
ConsumeTile<false>(TILE_ITEMS, tile_idx, tile_offset, tile_state, accum);
}
else
{
// The last tile (possibly partially-full)
const OffsetT num_remaining = num_items - tile_offset;
ConsumeTile<true>(num_remaining, tile_idx, tile_offset, tile_state, accum);
if (threadIdx.x == 0)
{
// Update the number of selected items with this partition's selections
streaming_context.update_num_selected(
d_num_selected_out, AccumPackHelperT::first(accum), AccumPackHelperT::second(accum), num_items);
}
}
}
};
} // namespace detail::three_way_partition
CUB_NAMESPACE_END

View File

@@ -0,0 +1,835 @@
// SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//! @file
//! cub::AgentTopK implements a stateful abstraction of CUDA thread blocks for participating in device-wide topK.
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/block/block_store.cuh>
#include <cub/block/radix_rank_sort_operations.cuh>
#include <cub/util_type.cuh>
#include <cuda/__cmath/ceil_div.h>
CUB_NAMESPACE_BEGIN
namespace detail::topk
{
//! @brief Parameterizable tuning policy type for agent_topk
//!
//! @tparam ThreadsPerBlock
//! Threads per thread block
//!
//! @tparam ItemsPerThread
//! Items per thread (per tile of input)
//!
//! @tparam BitsPerPass
//! Number of bits processed per pass
//!
//! @tparam LoadAlgorithm
//! The BlockLoad algorithm to use
//!
//! @tparam ScanAlgorithm
//! The BlockScan algorithm to use
//!
template <int ThreadsPerBlock,
int ItemsPerThread,
int BitsPerPass,
BlockLoadAlgorithm LoadAlgorithm,
BlockScanAlgorithm ScanAlgorithm>
struct agent_topk_policy
{
static constexpr int threads_per_block = ThreadsPerBlock;
static constexpr int items_per_thread = ItemsPerThread;
static constexpr int bits_per_pass = BitsPerPass;
static constexpr BlockLoadAlgorithm load_algorithm = LoadAlgorithm;
static constexpr BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
};
template <typename KeyT, bool CanTwiddle = detail::radix::can_twiddle<KeyT>>
struct key_prefix_storage_t;
template <typename KeyT>
struct key_prefix_storage_t<KeyT, true>
{
using bits_t = typename Traits<KeyT>::UnsignedBits;
bits_t bits;
};
// Calculates the number of passes needed for a type T with BitsPerPass bits processed per pass.
template <typename T>
[[nodiscard]] _CCCL_HOST_DEVICE _CCCL_FORCEINLINE constexpr int calc_num_passes(int bits_per_pass)
{
return ::cuda::ceil_div<int>(sizeof(T) * 8, bits_per_pass);
}
template <int BitsPerPass>
[[nodiscard]] _CCCL_HOST_DEVICE _CCCL_FORCEINLINE int calc_num_passes(const int total_bits)
{
return ::cuda::ceil_div<int>(total_bits, BitsPerPass);
}
// Calculates the starting bit for a given pass (bit 0 is the least significant (rightmost) bit).
// We process the input from the most to the least significant bit. This way, we can skip some passes in the end.
template <typename T, int BitsPerPass>
[[nodiscard]] _CCCL_HOST_DEVICE _CCCL_FORCEINLINE constexpr int calc_start_bit(const int pass)
{
int start_bit = int{sizeof(T)} * 8 - (pass + 1) * BitsPerPass;
if (start_bit < 0)
{
start_bit = 0;
}
return start_bit;
}
template <int BitsPerPass>
[[nodiscard]] _CCCL_HOST_DEVICE _CCCL_FORCEINLINE int calc_start_bit(const int total_bits, const int pass)
{
int start_bit = total_bits - (pass + 1) * BitsPerPass;
if (start_bit < 0)
{
start_bit = 0;
}
return start_bit;
}
// Bit-vector for accumulating prefix digits via funnel shift. Each pass shifts the existing
// contents left by BitsPerPass and ORs the new bucket at the bottom. Sized to hold all
// decomposed bits of KeyT plus headroom for the shift padding of the last pass.
template <typename KeyT>
struct key_prefix_storage_t<KeyT, false>
{
static constexpr int num_words = ::cuda::ceil_div<int>(sizeof(KeyT) * 8 + 31, 32);
unsigned int words[num_words];
// Funnel-shifts the entire bit-vector left by `shift` positions and inserts `value` into the
// vacated low bits. Each word receives carry bits from its lower neighbor (high-to-low order
// so each word reads its neighbor's original value). The final word is filled from `value`.
_CCCL_DEVICE _CCCL_FORCEINLINE void shift_or(int shift, unsigned int value)
{
_CCCL_ASSERT(shift > 0 && shift < 32, "shift_or requires 0 < shift < 32");
_CCCL_PRAGMA_UNROLL_FULL()
for (int i = num_words - 1; i > 0; --i)
{
words[i] = __funnelshift_l(words[i - 1], words[i], shift);
}
words[0] = (words[0] << shift) | value;
}
};
template <typename KeyT, int BitsPerPass>
_CCCL_DEVICE _CCCL_FORCEINLINE void
set_kth_key_bits(key_prefix_storage_t<KeyT>& prefix, const int pass, const int bin_index)
{
if constexpr (detail::radix::can_twiddle<KeyT>)
{
using bits_t = typename Traits<KeyT>::UnsignedBits;
const int start_bit = calc_start_bit<KeyT, BitsPerPass>(pass);
bits_t bucket = bin_index;
prefix.bits |= static_cast<bits_t>(bucket) << start_bit;
}
else
{
prefix.shift_or(BitsPerPass, bin_index);
}
}
template <typename KeyInT, typename OffsetT, typename OutOffsetT>
struct alignas(128) Counter
{
// We are processing the items in multiple passes, from most-significant to least-significant bits. In each pass, we
// keep the length of input (`len`) and the `k` of current pass, and update them at the end of the pass.
OutOffsetT k;
OffsetT len;
// `previous_len` is the length of the input in the previous pass. Note that `previous_len` rather than `len` is used
// for the filtering step because filtering is indeed for previous pass.
OffsetT previous_len;
// We determine the bits of the k_th key inside the mask processed by the pass. The
// already known bits are stored in `kth_key_bits`. It's used to discriminate a
// element is a result (written to `out`), a candidate for next pass (written to
// `out_buf`), or not useful (discarded). The bits that are not yet processed do not
// matter for this purpose.
key_prefix_storage_t<KeyInT> kth_key_bits;
// Record how many elements have passed filtering. It's used to determine the position
// in the `out_buf` where an element should be written.
alignas(128) OffsetT filter_cnt;
// For a row inside a batch, we may launch multiple thread blocks. This counter is
// used to determine if the current block is the last running block. If so, this block
// will execute compute_bin_offsets() and choose_bucket().
alignas(128) unsigned int finished_block_cnt;
// Record how many elements have been written to the front of `out`. Elements less (if
// SelectMin==true) than the k-th key are written from front to back.
alignas(128) OutOffsetT out_cnt;
// Record how many elements have been written to the back of `out`. Elements equal to
// the k-th key are written from back to front. We need to keep count of them
// separately because the number of elements that <= the k-th key might exceed k.
alignas(128) OutOffsetT out_back_cnt;
// The 'alignas' is necessary to improve the performance of global memory accessing by isolating the request,
// especially for the segment version.
};
enum class candidate_class
{
// The given candidate is definitely amongst the top-k items
selected,
// The given candidate may or may not be amongst the top-k items
candidate,
// The given candidate is definitely not amongst the top-k items
rejected
};
//! @brief AgentTopK implements a stateful abstraction of CUDA thread blocks for participating in
//! device-wide topK
//!
//! @tparam AgentTopKPolicyT
//! Parameterized agent_topk_policy tuning policy type
//!
//! @tparam KeyInputIteratorT
//! **[inferred]** Random-access input iterator type for reading input keys @iterator
//!
//! @tparam KeyOutputIteratorT
//! **[inferred]** Random-access output iterator type for writing output keys @iterator
//!
//! @tparam ValueInputIteratorT
//! **[inferred]** Random-access input iterator type for reading input values @iterator
//!
//! @tparam ValueOutputIteratorT
//! **[inferred]** Random-access output iterator type for writing output values @iterator
//!
//! @tparam ExtractBinOpT
//! Operations to extract the bin from the input key values
//!
//! @tparam IdentifyCandidatesOpT
//! Operations to filter the input key values
//!
//! @tparam OffsetT
//! Type of variable num_items
//!
//! @tparam OutOffsetT
//! Type of variable k
//!
template <typename AgentTopKPolicyT,
typename KeyInputIteratorT,
typename KeyOutputIteratorT,
typename ValueInputIteratorT,
typename ValueOutputIteratorT,
typename ExtractBinOpT,
typename IdentifyCandidatesOpT,
typename OffsetT,
typename OutOffsetT>
struct AgentTopK
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
// The key and value type
using key_in_t = it_value_t<KeyInputIteratorT>;
using value_in_t = it_value_t<ValueInputIteratorT>;
static constexpr int threads_per_block = AgentTopKPolicyT::threads_per_block;
static constexpr int items_per_thread = AgentTopKPolicyT::items_per_thread;
static constexpr int bits_per_pass = AgentTopKPolicyT::bits_per_pass;
static constexpr int tile_items = threads_per_block * items_per_thread;
static constexpr int num_buckets = 1 << bits_per_pass;
static constexpr bool keys_only = ::cuda::std::is_same_v<value_in_t, NullType>;
static constexpr int bins_per_thread = ::cuda::ceil_div(num_buckets, threads_per_block);
// Parameterized BlockLoad type for input data
using block_load_input_t = BlockLoad<key_in_t, threads_per_block, items_per_thread, AgentTopKPolicyT::load_algorithm>;
using block_load_trans_t = BlockLoad<OffsetT, threads_per_block, bins_per_thread, BLOCK_LOAD_TRANSPOSE>;
// Parameterized BlockScan type
using block_scan_t = BlockScan<OffsetT, threads_per_block, AgentTopKPolicyT::SCAN_ALGORITHM>;
// Parameterized BlockStore type
using block_store_trans_t = BlockStore<OffsetT, threads_per_block, bins_per_thread, BLOCK_STORE_TRANSPOSE>;
// Shared memory
struct _TempStorage
{
union
{
// Smem needed for loading
typename block_load_input_t::TempStorage load_input;
typename block_load_trans_t::TempStorage load_trans;
// Smem needed for scan
typename block_scan_t::TempStorage scan;
// Smem needed for storing
typename block_store_trans_t::TempStorage store_trans;
};
OffsetT histogram[num_buckets];
};
/// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
_TempStorage& temp_storage; // Reference to temp_storage
KeyInputIteratorT d_keys_in; // Input keys
KeyOutputIteratorT d_keys_out; // Output keys
ValueInputIteratorT d_values_in; // Input values
ValueOutputIteratorT d_values_out; // Output values
OffsetT num_items; // Total number of input items
OutOffsetT k; // Total number of output items
OffsetT buffer_length; // Size of the buffer for storing intermediate candidates
ExtractBinOpT extract_bin_op; // The operation for bin
IdentifyCandidatesOpT identify_candidates_op; // The operation for filtering
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
//! @param temp_storage
//! Reference to temp_storage
//!
//! @param d_keys_in
//! Input data, keys
//!
//! @param d_keys_out
//! Output data, keys
//!
//! @param d_values_in
//! Input data, values
//!
//! @param d_values_out
//! Output data, values
//!
//! @param num_items
//! Total number of input items
//!
//! @param k
//! The K value. Will find K elements from num_items elements
//!
//! @param buffer_length
//! The size of the buffer for storing intermediate candidates
//!
//! @param extract_bin_op
//! Extract bin operator
//!
//! @param identify_candidates_op
//! Filter operator
//!
_CCCL_DEVICE _CCCL_FORCEINLINE AgentTopK(
TempStorage& temp_storage,
const KeyInputIteratorT d_keys_in,
KeyOutputIteratorT d_keys_out,
const ValueInputIteratorT d_values_in,
ValueOutputIteratorT d_values_out,
OffsetT num_items,
OutOffsetT k,
OffsetT buffer_length,
ExtractBinOpT extract_bin_op,
IdentifyCandidatesOpT identify_candidates_op)
: temp_storage(temp_storage.Alias())
, d_keys_in(d_keys_in)
, d_keys_out(d_keys_out)
, d_values_in(d_values_in)
, d_values_out(d_values_out)
, num_items(num_items)
, k(k)
, buffer_length(buffer_length)
, extract_bin_op(extract_bin_op)
, identify_candidates_op(identify_candidates_op)
{}
//---------------------------------------------------------------------
// Utility methods for device topK
//---------------------------------------------------------------------
// Process a range of input data in tiles, calling f(key, index) for each element
template <typename InputItT, typename FuncT>
_CCCL_DEVICE _CCCL_FORCEINLINE void process_range(InputItT in, const OffsetT num_items, FuncT f)
{
key_in_t thread_data[items_per_thread];
const OffsetT items_per_pass =
static_cast<OffsetT>(tile_items * gridDim.x); // NOLINT(bugprone-misplaced-widening-cast)
const OffsetT total_num_blocks = ::cuda::ceil_div(num_items, tile_items);
const OffsetT num_remaining_elements = num_items % tile_items;
const OffsetT last_block_id = (total_num_blocks - 1) % gridDim.x;
OffsetT tile_base = static_cast<OffsetT>(blockIdx.x * tile_items); // NOLINT(bugprone-misplaced-widening-cast)
OffsetT offset = threadIdx.x * items_per_thread + tile_base;
for (int i_block = static_cast<int>(blockIdx.x); i_block < total_num_blocks - 1;
i_block += static_cast<int>(gridDim.x))
{
// Ensure that the temporary storage from previous iteration can be reused
__syncthreads();
block_load_input_t(temp_storage.load_input).Load(in + tile_base, thread_data);
for (int j = 0; j < items_per_thread; ++j)
{
f(thread_data[j], offset + j);
}
tile_base += items_per_pass;
offset += items_per_pass;
}
// Last tile specialized code-path
if (blockIdx.x == last_block_id)
{
// Ensure that the temporary storage from the previous loop can be reused
__syncthreads();
if (num_remaining_elements == 0)
{
block_load_input_t(temp_storage.load_input).Load(in + tile_base, thread_data);
}
else
{
block_load_input_t(temp_storage.load_input).Load(in + tile_base, thread_data, num_remaining_elements);
}
for (int j = 0; j < items_per_thread; ++j)
{
if ((offset + j) < num_items)
{
f(thread_data[j], offset + j);
}
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void init_histograms(OffsetT* histogram)
{
// Initialize histogram bin counts to zeros
int histo_offset = 0;
// Loop unrolling is beneficial for performance here
_CCCL_PRAGMA_UNROLL_FULL()
for (; histo_offset + threads_per_block <= num_buckets; histo_offset += threads_per_block)
{
histogram[histo_offset + threadIdx.x] = 0;
}
// Finish up with guarded initialization if necessary
if ((num_buckets % threads_per_block != 0) && (histo_offset + threadIdx.x < num_buckets))
{
histogram[histo_offset + threadIdx.x] = 0;
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void merge_histograms(OffsetT* global_histogram)
{
int histo_offset = 0;
// Loop unrolling is beneficial for performance here
_CCCL_PRAGMA_UNROLL_FULL()
for (; histo_offset + threads_per_block <= num_buckets; histo_offset += threads_per_block)
{
if (temp_storage.histogram[histo_offset + threadIdx.x] != 0)
{
atomicAdd(global_histogram + (histo_offset + threadIdx.x), temp_storage.histogram[histo_offset + threadIdx.x]);
}
}
// Finish up with guarded merging if necessary
if ((num_buckets % threads_per_block != 0) && (histo_offset + threadIdx.x < num_buckets))
{
atomicAdd(global_histogram + (histo_offset + threadIdx.x), temp_storage.histogram[histo_offset + threadIdx.x]);
}
}
// Fused filtering of the current pass and building histogram for the next pass
_CCCL_DEVICE _CCCL_FORCEINLINE void filter_and_histogram(
key_in_t* in_buf,
OffsetT* in_idx_buf,
key_in_t* out_buf,
OffsetT* out_idx_buf,
OffsetT previous_len,
Counter<key_in_t, OffsetT, OutOffsetT>* counter,
OffsetT* histogram,
bool early_stop,
bool load_from_original_input)
{
// Initialize shared memory histogram
init_histograms(temp_storage.histogram);
// Make sure the histogram was initialized
__syncthreads();
OffsetT* p_filter_cnt = &counter->filter_cnt;
OutOffsetT* p_out_cnt = &counter->out_cnt;
// Lambda for early_stop = true (i.e., we have identified the exact "splitter" key):
// Select all items that fall into the bin of the k-th item (i.e., the 'candidates') and the ones that fall into
// bins preceding the k-th item bin (i.e., 'selected' items), write them to output.
// We can skip histogram computation because we don't need to further passes to refine the candidates.
auto f_early_stop = [load_from_original_input, in_idx_buf, p_out_cnt, this](key_in_t key, OffsetT i) {
const candidate_class pre_res = identify_candidates_op(key);
if (pre_res == candidate_class::candidate || pre_res == candidate_class::selected)
{
const OutOffsetT pos = atomicAdd(p_out_cnt, OutOffsetT{1});
d_keys_out[pos] = key;
if constexpr (!keys_only)
{
const OffsetT index = load_from_original_input ? i : in_idx_buf[i];
d_values_out[pos] = d_values_in[index];
}
}
};
// Lambda for early_stop = false, out_buf != nullptr (i.e., we need to further refine the candidates in the next
// pass): Write out selected items to output, write candidates to out_buf, and build histogram for candidates.
auto f_with_out_buf = [load_from_original_input, in_idx_buf, out_buf, out_idx_buf, p_filter_cnt, p_out_cnt, this](
key_in_t key, OffsetT i) {
const candidate_class pre_res = identify_candidates_op(key);
if (pre_res == candidate_class::candidate)
{
const OffsetT pos = atomicAdd(p_filter_cnt, OffsetT{1});
out_buf[pos] = key;
if constexpr (!keys_only)
{
const OffsetT index = load_from_original_input ? i : in_idx_buf[i];
out_idx_buf[pos] = index;
}
const int bucket = extract_bin_op(key);
atomicAdd(temp_storage.histogram + bucket, OffsetT{1});
}
else if (pre_res == candidate_class::selected)
{
const OutOffsetT pos = atomicAdd(p_out_cnt, OutOffsetT{1});
d_keys_out[pos] = key;
if constexpr (!keys_only)
{
const OffsetT index = in_idx_buf ? in_idx_buf[i] : i;
d_values_out[pos] = d_values_in[index];
}
}
};
// Lambda for early_stop = false, out_buf = nullptr (i.e., we need to further refine the candidates in the next
// pass, but we skip writing candidates to out_buf):
// Just build histogram for candidates.
// Note: We will only begin writing to d_keys_out starting from the pass in which the number of output-candidates
// is small enough to fit into the output buffer (otherwise, we would be writing the same items to d_keys_out
// multiple times).
auto f_no_out_buf = [this](key_in_t key, OffsetT i) {
const candidate_class pre_res = identify_candidates_op(key);
if (pre_res == candidate_class::candidate)
{
const int bucket = extract_bin_op(key);
atomicAdd(temp_storage.histogram + bucket, OffsetT{1});
}
};
// Choose and invoke the appropriate lambda with the correct input source
// If the input size exceeds the allocated buffer size, we know for sure we haven't started writing candidates to
// the output buffer yet
if (load_from_original_input)
{
if (early_stop)
{
process_range(d_keys_in, previous_len, f_early_stop);
}
else if (out_buf)
{
process_range(d_keys_in, previous_len, f_with_out_buf);
}
else
{
process_range(d_keys_in, previous_len, f_no_out_buf);
}
}
else
{
if (early_stop)
{
process_range(in_buf, previous_len, f_early_stop);
}
else if (out_buf)
{
process_range(in_buf, previous_len, f_with_out_buf);
}
else
{
process_range(in_buf, previous_len, f_no_out_buf);
}
}
// Early stop means that subsequent passes are not needed
if (early_stop)
{
return;
}
// Ensure all threads have contributed to the histogram before accumulating in the global memory
__syncthreads();
// Merge the locally aggregated histogram into the global histogram
merge_histograms(histogram);
}
// Replace histogram with its own prefix sum
_CCCL_DEVICE _CCCL_FORCEINLINE void compute_bin_offsets(volatile OffsetT* histogram)
{
OffsetT thread_data[bins_per_thread]{};
// Load global histogram (we can skip initializing oob-items to zero because they won't be stored back)
block_load_trans_t(temp_storage.load_trans).Load(histogram, thread_data, num_buckets);
__syncthreads();
block_scan_t(temp_storage.scan).InclusiveSum(thread_data, thread_data);
__syncthreads();
block_store_trans_t(temp_storage.store_trans).Store(temp_storage.histogram, thread_data, num_buckets);
}
// Identify the bucket that the k-th value falls into
_CCCL_DEVICE _CCCL_FORCEINLINE void
choose_bucket(Counter<key_in_t, OffsetT, OutOffsetT>* counter, const OutOffsetT k, const int pass)
{
// Initialize histogram bin counts to zeros
int histo_offset = 0;
auto body = [&] {
const int bin_idx = static_cast<int>(histo_offset + threadIdx.x);
const OffsetT prev = (bin_idx == 0) ? 0 : temp_storage.histogram[bin_idx - 1];
const OffsetT cur = temp_storage.histogram[bin_idx];
// Identify the bin that the k-th item falls into. One and only one thread will satisfy this condition, so counter
// is written by only one thread
if (prev < k && cur >= k)
{
// The number of items that are yet to be identified
counter->k = k - prev;
// The number of candidates in the next pass
counter->len = cur - prev;
const unsigned int bucket = static_cast<unsigned int>(bin_idx);
// Update the "splitter" key by adding the radix digit of the k-th item bin of this pass
set_kth_key_bits<key_in_t, bits_per_pass>(counter->kth_key_bits, pass, bucket);
}
};
_CCCL_PRAGMA_UNROLL_FULL()
for (; histo_offset + threads_per_block <= num_buckets; histo_offset += threads_per_block)
{
body();
}
// Finish up with guarded initialization if necessary
if ((num_buckets % threads_per_block != 0) && (histo_offset + threadIdx.x < num_buckets))
{
body();
}
}
// Performs the last-block coordination after histogram accumulation: ensures global visibility,
// detects the last finishing block, runs the prefix sum, identifies the k-th bucket, and resets
// the histogram for the next pass. The caller-supplied counter_update_fn runs on thread 0 of the
// last block to update pass-specific counter state.
template <typename CounterUpdateFn>
_CCCL_DEVICE _CCCL_FORCEINLINE void finalize_pass(
Counter<key_in_t, OffsetT, OutOffsetT>* counter,
OffsetT* histogram,
OutOffsetT current_k,
int pass,
bool is_last_pass,
CounterUpdateFn counter_update_fn)
{
// Ensure all writes to the global memory-histogram are visible to all threads before
// proceeding to compute the prefix sum over the histogram.
__threadfence();
// Identify the last block in the grid to perform the prefix sum over the histogram
bool is_last_block = false;
if (threadIdx.x == 0)
{
unsigned int finished = atomicInc(&counter->finished_block_cnt, gridDim.x - 1);
is_last_block = (finished == (gridDim.x - 1));
}
// syncthreads ensures that the BlockLoad for loading the global histogram can reuse the temporary storage
if (__syncthreads_or(is_last_block))
{
if (threadIdx.x == 0)
{
counter_update_fn();
}
// Compute prefix sum over the histogram's bin counts
compute_bin_offsets(histogram);
// Make sure the prefix sum has been written to shared memory before choose_bucket()
__syncthreads();
// Identify the bucket that the k-th item falls into
choose_bucket(counter, current_k, pass);
if (!is_last_pass)
{
init_histograms(histogram);
}
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void invoke_last_filter(
key_in_t* in_buf, OffsetT* in_idx_buf, Counter<key_in_t, OffsetT, OutOffsetT>* counter, OutOffsetT k, int pass)
{
const bool load_from_original_input = (pass <= 1) || counter->previous_len > buffer_length;
const OffsetT current_len = load_from_original_input ? num_items : counter->previous_len;
in_idx_buf = load_from_original_input ? nullptr : in_idx_buf; // ? out_idx_buf : in_idx_buf;
if (current_len == 0)
{
return;
}
// changed in choose_bucket(); need to reload
OffsetT num_of_kth_needed = counter->k;
OutOffsetT* p_out_cnt = &counter->out_cnt;
OutOffsetT* p_out_back_cnt = &counter->out_back_cnt;
auto f = [this, p_out_cnt, in_idx_buf, p_out_back_cnt, num_of_kth_needed, k, load_from_original_input](
key_in_t key, OffsetT i) {
const candidate_class res = identify_candidates_op(key);
if (res == candidate_class::selected)
{
const OutOffsetT pos = atomicAdd(p_out_cnt, OffsetT{1});
d_keys_out[pos] = key;
if constexpr (!keys_only)
{
// If writing has been skipped up to this point, `in_idx_buf` is nullptr
const OffsetT index = load_from_original_input ? i : in_idx_buf[i];
d_values_out[pos] = d_values_in[index];
}
}
else if (res == candidate_class::candidate)
{
const OutOffsetT back_pos = atomicAdd(p_out_back_cnt, OffsetT{1});
if (back_pos < num_of_kth_needed)
{
const OutOffsetT pos = k - 1 - back_pos;
d_keys_out[pos] = key;
if constexpr (!keys_only)
{
const OffsetT new_idx = load_from_original_input ? i : in_idx_buf[i];
d_values_out[pos] = d_values_in[new_idx];
}
}
}
};
if (load_from_original_input)
{
process_range(d_keys_in, current_len, f);
}
else
{
process_range(in_buf, current_len, f);
}
}
_CCCL_DEVICE _CCCL_FORCEINLINE void invoke_filter_and_histogram(
key_in_t* in_buf,
OffsetT* in_idx_buf,
key_in_t* out_buf,
OffsetT* out_idx_buf,
Counter<key_in_t, OffsetT, OutOffsetT>* counter,
OffsetT* histogram,
int pass,
bool is_last_pass)
{
const OutOffsetT current_k = counter->k;
const OffsetT current_len = counter->len;
OffsetT previous_len = counter->previous_len;
// If current_len is 0, it means all the candidates have been found in previous passes.
if (current_len == 0)
{
return;
}
// Early stop means that the bin containing the k-th element has been identified, and all
// the elements in this bin are exactly the remaining k items we need to find. So we can
// stop the process after this filtering pass.
const bool early_stop = (current_len == static_cast<OffsetT>(current_k));
// If previous_len > buffer_length, it means we haven't started writing candidates to out_buf yet,
// so have to make sure to load input directly from the original input.
// Also, unless we've had the chance to do at least one filtering pass, our input is definitely the original input
// (this is to guard against edge cases, e.g., buffer_length=num_items=1).
const bool load_from_original_input = (pass <= 1) || previous_len > buffer_length;
if (load_from_original_input)
{
in_idx_buf = nullptr;
previous_len = num_items;
}
// "current_len > buffer_length" means current pass will skip writing buffer
if (current_len > buffer_length)
{
out_buf = nullptr;
out_idx_buf = nullptr;
}
// Fused filtering of candidates and histogram computation over the output-candidates
filter_and_histogram(
in_buf, in_idx_buf, out_buf, out_idx_buf, previous_len, counter, histogram, early_stop, load_from_original_input);
finalize_pass(counter, histogram, current_k, pass, is_last_pass, [counter, current_len, early_stop] {
if (early_stop)
{
counter->previous_len = 0;
counter->len = 0;
}
else
{
counter->previous_len = current_len;
counter->filter_cnt = 0;
}
});
}
// Histogram-only pass: computes the histogram over the full input without filtering.
// Used for the first radix pass before any candidates have been identified.
_CCCL_DEVICE _CCCL_FORCEINLINE void invoke_histogram_only(
Counter<key_in_t, OffsetT, OutOffsetT>* counter, OffsetT* histogram, int pass, bool is_last_pass)
{
// Initialize shared memory histogram
init_histograms(temp_storage.histogram);
__syncthreads();
// Compute per-thread block histograms over the full input
auto f = [this](key_in_t key, OffsetT /*index*/) {
const int bucket = extract_bin_op(key);
atomicAdd(temp_storage.histogram + bucket, OffsetT{1});
};
process_range(d_keys_in, num_items, f);
// Ensure all threads have contributed to the histogram before accumulating in global memory
__syncthreads();
// Merge the locally aggregated histogram into the global histogram
merge_histograms(histogram);
finalize_pass(counter, histogram, k, pass, is_last_pass, [counter, this] {
counter->previous_len = num_items;
counter->filter_cnt = 0;
});
}
};
} // namespace detail::topk
CUB_NAMESPACE_END

View File

@@ -0,0 +1,586 @@
// SPDX-FileCopyrightText: Copyright (c), NVIDIA CORPORATION. All rights reserved.
// SPDX-License-Identifier: BSD-3
/**
* @file
* cub::AgentUniqueByKey implements a stateful abstraction of CUDA thread blocks for participating in device-wide
* unique-by-key.
*/
#pragma once
#include <cub/config.cuh>
#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)
# pragma GCC system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)
# pragma clang system_header
#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)
# pragma system_header
#endif // no system header
#include <cub/agent/single_pass_scan_operators.cuh>
#include <cub/block/block_discontinuity.cuh>
#include <cub/block/block_load.cuh>
#include <cub/block/block_scan.cuh>
#include <cub/thread/thread_operators.cuh>
CUB_NAMESPACE_BEGIN
/******************************************************************************
* Tuning policy types
******************************************************************************/
namespace detail
{
// TODO(bgruber): remove this when C++20 is the minimum, since then we can pass policy values as NTTP
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
cub::BlockLoadAlgorithm LoadAlgorithm = cub::BLOCK_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifier = cub::LOAD_LDG,
cub::BlockScanAlgorithm ScanAlgorithm = cub::BLOCK_SCAN_WARP_SCANS,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
struct agent_unique_by_key_policy
{
static constexpr int BLOCK_THREADS = ThreadsPerBlock;
static constexpr int ITEMS_PER_THREAD = ItemsPerThread;
static constexpr cub::BlockLoadAlgorithm LOAD_ALGORITHM = LoadAlgorithm;
static constexpr cub::CacheLoadModifier LOAD_MODIFIER = LoadModifier;
static constexpr cub::BlockScanAlgorithm SCAN_ALGORITHM = ScanAlgorithm;
struct detail
{
using delay_constructor_t = DelayConstructorT;
};
};
} // namespace detail
//! Deprecated [Since 3.5]
template <int ThreadsPerBlock,
int ItemsPerThread = 1,
cub::BlockLoadAlgorithm LoadAlgorithm = cub::BLOCK_LOAD_DIRECT,
cub::CacheLoadModifier LoadModifier = cub::LOAD_LDG,
cub::BlockScanAlgorithm ScanAlgorithm = cub::BLOCK_SCAN_WARP_SCANS,
typename DelayConstructorT = detail::fixed_delay_constructor_t<350, 450>>
using AgentUniqueByKeyPolicy CCCL_DEPRECATED_BECAUSE("Use the tuning API for DeviceSelect") = detail::
agent_unique_by_key_policy<ThreadsPerBlock, ItemsPerThread, LoadAlgorithm, LoadModifier, ScanAlgorithm, DelayConstructorT>;
/******************************************************************************
* Thread block abstractions
******************************************************************************/
namespace detail::unique_by_key
{
/**
* @brief AgentUniqueByKey implements a stateful abstraction of CUDA thread blocks for participating
* in device-wide unique-by-key
*
* @tparam AgentUniqueByKeyPolicyT
* Parameterized AgentUniqueByKeyPolicy tuning policy type
*
* @tparam KeyInputIteratorT
* Random-access input iterator type for keys
*
* @tparam ValueInputIteratorT
* Random-access input iterator type for values
*
* @tparam KeyOutputIteratorT
* Random-access output iterator type for keys
*
* @tparam ValueOutputIteratorT
* Random-access output iterator type for values
*
* @tparam EqualityOpT
* Equality operator type
*
* @tparam OffsetT
* Signed integer type for global offsets
*/
template <typename AgentUniqueByKeyPolicyT,
typename KeyInputIteratorT,
typename ValueInputIteratorT,
typename KeyOutputIteratorT,
typename ValueOutputIteratorT,
typename EqualityOpT,
typename OffsetT>
struct AgentUniqueByKey
{
//---------------------------------------------------------------------
// Types and constants
//---------------------------------------------------------------------
// The input key and value type
using KeyT = cub::detail::it_value_t<KeyInputIteratorT>;
using ValueT = cub::detail::it_value_t<ValueInputIteratorT>;
// Tile status descriptor interface type
using ScanTileStateT = ScanTileState<OffsetT>;
// Constants
static constexpr int BLOCK_THREADS = AgentUniqueByKeyPolicyT::BLOCK_THREADS;
static constexpr int ITEMS_PER_THREAD = AgentUniqueByKeyPolicyT::ITEMS_PER_THREAD;
static constexpr int ITEMS_PER_TILE = BLOCK_THREADS * ITEMS_PER_THREAD;
// Cache-modified Input iterator wrapper type (for applying cache modifier) for keys
using WrappedKeyInputIteratorT = ::cuda::std::conditional_t<
::cuda::std::is_pointer_v<KeyInputIteratorT>,
CacheModifiedInputIterator<AgentUniqueByKeyPolicyT::LOAD_MODIFIER, KeyT, OffsetT>, // Wrap the native input pointer
// with
// CacheModifiedValuesInputIterator
KeyInputIteratorT>; // Directly use the supplied input iterator type
// Cache-modified Input iterator wrapper type (for applying cache modifier) for values
using WrappedValueInputIteratorT = ::cuda::std::conditional_t<
::cuda::std::is_pointer_v<ValueInputIteratorT>,
CacheModifiedInputIterator<AgentUniqueByKeyPolicyT::LOAD_MODIFIER, ValueT, OffsetT>, // Wrap the native input
// pointer with
// CacheModifiedValuesInputIterator
ValueInputIteratorT>; // Directly use the supplied input iterator type
// Parameterized BlockLoad type for input data
using BlockLoadKeys = BlockLoad<KeyT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentUniqueByKeyPolicyT::LOAD_ALGORITHM>;
// Parameterized BlockLoad type for flags
using BlockLoadValues = BlockLoad<ValueT, BLOCK_THREADS, ITEMS_PER_THREAD, AgentUniqueByKeyPolicyT::LOAD_ALGORITHM>;
// Parameterized BlockDiscontinuity type for items
using BlockDiscontinuityKeys = cub::BlockDiscontinuity<KeyT, BLOCK_THREADS>;
// Parameterized BlockScan type
using BlockScanT = cub::BlockScan<OffsetT, BLOCK_THREADS, AgentUniqueByKeyPolicyT::SCAN_ALGORITHM>;
// Parameterized BlockDiscontinuity type for items
using DelayConstructorT = typename AgentUniqueByKeyPolicyT::detail::delay_constructor_t;
using TilePrefixCallback = cub::TilePrefixCallbackOp<OffsetT, ::cuda::std::plus<>, ScanTileStateT, DelayConstructorT>;
// Key exchange type
using KeyExchangeT = KeyT[ITEMS_PER_TILE];
// Value exchange type
using ValueExchangeT = ValueT[ITEMS_PER_TILE];
// Shared memory type for this thread block
union _TempStorage
{
struct ScanStorage
{
typename BlockScanT::TempStorage scan;
typename TilePrefixCallback::TempStorage prefix;
typename BlockDiscontinuityKeys::TempStorage discontinuity;
} scan_storage;
// Smem needed for loading keys
typename BlockLoadKeys::TempStorage load_keys;
// Smem needed for loading values
typename BlockLoadValues::TempStorage load_values;
// Smem needed for compacting items (allows non POD items in this union)
Uninitialized<KeyExchangeT> shared_keys;
Uninitialized<ValueExchangeT> shared_values;
};
// Alias wrapper allowing storage to be unioned
struct TempStorage : Uninitialized<_TempStorage>
{};
//---------------------------------------------------------------------
// Per-thread fields
//---------------------------------------------------------------------
_TempStorage& temp_storage;
WrappedKeyInputIteratorT d_keys_in;
WrappedValueInputIteratorT d_values_in;
KeyOutputIteratorT d_keys_out;
ValueOutputIteratorT d_values_out;
cub::InequalityWrapper<EqualityOpT> inequality_op;
OffsetT num_items;
//---------------------------------------------------------------------
// Constructor
//---------------------------------------------------------------------
// Constructor
_CCCL_DEVICE _CCCL_FORCEINLINE AgentUniqueByKey(
TempStorage& temp_storage_,
WrappedKeyInputIteratorT d_keys_in_,
WrappedValueInputIteratorT d_values_in_,
KeyOutputIteratorT d_keys_out_,
ValueOutputIteratorT d_values_out_,
EqualityOpT equality_op_,
OffsetT num_items_)
: temp_storage(temp_storage_.Alias())
, d_keys_in(d_keys_in_)
, d_values_in(d_values_in_)
, d_keys_out(d_keys_out_)
, d_values_out(d_values_out_)
, inequality_op(equality_op_)
, num_items(num_items_)
{}
//---------------------------------------------------------------------
// Utility functions
//---------------------------------------------------------------------
struct KeyTagT
{};
struct ValueTagT
{};
_CCCL_DEVICE _CCCL_FORCEINLINE KeyExchangeT& GetShared(KeyTagT)
{
return temp_storage.shared_keys.Alias();
}
_CCCL_DEVICE _CCCL_FORCEINLINE ValueExchangeT& GetShared(ValueTagT)
{
return temp_storage.shared_values.Alias();
}
//---------------------------------------------------------------------
// Scatter utility methods
//---------------------------------------------------------------------
template <typename Tag, typename OutputIt, typename T>
_CCCL_DEVICE _CCCL_FORCEINLINE void Scatter(
Tag tag,
OutputIt items_out,
T (&items)[ITEMS_PER_THREAD],
OffsetT (&selection_flags)[ITEMS_PER_THREAD],
OffsetT (&selection_indices)[ITEMS_PER_THREAD],
int /*num_tile_items*/,
int num_tile_selections,
OffsetT num_selections_prefix,
OffsetT /*num_selections*/)
{
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
int local_scatter_offset = selection_indices[ITEM] - num_selections_prefix;
if (selection_flags[ITEM])
{
GetShared(tag)[local_scatter_offset] = items[ITEM];
}
}
__syncthreads();
// Preventing loop unrolling helps avoid perf degradation when switching from signed to unsigned 32-bit offset
// types
_CCCL_PRAGMA_NOUNROLL()
for (int item = static_cast<int>(threadIdx.x); item < num_tile_selections; item += BLOCK_THREADS)
{
items_out[num_selections_prefix + item] = GetShared(tag)[item]; // NOLINT(bugprone-misplaced-widening-cast)
}
__syncthreads();
}
//---------------------------------------------------------------------
// Cooperatively scan a device-wide sequence of tiles with other CTAs
//---------------------------------------------------------------------
/**
* @brief Process first tile of input (dynamic chained scan).
*
* @param num_tile_items
* Number of input items comprising this tile
*
* @param tile_offset
* Tile offset
*
* @param tile_state
* Global tile state descriptor
*
* @return The running count of selections (including this tile)
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE OffsetT
ConsumeFirstTile(int num_tile_items, OffsetT tile_offset, ScanTileStateT& tile_state)
{
KeyT keys[ITEMS_PER_THREAD];
OffsetT selection_flags[ITEMS_PER_THREAD];
OffsetT selection_idx[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last elements with the first element
// because collectives are not suffix guarded
BlockLoadKeys(temp_storage.load_keys)
.Load(d_keys_in + tile_offset, keys, num_tile_items, *(d_keys_in + tile_offset));
}
else
{
BlockLoadKeys(temp_storage.load_keys).Load(d_keys_in + tile_offset, keys);
}
__syncthreads();
ValueT values[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last elements with the first element
// because collectives are not suffix guarded
BlockLoadValues(temp_storage.load_values)
.Load(d_values_in + tile_offset, values, num_tile_items, *(d_values_in + tile_offset));
}
else
{
BlockLoadValues(temp_storage.load_values).Load(d_values_in + tile_offset, values);
}
__syncthreads();
BlockDiscontinuityKeys(temp_storage.scan_storage.discontinuity).FlagHeads(selection_flags, keys, inequality_op);
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
// Set selection_flags for out-of-bounds items
if ((IS_LAST_TILE) && (OffsetT(threadIdx.x * ITEMS_PER_THREAD) + ITEM >= num_tile_items))
{
selection_flags[ITEM] = 1;
}
}
__syncthreads();
OffsetT num_tile_selections = 0;
OffsetT num_selections = 0;
OffsetT num_selections_prefix = 0;
BlockScanT(temp_storage.scan_storage.scan).ExclusiveSum(selection_flags, selection_idx, num_tile_selections);
if (threadIdx.x == 0)
{
// Update tile status if this is not the last tile
if constexpr (!IS_LAST_TILE)
{
tile_state.SetInclusive(0, num_tile_selections);
}
}
// Do not count any out-of-bounds selections
if constexpr (IS_LAST_TILE)
{
int num_discount = ITEMS_PER_TILE - num_tile_items;
num_tile_selections -= num_discount;
}
num_selections = num_tile_selections;
__syncthreads();
Scatter(KeyTagT(),
d_keys_out,
keys,
selection_flags,
selection_idx,
num_tile_items,
num_tile_selections,
num_selections_prefix,
num_selections);
__syncthreads();
Scatter(ValueTagT(),
d_values_out,
values,
selection_flags,
selection_idx,
num_tile_items,
num_tile_selections,
num_selections_prefix,
num_selections);
return num_selections;
}
/**
* @brief Process subsequent tile of input (dynamic chained scan).
*
* @param num_tile_items
* Number of input items comprising this tile
*
* @param tile_idx
* Tile index
*
* @param tile_offset
* Tile offset
*
* @param tile_state
* Global tile state descriptor
*
* @return Returns the running count of selections (including this tile)
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE OffsetT
ConsumeSubsequentTile(int num_tile_items, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state)
{
KeyT keys[ITEMS_PER_THREAD];
OffsetT selection_flags[ITEMS_PER_THREAD];
OffsetT selection_idx[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last elements with the first element
// because collectives are not suffix guarded
BlockLoadKeys(temp_storage.load_keys)
.Load(d_keys_in + tile_offset, keys, num_tile_items, *(d_keys_in + tile_offset));
}
else
{
BlockLoadKeys(temp_storage.load_keys).Load(d_keys_in + tile_offset, keys);
}
__syncthreads();
ValueT values[ITEMS_PER_THREAD];
if constexpr (IS_LAST_TILE)
{
// Fill last elements with the first element
// because collectives are not suffix guarded
BlockLoadValues(temp_storage.load_values)
.Load(d_values_in + tile_offset, values, num_tile_items, *(d_values_in + tile_offset));
}
else
{
BlockLoadValues(temp_storage.load_values).Load(d_values_in + tile_offset, values);
}
__syncthreads();
KeyT tile_predecessor = d_keys_in[tile_offset - 1];
BlockDiscontinuityKeys(temp_storage.scan_storage.discontinuity)
.FlagHeads(selection_flags, keys, inequality_op, tile_predecessor);
_CCCL_PRAGMA_UNROLL_FULL()
for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)
{
// Set selection_flags for out-of-bounds items
if ((IS_LAST_TILE) && (OffsetT(threadIdx.x * ITEMS_PER_THREAD) + ITEM >= num_tile_items))
{
selection_flags[ITEM] = 1;
}
}
__syncthreads();
OffsetT num_tile_selections = 0;
OffsetT num_selections = 0;
OffsetT num_selections_prefix = 0;
TilePrefixCallback prefix_cb(tile_state, temp_storage.scan_storage.prefix, ::cuda::std::plus<>{}, tile_idx);
BlockScanT(temp_storage.scan_storage.scan).ExclusiveSum(selection_flags, selection_idx, prefix_cb);
num_selections = prefix_cb.GetInclusivePrefix();
num_tile_selections = prefix_cb.GetBlockAggregate();
num_selections_prefix = prefix_cb.GetExclusivePrefix();
if constexpr (IS_LAST_TILE)
{
int num_discount = ITEMS_PER_TILE - num_tile_items;
num_tile_selections -= num_discount;
num_selections -= num_discount;
}
__syncthreads();
Scatter(KeyTagT(),
d_keys_out,
keys,
selection_flags,
selection_idx,
num_tile_items,
num_tile_selections,
num_selections_prefix,
num_selections);
__syncthreads();
Scatter(ValueTagT(),
d_values_out,
values,
selection_flags,
selection_idx,
num_tile_items,
num_tile_selections,
num_selections_prefix,
num_selections);
return num_selections;
}
/**
* @brief Process a tile of input
*
* @param num_tile_items
* Number of input items comprising this tile
*
* @param tile_idx
* Tile index
*
* @param tile_offset
* Tile offset
*
* @param tile_state
* Global tile state descriptor
*/
template <bool IS_LAST_TILE>
_CCCL_DEVICE _CCCL_FORCEINLINE OffsetT
ConsumeTile(int num_tile_items, int tile_idx, OffsetT tile_offset, ScanTileStateT& tile_state)
{
OffsetT num_selections;
if (tile_idx == 0)
{
num_selections = ConsumeFirstTile<IS_LAST_TILE>(num_tile_items, tile_offset, tile_state);
}
else
{
num_selections = ConsumeSubsequentTile<IS_LAST_TILE>(num_tile_items, tile_idx, tile_offset, tile_state);
}
return num_selections;
}
/**
* @brief Scan tiles of items as part of a dynamic chained scan
*
* @param num_tiles
* Total number of input tiles
*
* @param tile_state
* Global tile state descriptor
*
* @param d_num_selected_out
* Output total number selection_flags
*
* @tparam NumSelectedIteratorT
* Output iterator type for recording number of items selection_flags
*
*/
template <typename NumSelectedIteratorT>
_CCCL_DEVICE _CCCL_FORCEINLINE void
ConsumeRange(int num_tiles, ScanTileStateT& tile_state, NumSelectedIteratorT d_num_selected_out)
{
// Blocks are launched in increasing order, so just assign one tile per block
int tile_idx = static_cast<int>((blockIdx.x * gridDim.y) + blockIdx.y); // Current tile index
// Global offset for the current tile
OffsetT tile_offset = static_cast<OffsetT>(tile_idx) * static_cast<OffsetT>(ITEMS_PER_TILE);
if (tile_idx < num_tiles - 1)
{
ConsumeTile<false>(ITEMS_PER_TILE, tile_idx, tile_offset, tile_state);
}
else
{
int num_remaining = static_cast<int>(num_items - tile_offset);
OffsetT num_selections = ConsumeTile<true>(num_remaining, tile_idx, tile_offset, tile_state);
if (threadIdx.x == 0)
{
*d_num_selected_out = num_selections;
}
}
}
};
} // namespace detail::unique_by_key
CUB_NAMESPACE_END

File diff suppressed because it is too large Load Diff