Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
agent_segment_fixup.cuh468 linesDownload Raw Back to agent
1/******************************************************************************2 * Copyright (c) 2011, Duane Merrill.  All rights reserved.3 * Copyright (c) 2011-2018, NVIDIA CORPORATION.  All rights reserved.4 *5 * Redistribution and use in source and binary forms, with or without6 * modification, are permitted provided that the following conditions are met:7 *     * Redistributions of source code must retain the above copyright8 *       notice, this list of conditions and the following disclaimer.9 *     * Redistributions in binary form must reproduce the above copyright10 *       notice, this list of conditions and the following disclaimer in the11 *       documentation and/or other materials provided with the distribution.12 *     * Neither the name of the NVIDIA CORPORATION nor the13 *       names of its contributors may be used to endorse or promote products14 *       derived from this software without specific prior written permission.15 *16 * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND17 * ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED18 * WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE19 * DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE FOR ANY20 * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES21 * (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;22 * LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND23 * ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT24 * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS25 * SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.26 *27 ******************************************************************************/28 29/**30 * \file31 * cub::AgentSegmentFixup implements a stateful abstraction of CUDA thread blocks for participating in device-wide reduce-value-by-key.32 */33 34#pragma once35 36#include <cub/config.cuh>37 38#if defined(_CCCL_IMPLICIT_SYSTEM_HEADER_GCC)39#  pragma GCC system_header40#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_CLANG)41#  pragma clang system_header42#elif defined(_CCCL_IMPLICIT_SYSTEM_HEADER_MSVC)43#  pragma system_header44#endif // no system header45 46#include <cub/agent/single_pass_scan_operators.cuh>47#include <cub/block/block_discontinuity.cuh>48#include <cub/block/block_load.cuh>49#include <cub/block/block_scan.cuh>50#include <cub/block/block_store.cuh>51#include <cub/iterator/cache_modified_input_iterator.cuh>52#include <cub/iterator/constant_input_iterator.cuh>53 54#include <iterator>55 56CUB_NAMESPACE_BEGIN57 58 59/******************************************************************************60 * Tuning policy types61 ******************************************************************************/62 63/**64 * @brief Parameterizable tuning policy type for AgentSegmentFixup65 *66 * @tparam _BLOCK_THREADS67 *   Threads per thread block68 *69 * @tparam _ITEMS_PER_THREAD70 *   Items per thread (per tile of input)71 *72 * @tparam _LOAD_ALGORITHM73 *   The BlockLoad algorithm to use74 *75 * @tparam _LOAD_MODIFIER76 *   Cache load modifier for reading input elements77 *78 * @tparam _SCAN_ALGORITHM79 *   The BlockScan algorithm to use80 */81template <int _BLOCK_THREADS,82          int _ITEMS_PER_THREAD,83          BlockLoadAlgorithm _LOAD_ALGORITHM,84          CacheLoadModifier _LOAD_MODIFIER,85          BlockScanAlgorithm _SCAN_ALGORITHM>86struct AgentSegmentFixupPolicy87{88  enum89  {90    /// Threads per thread block91    BLOCK_THREADS = _BLOCK_THREADS,92 93    /// Items per thread (per tile of input)94    ITEMS_PER_THREAD = _ITEMS_PER_THREAD,95  };96 97  /// The BlockLoad algorithm to use98  static constexpr BlockLoadAlgorithm LOAD_ALGORITHM = _LOAD_ALGORITHM;99 100  /// Cache load modifier for reading input elements101  static constexpr CacheLoadModifier LOAD_MODIFIER = _LOAD_MODIFIER;102 103  /// The BlockScan algorithm to use104  static constexpr BlockScanAlgorithm SCAN_ALGORITHM = _SCAN_ALGORITHM;105};106 107/******************************************************************************108 * Thread block abstractions109 ******************************************************************************/110 111/**112 * @brief AgentSegmentFixup implements a stateful abstraction of CUDA thread blocks for113 * participating in device-wide reduce-value-by-key114 *115 * @tparam AgentSegmentFixupPolicyT116 *   Parameterized AgentSegmentFixupPolicy tuning policy type117 *118 * @tparam PairsInputIteratorT119 *   Random-access input iterator type for keys120 *121 * @tparam AggregatesOutputIteratorT122 *   Random-access output iterator type for values123 *124 * @tparam EqualityOpT125 *   KeyT equality operator type126 *127 * @tparam ReductionOpT128 *   ValueT reduction operator type129 *130 * @tparam OffsetT131 *   Signed integer type for global offsets132 */133template <typename AgentSegmentFixupPolicyT,134          typename PairsInputIteratorT,135          typename AggregatesOutputIteratorT,136          typename EqualityOpT,137          typename ReductionOpT,138          typename OffsetT>139struct AgentSegmentFixup140{141    //---------------------------------------------------------------------142    // Types and constants143    //---------------------------------------------------------------------144 145    // Data type of key-value input iterator146    using KeyValuePairT = cub::detail::value_t<PairsInputIteratorT>;147 148    // Value type149    using ValueT = typename KeyValuePairT::Value;150 151    // Tile status descriptor interface type152    using ScanTileStateT = ReduceByKeyScanTileState<ValueT, OffsetT>;153 154    // Constants155    enum156    {157      BLOCK_THREADS    = AgentSegmentFixupPolicyT::BLOCK_THREADS,158      ITEMS_PER_THREAD = AgentSegmentFixupPolicyT::ITEMS_PER_THREAD,159      TILE_ITEMS       = BLOCK_THREADS * ITEMS_PER_THREAD,160 161      // Whether or not do fixup using RLE + global atomics162      USE_ATOMIC_FIXUP = (std::is_same<ValueT, float>::value ||163                          std::is_same<ValueT, int>::value ||164                          std::is_same<ValueT, unsigned int>::value ||165                          std::is_same<ValueT, unsigned long long>::value),166 167      // Whether or not the scan operation has a zero-valued identity value168      // (true if we're performing addition on a primitive type)169      HAS_IDENTITY_ZERO = (std::is_same<ReductionOpT, cub::Sum>::value) &&170                          (Traits<ValueT>::PRIMITIVE),171    };172 173    // Cache-modified Input iterator wrapper type (for applying cache modifier) for keys174    // Wrap the native input pointer with CacheModifiedValuesInputIterator175    // or directly use the supplied input iterator type176    using WrappedPairsInputIteratorT = cub::detail::conditional_t<177      std::is_pointer<PairsInputIteratorT>::value,178      CacheModifiedInputIterator<AgentSegmentFixupPolicyT::LOAD_MODIFIER,179                                 KeyValuePairT,180                                 OffsetT>,181      PairsInputIteratorT>;182 183    // Cache-modified Input iterator wrapper type (for applying cache modifier) for fixup values184    // Wrap the native input pointer with CacheModifiedValuesInputIterator185    // or directly use the supplied input iterator type186    using WrappedFixupInputIteratorT = cub::detail::conditional_t<187      std::is_pointer<AggregatesOutputIteratorT>::value,188      CacheModifiedInputIterator<AgentSegmentFixupPolicyT::LOAD_MODIFIER,189                                 ValueT,190                                 OffsetT>,191      AggregatesOutputIteratorT>;192 193    // Reduce-value-by-segment scan operator194    using ReduceBySegmentOpT = ReduceByKeyOp<cub::Sum>;195 196    // Parameterized BlockLoad type for pairs197    using BlockLoadPairs = BlockLoad<KeyValuePairT,198                                     BLOCK_THREADS,199                                     ITEMS_PER_THREAD,200                                     AgentSegmentFixupPolicyT::LOAD_ALGORITHM>;201 202    // Parameterized BlockScan type203    using BlockScanT = BlockScan<KeyValuePairT,204                                 BLOCK_THREADS,205                                 AgentSegmentFixupPolicyT::SCAN_ALGORITHM>;206 207    // Callback type for obtaining tile prefix during block scan208    using TilePrefixCallbackOpT =209      TilePrefixCallbackOp<KeyValuePairT, ReduceBySegmentOpT, ScanTileStateT>;210 211    // Shared memory type for this thread block212    union _TempStorage213    {214        struct ScanStorage215        {216          // Smem needed for tile scanning217          typename BlockScanT::TempStorage scan;218 219          // Smem needed for cooperative prefix callback220          typename TilePrefixCallbackOpT::TempStorage prefix;221        } scan_storage;222 223        // Smem needed for loading keys224        typename BlockLoadPairs::TempStorage load_pairs;225    };226 227    // Alias wrapper allowing storage to be unioned228    struct TempStorage : Uninitialized<_TempStorage> {};229 230 231    //---------------------------------------------------------------------232    // Per-thread fields233    //---------------------------------------------------------------------234 235    _TempStorage &temp_storage;                   ///< Reference to temp_storage236    WrappedPairsInputIteratorT d_pairs_in;        ///< Input keys237    AggregatesOutputIteratorT d_aggregates_out;   ///< Output value aggregates238    WrappedFixupInputIteratorT d_fixup_in;        ///< Fixup input values239    InequalityWrapper<EqualityOpT> inequality_op; ///< KeyT inequality operator240    ReductionOpT reduction_op;                    ///< Reduction operator241    ReduceBySegmentOpT scan_op;                   ///< Reduce-by-segment scan operator242 243    //---------------------------------------------------------------------244    // Constructor245    //---------------------------------------------------------------------246 247    /**248     * @param temp_storage249     *   Reference to temp_storage250     *251     * @param d_pairs_in252     *   Input keys253     *254     * @param d_aggregates_out255     *   Output value aggregates256     *257     * @param equality_op258     *   KeyT equality operator259     *260     * @param reduction_op261     *   ValueT reduction operator262     */263    __device__ __forceinline__ AgentSegmentFixup(TempStorage &temp_storage,264                                                 PairsInputIteratorT d_pairs_in,265                                                 AggregatesOutputIteratorT d_aggregates_out,266                                                 EqualityOpT equality_op,267                                                 ReductionOpT reduction_op)268        : temp_storage(temp_storage.Alias())269        , d_pairs_in(d_pairs_in)270        , d_aggregates_out(d_aggregates_out)271        , d_fixup_in(d_aggregates_out)272        , inequality_op(equality_op)273        , reduction_op(reduction_op)274        , scan_op(reduction_op)275    {}276 277    //---------------------------------------------------------------------278    // Cooperatively scan a device-wide sequence of tiles with other CTAs279    //---------------------------------------------------------------------280 281 282    /**283     * @brief Process input tile. Specialized for atomic-fixup284     *285     * @param num_remaining286     *   Number of global input items remaining (including this tile)287     *288     * @param tile_idx289     *   Tile index290     *291     * @param tile_offset292     *   Tile offset293     *294     * @param tile_state295     *   Global tile state descriptor296     *297     * @param use_atomic_fixup298     *   Marker whether to use atomicAdd (instead of reduce-by-key)299     */300    template <bool IS_LAST_TILE>301    __device__ __forceinline__ void ConsumeTile(OffsetT num_remaining,302                                                int tile_idx,303                                                OffsetT tile_offset,304                                                ScanTileStateT &tile_state,305                                                Int2Type<true> use_atomic_fixup)306    {307        KeyValuePairT   pairs[ITEMS_PER_THREAD];308 309        // Load pairs310        KeyValuePairT oob_pair;311        oob_pair.key = -1;312 313        if (IS_LAST_TILE)314            BlockLoadPairs(temp_storage.load_pairs).Load(d_pairs_in + tile_offset, pairs, num_remaining, oob_pair);315        else316            BlockLoadPairs(temp_storage.load_pairs).Load(d_pairs_in + tile_offset, pairs);317 318        // RLE319        #pragma unroll320        for (int ITEM = 1; ITEM < ITEMS_PER_THREAD; ++ITEM)321        {322            ValueT* d_scatter = d_aggregates_out + pairs[ITEM - 1].key;323            if (pairs[ITEM].key != pairs[ITEM - 1].key)324                atomicAdd(d_scatter, pairs[ITEM - 1].value);325            else326                pairs[ITEM].value = reduction_op(pairs[ITEM - 1].value, pairs[ITEM].value);327        }328 329        // Flush last item if valid330        ValueT* d_scatter = d_aggregates_out + pairs[ITEMS_PER_THREAD - 1].key;331        if ((!IS_LAST_TILE) || (pairs[ITEMS_PER_THREAD - 1].key >= 0))332            atomicAdd(d_scatter, pairs[ITEMS_PER_THREAD - 1].value);333    }334 335    /**336     * @brief Process input tile. Specialized for reduce-by-key fixup337     *338     * @param num_remaining339     *   Number of global input items remaining (including this tile)340     *341     * @param tile_idx342     *   Tile index343     *344     * @param tile_offset345     *   Tile offset346     *347     * @param tile_state348     *   Global tile state descriptor349     *350     * @param use_atomic_fixup351     *   Marker whether to use atomicAdd (instead of reduce-by-key)352     */353    template <bool IS_LAST_TILE>354    __device__ __forceinline__ void ConsumeTile(OffsetT num_remaining,355                                                int tile_idx,356                                                OffsetT tile_offset,357                                                ScanTileStateT &tile_state,358                                                Int2Type<false> use_atomic_fixup)359    {360        KeyValuePairT   pairs[ITEMS_PER_THREAD];361        KeyValuePairT   scatter_pairs[ITEMS_PER_THREAD];362 363        // Load pairs364        KeyValuePairT oob_pair;365        oob_pair.key = -1;366 367        if (IS_LAST_TILE)368            BlockLoadPairs(temp_storage.load_pairs).Load(d_pairs_in + tile_offset, pairs, num_remaining, oob_pair);369        else370            BlockLoadPairs(temp_storage.load_pairs).Load(d_pairs_in + tile_offset, pairs);371 372        CTA_SYNC();373 374        KeyValuePairT tile_aggregate;375        if (tile_idx == 0)376        {377            // Exclusive scan of values and segment_flags378            BlockScanT(temp_storage.scan_storage.scan).ExclusiveScan(pairs, scatter_pairs, scan_op, tile_aggregate);379 380            // Update tile status if this is not the last tile381            if (threadIdx.x == 0)382            {383                // Set first segment id to not trigger a flush (invalid from exclusive scan)384                scatter_pairs[0].key = pairs[0].key;385 386                if (!IS_LAST_TILE)387                    tile_state.SetInclusive(0, tile_aggregate);388 389            }390        }391        else392        {393            // Exclusive scan of values and segment_flags394            TilePrefixCallbackOpT prefix_op(tile_state, temp_storage.scan_storage.prefix, scan_op, tile_idx);395            BlockScanT(temp_storage.scan_storage.scan).ExclusiveScan(pairs, scatter_pairs, scan_op, prefix_op);396            tile_aggregate = prefix_op.GetBlockAggregate();397        }398 399        // Scatter updated values400        #pragma unroll401        for (int ITEM = 0; ITEM < ITEMS_PER_THREAD; ++ITEM)402        {403            if (scatter_pairs[ITEM].key != pairs[ITEM].key)404            {405                // Update the value at the key location406                ValueT value    = d_fixup_in[scatter_pairs[ITEM].key];407                value           = reduction_op(value, scatter_pairs[ITEM].value);408 409                d_aggregates_out[scatter_pairs[ITEM].key] = value;410            }411        }412 413        // Finalize the last item414        if (IS_LAST_TILE)415        {416            // Last thread will output final count and last item, if necessary417            if (threadIdx.x == BLOCK_THREADS - 1)418            {419                // If the last tile is a whole tile, the inclusive prefix contains accumulated value reduction for the last segment420                if (num_remaining == TILE_ITEMS)421                {422                    // Update the value at the key location423                    OffsetT last_key = pairs[ITEMS_PER_THREAD - 1].key;424                    d_aggregates_out[last_key] = reduction_op(tile_aggregate.value, d_fixup_in[last_key]);425                }426            }427        }428    }429 430    /**431     * @brief Scan tiles of items as part of a dynamic chained scan432     *433     * @param num_items434     *   Total number of input items435     *436     * @param num_tiles437     *   Total number of input tiles438     *439     * @param tile_state440     *   Global tile state descriptor441     */442    __device__ __forceinline__ void ConsumeRange(OffsetT num_items,443                                                 int num_tiles,444                                                 ScanTileStateT &tile_state)445    {446        // Blocks are launched in increasing order, so just assign one tile per block447        int tile_idx          = (blockIdx.x * gridDim.y) + blockIdx.y; // Current tile index448        OffsetT tile_offset   = tile_idx * TILE_ITEMS;   // Global offset for the current tile449        OffsetT num_remaining = num_items - tile_offset; // Remaining items (including this tile)450 451        if (num_remaining > TILE_ITEMS)452        {453            // Not the last tile (full)454            ConsumeTile<false>(num_remaining, tile_idx, tile_offset, tile_state, Int2Type<USE_ATOMIC_FIXUP>());455        }456        else if (num_remaining > 0)457        {458            // The last tile (possibly partially-full)459            ConsumeTile<true>(num_remaining, tile_idx, tile_offset, tile_state, Int2Type<USE_ATOMIC_FIXUP>());460        }461    }462 463};464 465 466CUB_NAMESPACE_END467 468 
codekingpro/portable-devtools · Team Ai