codekingpro/portable-devtools
114k
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 