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:
@@ -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
@@ -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
|
||||
201
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_find.cuh
Normal file
201
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_find.cuh
Normal 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
|
||||
@@ -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
|
||||
56
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_for.cuh
Normal file
56
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_for.cuh
Normal 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
|
||||
@@ -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
|
||||
335
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_merge.cuh
Normal file
335
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_merge.cuh
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
611
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_reduce.cuh
Normal file
611
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_reduce.cuh
Normal 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
|
||||
@@ -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
|
||||
1072
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_rle.cuh
Normal file
1072
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_rle.cuh
Normal file
File diff suppressed because it is too large
Load Diff
567
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_scan.cuh
Normal file
567
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_scan.cuh
Normal 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
|
||||
@@ -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
|
||||
@@ -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
|
||||
1074
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_select_if.cuh
Normal file
1074
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_select_if.cuh
Normal file
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
835
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_topk.cuh
Normal file
835
qwen3_6_scripts/cccl_preload/include/cub/agent/agent_topk.cuh
Normal 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
|
||||
@@ -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
Reference in New Issue
Block a user