codekingpro/portable-devtools
114k
1/******************************************************************************2 * Copyright (c) 2011-2021, NVIDIA CORPORATION. All rights reserved.3 *4 * Redistribution and use in source and binary forms, with or without5 * modification, are permitted provided that the following conditions are met:6 * * Redistributions of source code must retain the above copyright7 * notice, this list of conditions and the following disclaimer.8 * * Redistributions in binary form must reproduce the above copyright9 * notice, this list of conditions and the following disclaimer in the10 * documentation and/or other materials provided with the distribution.11 * * Neither the name of the NVIDIA CORPORATION nor the12 * names of its contributors may be used to endorse or promote products13 * derived from this software without specific prior written permission.14 *15 * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND16 * ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED17 * WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE18 * DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE FOR ANY19 * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES20 * (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;21 * LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND22 * ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT23 * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS24 * SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.25 *26 ******************************************************************************/27 28#pragma once29 30#include <cub/config.cuh>31 32#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)33# pragma GCC system_header34#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)35# pragma clang system_header36#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)37# pragma system_header38#endif // no system header39 40#include <cub/agent/agent_radix_sort_downsweep.cuh>41#include <cub/agent/agent_radix_sort_upsweep.cuh>42#include <cub/block/block_radix_sort.cuh>43#include <cub/util_namespace.cuh>44#include <cub/util_type.cuh>45 46 47CUB_NAMESPACE_BEGIN48 49 50/**51 * This agent will be implementing the `DeviceSegmentedRadixSort` when the52 * https://github.com/NVIDIA/cub/issues/383 is addressed.53 *54 * @tparam IS_DESCENDING55 * Whether or not the sorted-order is high-to-low56 *57 * @tparam SegmentedPolicyT58 * Chained tuning policy59 *60 * @tparam KeyT61 * Key type62 *63 * @tparam ValueT64 * Value type65 *66 * @tparam OffsetT67 * Signed integer type for global offsets68 */69template <bool IS_DESCENDING,70 typename SegmentedPolicyT,71 typename KeyT,72 typename ValueT,73 typename OffsetT,74 typename DecomposerT = detail::identity_decomposer_t>75struct AgentSegmentedRadixSort76{77 OffsetT num_items;78 79 static constexpr int ITEMS_PER_THREAD = SegmentedPolicyT::ITEMS_PER_THREAD;80 static constexpr int BLOCK_THREADS = SegmentedPolicyT::BLOCK_THREADS;81 static constexpr int RADIX_BITS = SegmentedPolicyT::RADIX_BITS;82 static constexpr int RADIX_DIGITS = 1 << RADIX_BITS;83 static constexpr int KEYS_ONLY = std::is_same<ValueT, NullType>::value;84 85 using traits = detail::radix::traits_t<KeyT>;86 using bit_ordered_type = typename traits::bit_ordered_type;87 88 // Huge segment handlers89 using BlockUpsweepT = AgentRadixSortUpsweep<SegmentedPolicyT, KeyT, OffsetT, DecomposerT>;90 using DigitScanT = BlockScan<OffsetT, BLOCK_THREADS>;91 using BlockDownsweepT =92 AgentRadixSortDownsweep<SegmentedPolicyT, IS_DESCENDING, KeyT, ValueT, OffsetT, DecomposerT>;93 94 /// Number of bin-starting offsets tracked per thread95 static constexpr int BINS_TRACKED_PER_THREAD = BlockDownsweepT::BINS_TRACKED_PER_THREAD;96 97 // Small segment handlers98 using BlockRadixSortT =99 BlockRadixSort<KeyT,100 BLOCK_THREADS,101 ITEMS_PER_THREAD,102 ValueT,103 RADIX_BITS,104 (SegmentedPolicyT::RANK_ALGORITHM == RADIX_RANK_MEMOIZE),105 SegmentedPolicyT::SCAN_ALGORITHM>;106 107 using BlockKeyLoadT = BlockLoad<KeyT,108 BLOCK_THREADS,109 ITEMS_PER_THREAD,110 SegmentedPolicyT::LOAD_ALGORITHM>;111 112 using BlockValueLoadT = BlockLoad<ValueT,113 BLOCK_THREADS,114 ITEMS_PER_THREAD,115 SegmentedPolicyT::LOAD_ALGORITHM>;116 117 union _TempStorage118 {119 // Huge segment handlers120 typename BlockUpsweepT::TempStorage upsweep;121 typename BlockDownsweepT::TempStorage downsweep;122 123 struct UnboundBlockSort124 {125 OffsetT reverse_counts_in[RADIX_DIGITS];126 OffsetT reverse_counts_out[RADIX_DIGITS];127 typename DigitScanT::TempStorage scan;128 } unbound_sort;129 130 // Small segment handlers131 typename BlockKeyLoadT::TempStorage keys_load;132 typename BlockValueLoadT::TempStorage values_load;133 typename BlockRadixSortT::TempStorage sort;134 };135 136 using TempStorage = Uninitialized<_TempStorage>;137 _TempStorage &temp_storage;138 139 DecomposerT decomposer;140 141 __device__ __forceinline__142 AgentSegmentedRadixSort(OffsetT num_items,143 TempStorage &temp_storage,144 DecomposerT decomposer = {})145 : num_items(num_items)146 , temp_storage(temp_storage.Alias())147 , decomposer(decomposer)148 {}149 150 __device__ __forceinline__ void ProcessSinglePass(int begin_bit,151 int end_bit,152 const KeyT *d_keys_in,153 const ValueT *d_values_in,154 KeyT *d_keys_out,155 ValueT *d_values_out)156 {157 KeyT thread_keys[ITEMS_PER_THREAD];158 ValueT thread_values[ITEMS_PER_THREAD];159 160 // For FP64 the difference is:161 // Lowest() -> -1.79769e+308 = 00...00b -> TwiddleIn -> -0 = 10...00b162 // LOWEST -> -nan = 11...11b -> TwiddleIn -> 0 = 00...00b163 164 bit_ordered_type default_key_bits = IS_DESCENDING165 ? traits::min_raw_binary_key(decomposer)166 : traits::max_raw_binary_key(decomposer);167 KeyT oob_default = reinterpret_cast<KeyT &>(default_key_bits);168 169 if (!KEYS_ONLY)170 {171 BlockValueLoadT(temp_storage.values_load)172 .Load(d_values_in, thread_values, num_items);173 174 CTA_SYNC();175 }176 177 {178 BlockKeyLoadT(temp_storage.keys_load)179 .Load(d_keys_in, thread_keys, num_items, oob_default);180 181 CTA_SYNC();182 }183 184 BlockRadixSortT(temp_storage.sort).SortBlockedToStriped(185 thread_keys,186 thread_values,187 begin_bit,188 end_bit,189 Int2Type<IS_DESCENDING>(),190 Int2Type<KEYS_ONLY>(),191 decomposer);192 193 cub::StoreDirectStriped<BLOCK_THREADS>(194 threadIdx.x, d_keys_out, thread_keys, num_items);195 196 if (!KEYS_ONLY)197 {198 cub::StoreDirectStriped<BLOCK_THREADS>(199 threadIdx.x, d_values_out, thread_values, num_items);200 }201 }202 203 __device__ __forceinline__ void ProcessIterative(int current_bit,204 int pass_bits,205 const KeyT *d_keys_in,206 const ValueT *d_values_in,207 KeyT *d_keys_out,208 ValueT *d_values_out)209 {210 // Upsweep211 BlockUpsweepT upsweep(temp_storage.upsweep,212 d_keys_in,213 current_bit,214 pass_bits,215 decomposer);216 upsweep.ProcessRegion(OffsetT{}, num_items);217 218 CTA_SYNC();219 220 // The count of each digit value in this pass (valid in the first RADIX_DIGITS threads)221 OffsetT bin_count[BINS_TRACKED_PER_THREAD];222 upsweep.ExtractCounts(bin_count);223 224 CTA_SYNC();225 226 if (IS_DESCENDING)227 {228 // Reverse bin counts229 #pragma unroll230 for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)231 {232 int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;233 234 if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))235 {236 temp_storage.unbound_sort.reverse_counts_in[bin_idx] = bin_count[track];237 }238 }239 240 CTA_SYNC();241 242 #pragma unroll243 for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)244 {245 int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;246 247 if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))248 {249 bin_count[track] = temp_storage.unbound_sort.reverse_counts_in[RADIX_DIGITS - bin_idx - 1];250 }251 }252 }253 254 // Scan255 // The global scatter base offset for each digit value in this pass256 // (valid in the first RADIX_DIGITS threads)257 OffsetT bin_offset[BINS_TRACKED_PER_THREAD];258 DigitScanT(temp_storage.unbound_sort.scan).ExclusiveSum(bin_count, bin_offset);259 260 if (IS_DESCENDING)261 {262 // Reverse bin offsets263 #pragma unroll264 for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)265 {266 int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;267 268 if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))269 {270 temp_storage.unbound_sort.reverse_counts_out[threadIdx.x] = bin_offset[track];271 }272 }273 274 CTA_SYNC();275 276 #pragma unroll277 for (int track = 0; track < BINS_TRACKED_PER_THREAD; ++track)278 {279 int bin_idx = (threadIdx.x * BINS_TRACKED_PER_THREAD) + track;280 281 if ((BLOCK_THREADS == RADIX_DIGITS) || (bin_idx < RADIX_DIGITS))282 {283 bin_offset[track] = temp_storage.unbound_sort.reverse_counts_out[RADIX_DIGITS - bin_idx - 1];284 }285 }286 }287 288 CTA_SYNC();289 290 // Downsweep291 BlockDownsweepT downsweep(temp_storage.downsweep,292 bin_offset,293 num_items,294 d_keys_in,295 d_keys_out,296 d_values_in,297 d_values_out,298 current_bit,299 pass_bits,300 decomposer);301 downsweep.ProcessRegion(OffsetT{}, num_items);302 }303};304 305 306CUB_NAMESPACE_END307 