Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
agent_segmented_radix_sort.cuh307 linesDownload Raw Back to agent
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 
codekingpro/portable-devtools · Team Ai