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 * Callback operator types for supplying BlockScan prefixes32 */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 <iterator>47 48#include <cub/detail/strong_load.cuh>49#include <cub/detail/strong_store.cuh>50#include <cub/detail/uninitialized_copy.cuh>51#include <cub/thread/thread_load.cuh>52#include <cub/thread/thread_store.cuh>53#include <cub/warp/warp_reduce.cuh>54#include <cub/util_temporary_storage.cuh>55 56#include <nv/target>57 58CUB_NAMESPACE_BEGIN59 60 61/******************************************************************************62 * Prefix functor type for maintaining a running prefix while scanning a63 * region independent of other thread blocks64 ******************************************************************************/65 66/**67 * Stateful callback operator type for supplying BlockScan prefixes.68 * Maintains a running prefix that can be applied to consecutive69 * BlockScan operations.70 *71 * @tparam T72 * BlockScan value type73 *74 * @tparam ScanOpT75 * Wrapped scan operator type76 */77template <typename T, typename ScanOpT>78struct BlockScanRunningPrefixOp79{80 /// Wrapped scan operator81 ScanOpT op;82 83 /// Running block-wide prefix84 T running_total;85 86 /// Constructor87 __device__ __forceinline__ BlockScanRunningPrefixOp(ScanOpT op)88 : op(op)89 {}90 91 /// Constructor92 __device__ __forceinline__ BlockScanRunningPrefixOp(T starting_prefix, ScanOpT op)93 : op(op)94 , running_total(starting_prefix)95 {}96 97 /**98 * Prefix callback operator. Returns the block-wide running_total in thread-0.99 *100 * @param block_aggregate101 * The aggregate sum of the BlockScan inputs102 */103 __device__ __forceinline__ T operator()(const T &block_aggregate)104 {105 T retval = running_total;106 running_total = op(running_total, block_aggregate);107 return retval;108 }109};110 111/******************************************************************************112 * Generic tile status interface types for block-cooperative scans113 ******************************************************************************/114 115/**116 * Enumerations of tile status117 */118enum ScanTileStatus119{120 SCAN_TILE_OOB, // Out-of-bounds (e.g., padding)121 SCAN_TILE_INVALID = 99, // Not yet processed122 SCAN_TILE_PARTIAL, // Tile aggregate is available123 SCAN_TILE_INCLUSIVE, // Inclusive tile prefix is available124};125 126namespace detail127{128 129template <int Delay, unsigned int GridThreshold = 500>130__device__ __forceinline__ void delay()131{132 NV_IF_TARGET(NV_PROVIDES_SM_70,133 (if (Delay > 0)134 {135 if (gridDim.x < GridThreshold)136 {137 __threadfence_block();138 }139 else140 {141 __nanosleep(Delay);142 }143 }));144}145 146template <unsigned int GridThreshold = 500>147__device__ __forceinline__ void delay(int ns)148{149 NV_IF_TARGET(NV_PROVIDES_SM_70,150 (if (ns > 0)151 {152 if (gridDim.x < GridThreshold)153 {154 __threadfence_block();155 }156 else157 {158 __nanosleep(ns);159 }160 }));161}162 163template <int Delay>164__device__ __forceinline__ void always_delay()165{166 NV_IF_TARGET(NV_PROVIDES_SM_70, (__nanosleep(Delay);));167}168 169__device__ __forceinline__ void always_delay(int ns)170{171 NV_IF_TARGET(NV_PROVIDES_SM_70, (__nanosleep(ns);), ((void)ns;));172}173 174template <unsigned int Delay = 350, unsigned int GridThreshold = 500>175__device__ __forceinline__ void delay_or_prevent_hoisting()176{177 NV_IF_TARGET(NV_PROVIDES_SM_70,178 (delay<Delay, GridThreshold>();),179 (__threadfence_block();));180}181 182template <unsigned int GridThreshold = 500>183__device__ __forceinline__ void delay_or_prevent_hoisting(int ns)184{185 NV_IF_TARGET(NV_PROVIDES_SM_70,186 (delay<GridThreshold>(ns);),187 ((void)ns; __threadfence_block();));188}189 190template <unsigned int Delay = 350>191__device__ __forceinline__ void always_delay_or_prevent_hoisting()192{193 NV_IF_TARGET(NV_PROVIDES_SM_70,194 (always_delay(Delay);),195 (__threadfence_block();));196}197 198__device__ __forceinline__ void always_delay_or_prevent_hoisting(int ns)199{200 NV_IF_TARGET(NV_PROVIDES_SM_70,201 (always_delay(ns);),202 ((void)ns; __threadfence_block();));203}204 205template <unsigned int L2WriteLatency>206struct no_delay_constructor_t207{208 struct delay_t209 {210 __device__ __forceinline__ void operator()()211 {212 NV_IF_TARGET(NV_PROVIDES_SM_70,213 (),214 (__threadfence_block();));215 }216 };217 218 __device__ __forceinline__ no_delay_constructor_t(unsigned int /* seed */)219 {220 delay<L2WriteLatency>();221 }222 223 __device__ __forceinline__ delay_t operator()() { return {}; }224};225 226template <unsigned int Delay, unsigned int L2WriteLatency, unsigned int GridThreshold = 500>227struct reduce_by_key_delay_constructor_t228{229 struct delay_t230 {231 __device__ __forceinline__ void operator()()232 {233 NV_DISPATCH_TARGET(234 NV_IS_EXACTLY_SM_80, (delay<Delay, GridThreshold>();),235 NV_PROVIDES_SM_70, (delay< 0, GridThreshold>();),236 NV_IS_DEVICE, (__threadfence_block();));237 }238 };239 240 __device__ __forceinline__ reduce_by_key_delay_constructor_t(unsigned int /* seed */)241 {242 delay<L2WriteLatency>();243 }244 245 __device__ __forceinline__ delay_t operator()() { return {}; }246};247 248template <unsigned int Delay, unsigned int L2WriteLatency>249struct fixed_delay_constructor_t250{251 struct delay_t252 {253 __device__ __forceinline__ void operator()() { delay_or_prevent_hoisting<Delay>(); }254 };255 256 __device__ __forceinline__ fixed_delay_constructor_t(unsigned int /* seed */)257 {258 delay<L2WriteLatency>();259 }260 261 __device__ __forceinline__ delay_t operator()() { return {}; }262};263 264template <unsigned int InitialDelay, unsigned int L2WriteLatency>265struct exponential_backoff_constructor_t266{267 struct delay_t268 {269 int delay;270 271 __device__ __forceinline__ void operator()()272 {273 always_delay_or_prevent_hoisting(delay);274 delay <<= 1;275 }276 };277 278 __device__ __forceinline__ exponential_backoff_constructor_t(unsigned int /* seed */)279 {280 always_delay<L2WriteLatency>();281 }282 283 __device__ __forceinline__ delay_t operator()() { return {InitialDelay}; }284};285 286template <unsigned int InitialDelay, unsigned int L2WriteLatency>287struct exponential_backoff_jitter_constructor_t288{289 struct delay_t290 {291 static constexpr unsigned int a = 16807;292 static constexpr unsigned int c = 0;293 static constexpr unsigned int m = 1u << 31;294 295 unsigned int max_delay;296 unsigned int &seed;297 298 __device__ __forceinline__ unsigned int next(unsigned int min, unsigned int max)299 {300 return (seed = (a * seed + c) % m) % (max + 1 - min) + min;301 }302 303 __device__ __forceinline__ void operator()()304 {305 always_delay_or_prevent_hoisting(next(0, max_delay));306 max_delay <<= 1;307 }308 };309 310 unsigned int seed;311 312 __device__ __forceinline__ exponential_backoff_jitter_constructor_t(unsigned int seed)313 : seed(seed)314 {315 always_delay<L2WriteLatency>();316 }317 318 __device__ __forceinline__ delay_t operator()() { return {InitialDelay, seed}; }319};320 321template <unsigned int InitialDelay, unsigned int L2WriteLatency>322struct exponential_backoff_jitter_window_constructor_t323{324 struct delay_t325 {326 static constexpr unsigned int a = 16807;327 static constexpr unsigned int c = 0;328 static constexpr unsigned int m = 1u << 31;329 330 unsigned int max_delay;331 unsigned int &seed;332 333 __device__ __forceinline__ unsigned int next(unsigned int min, unsigned int max)334 {335 return (seed = (a * seed + c) % m) % (max + 1 - min) + min;336 }337 338 __device__ __forceinline__ void operator()()339 {340 unsigned int next_max_delay = max_delay << 1;341 always_delay_or_prevent_hoisting(next(max_delay, next_max_delay));342 max_delay = next_max_delay;343 }344 };345 346 unsigned int seed;347 __device__ __forceinline__ exponential_backoff_jitter_window_constructor_t(unsigned int seed)348 : seed(seed)349 {350 always_delay<L2WriteLatency>();351 }352 353 __device__ __forceinline__ delay_t operator()() { return {InitialDelay, seed}; }354};355 356template <unsigned int InitialDelay, unsigned int L2WriteLatency>357struct exponential_backon_jitter_window_constructor_t358{359 struct delay_t360 {361 static constexpr unsigned int a = 16807;362 static constexpr unsigned int c = 0;363 static constexpr unsigned int m = 1u << 31;364 365 unsigned int max_delay;366 unsigned int &seed;367 368 __device__ __forceinline__ unsigned int next(unsigned int min, unsigned int max)369 {370 return (seed = (a * seed + c) % m) % (max + 1 - min) + min;371 }372 373 __device__ __forceinline__ void operator()()374 {375 int prev_delay = max_delay >> 1;376 always_delay_or_prevent_hoisting(next(prev_delay, max_delay));377 max_delay = prev_delay;378 }379 };380 381 unsigned int seed;382 unsigned int max_delay = InitialDelay;383 384 __device__ __forceinline__ exponential_backon_jitter_window_constructor_t(unsigned int seed)385 : seed(seed)386 {387 always_delay<L2WriteLatency>();388 }389 390 __device__ __forceinline__ delay_t operator()()391 {392 max_delay >>= 1;393 return {max_delay, seed};394 }395};396 397template <unsigned int InitialDelay, unsigned int L2WriteLatency>398struct exponential_backon_jitter_constructor_t399{400 struct delay_t401 {402 static constexpr unsigned int a = 16807;403 static constexpr unsigned int c = 0;404 static constexpr unsigned int m = 1u << 31;405 406 unsigned int max_delay;407 unsigned int &seed;408 409 __device__ __forceinline__ unsigned int next(unsigned int min, unsigned int max)410 {411 return (seed = (a * seed + c) % m) % (max + 1 - min) + min;412 }413 414 __device__ __forceinline__ void operator()()415 {416 always_delay_or_prevent_hoisting(next(0, max_delay));417 max_delay >>= 1;418 }419 };420 421 unsigned int seed;422 unsigned int max_delay = InitialDelay;423 424 __device__ __forceinline__ exponential_backon_jitter_constructor_t(unsigned int seed)425 : seed(seed)426 {427 always_delay<L2WriteLatency>();428 }429 430 __device__ __forceinline__ delay_t operator()()431 {432 max_delay >>= 1;433 return {max_delay, seed};434 }435};436 437template <unsigned int InitialDelay, unsigned int L2WriteLatency>438struct exponential_backon_constructor_t439{440 struct delay_t441 {442 unsigned int delay;443 444 __device__ __forceinline__ void operator()()445 {446 always_delay_or_prevent_hoisting(delay);447 delay >>= 1;448 }449 };450 451 unsigned int max_delay = InitialDelay;452 453 __device__ __forceinline__ exponential_backon_constructor_t(unsigned int /* seed */)454 {455 always_delay<L2WriteLatency>();456 }457 458 __device__ __forceinline__ delay_t operator()()459 {460 max_delay >>= 1;461 return {max_delay};462 }463};464 465using default_no_delay_constructor_t = no_delay_constructor_t<450>;466using default_no_delay_t = default_no_delay_constructor_t::delay_t;467 468template <class T>469using default_delay_constructor_t = cub::detail::conditional_t<Traits<T>::PRIMITIVE,470 fixed_delay_constructor_t<350, 450>,471 default_no_delay_constructor_t>;472 473template <class T>474using default_delay_t = typename default_delay_constructor_t<T>::delay_t;475 476template <class KeyT, class ValueT>477using default_reduce_by_key_delay_constructor_t =478 detail::conditional_t<(Traits<ValueT>::PRIMITIVE) && (sizeof(ValueT) + sizeof(KeyT) < 16),479 reduce_by_key_delay_constructor_t<350, 450>,480 default_delay_constructor_t<KeyValuePair<KeyT, ValueT>>>;481}482 483/**484 * Tile status interface.485 */486template <487 typename T,488 bool SINGLE_WORD = Traits<T>::PRIMITIVE>489struct ScanTileState;490 491 492/**493 * Tile status interface specialized for scan status and value types494 * that can be combined into one machine word that can be495 * read/written coherently in a single access.496 */497template <typename T>498struct ScanTileState<T, true>499{500 // Status word type501 using StatusWord = cub::detail::conditional_t<502 sizeof(T) == 8,503 unsigned long long,504 cub::detail::conditional_t<505 sizeof(T) == 4,506 unsigned int,507 cub::detail::conditional_t<sizeof(T) == 2, unsigned short, unsigned char>>>;508 509 // Unit word type510 using TxnWord = cub::detail::conditional_t<511 sizeof(T) == 8,512 ulonglong2,513 cub::detail::conditional_t<514 sizeof(T) == 4,515 uint2,516 unsigned int>>;517 518 // Device word type519 struct TileDescriptor520 {521 StatusWord status;522 T value;523 };524 525 526 // Constants527 enum528 {529 TILE_STATUS_PADDING = CUB_PTX_WARP_THREADS,530 };531 532 533 // Device storage534 TxnWord *d_tile_descriptors;535 536 /// Constructor537 __host__ __device__ __forceinline__538 ScanTileState()539 :540 d_tile_descriptors(NULL)541 {}542 543 /**544 * @brief Initializer545 *546 * @param[in] num_tiles547 * Number of tiles548 *549 * @param[in] d_temp_storage550 * Device-accessible allocation of temporary storage.551 * When NULL, the required allocation size is written to \p temp_storage_bytes and no work is552 * done.553 *554 * @param[in] temp_storage_bytes555 * Size in bytes of \t d_temp_storage allocation556 */557 __host__ __device__ __forceinline__ cudaError_t Init(int /*num_tiles*/,558 void *d_temp_storage,559 size_t /*temp_storage_bytes*/)560 {561 d_tile_descriptors = reinterpret_cast<TxnWord *>(d_temp_storage);562 return cudaSuccess;563 }564 565 /**566 * @brief Compute device memory needed for tile status567 *568 * @param[in] num_tiles569 * Number of tiles570 *571 * @param[out] temp_storage_bytes572 * Size in bytes of \t d_temp_storage allocation573 */574 __host__ __device__ __forceinline__ static cudaError_t575 AllocationSize(int num_tiles, size_t &temp_storage_bytes)576 {577 // bytes needed for tile status descriptors578 temp_storage_bytes = (num_tiles + TILE_STATUS_PADDING) * sizeof(TxnWord);579 return cudaSuccess;580 }581 582 /**583 * Initialize (from device)584 */585 __device__ __forceinline__ void InitializeStatus(int num_tiles)586 {587 int tile_idx = (blockIdx.x * blockDim.x) + threadIdx.x;588 589 TxnWord val = TxnWord();590 TileDescriptor *descriptor = reinterpret_cast<TileDescriptor*>(&val);591 592 if (tile_idx < num_tiles)593 {594 // Not-yet-set595 descriptor->status = StatusWord(SCAN_TILE_INVALID);596 d_tile_descriptors[TILE_STATUS_PADDING + tile_idx] = val;597 }598 599 if ((blockIdx.x == 0) && (threadIdx.x < TILE_STATUS_PADDING))600 {601 // Padding602 descriptor->status = StatusWord(SCAN_TILE_OOB);603 d_tile_descriptors[threadIdx.x] = val;604 }605 }606 607 608 /**609 * Update the specified tile's inclusive value and corresponding status610 */611 __device__ __forceinline__ void SetInclusive(int tile_idx, T tile_inclusive)612 {613 TileDescriptor tile_descriptor;614 tile_descriptor.status = SCAN_TILE_INCLUSIVE;615 tile_descriptor.value = tile_inclusive;616 617 TxnWord alias;618 *reinterpret_cast<TileDescriptor*>(&alias) = tile_descriptor;619 620 detail::store_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx, alias);621 }622 623 624 /**625 * Update the specified tile's partial value and corresponding status626 */627 __device__ __forceinline__ void SetPartial(int tile_idx, T tile_partial)628 {629 TileDescriptor tile_descriptor;630 tile_descriptor.status = SCAN_TILE_PARTIAL;631 tile_descriptor.value = tile_partial;632 633 TxnWord alias;634 *reinterpret_cast<TileDescriptor*>(&alias) = tile_descriptor;635 636 detail::store_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx, alias);637 }638 639 /**640 * Wait for the corresponding tile to become non-invalid641 */642 template <class DelayT = detail::default_delay_t<T>>643 __device__ __forceinline__ void WaitForValid(644 int tile_idx,645 StatusWord &status,646 T &value,647 DelayT delay_or_prevent_hoisting = {})648 {649 TileDescriptor tile_descriptor;650 651 {652 TxnWord alias = detail::load_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx);653 tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);654 }655 656 while (WARP_ANY((tile_descriptor.status == SCAN_TILE_INVALID), 0xffffffff))657 {658 delay_or_prevent_hoisting();659 TxnWord alias = detail::load_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx);660 tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);661 }662 663 status = tile_descriptor.status;664 value = tile_descriptor.value;665 }666 667 /**668 * Loads and returns the tile's value. The returned value is undefined if either (a) the tile's status is invalid or669 * (b) there is no memory fence between reading a non-invalid status and the call to LoadValid.670 */671 __device__ __forceinline__ T LoadValid(int tile_idx)672 {673 TxnWord alias = d_tile_descriptors[TILE_STATUS_PADDING + tile_idx];674 TileDescriptor tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);675 return tile_descriptor.value;676 }677};678 679 680 681/**682 * Tile status interface specialized for scan status and value types that683 * cannot be combined into one machine word.684 */685template <typename T>686struct ScanTileState<T, false>687{688 // Status word type689 using StatusWord = unsigned int;690 691 // Constants692 enum693 {694 TILE_STATUS_PADDING = CUB_PTX_WARP_THREADS,695 };696 697 // Device storage698 StatusWord *d_tile_status;699 T *d_tile_partial;700 T *d_tile_inclusive;701 702 /// Constructor703 __host__ __device__ __forceinline__704 ScanTileState()705 :706 d_tile_status(NULL),707 d_tile_partial(NULL),708 d_tile_inclusive(NULL)709 {}710 711 /**712 * @brief Initializer713 *714 * @param[in] num_tiles715 * Number of tiles716 *717 * @param[in] d_temp_storage718 * Device-accessible allocation of temporary storage.719 * When NULL, the required allocation size is written to \p temp_storage_bytes and no work is720 * done.721 *722 * @param[in] temp_storage_bytes723 * Size in bytes of \t d_temp_storage allocation724 */725 /// Initializer726 __host__ __device__ __forceinline__ cudaError_t Init(int num_tiles,727 void *d_temp_storage,728 size_t temp_storage_bytes)729 {730 cudaError_t error = cudaSuccess;731 do732 {733 void* allocations[3] = {};734 size_t allocation_sizes[3];735 736 // bytes needed for tile status descriptors737 allocation_sizes[0] = (num_tiles + TILE_STATUS_PADDING) * sizeof(StatusWord);738 739 // bytes needed for partials740 allocation_sizes[1] = (num_tiles + TILE_STATUS_PADDING) * sizeof(Uninitialized<T>);741 742 // bytes needed for inclusives743 allocation_sizes[2] = (num_tiles + TILE_STATUS_PADDING) * sizeof(Uninitialized<T>);744 745 // Compute allocation pointers into the single storage blob746 error = CubDebug(747 AliasTemporaries(d_temp_storage, temp_storage_bytes, allocations, allocation_sizes));748 749 if (cudaSuccess != error)750 {751 break;752 }753 754 // Alias the offsets755 d_tile_status = reinterpret_cast<StatusWord*>(allocations[0]);756 d_tile_partial = reinterpret_cast<T*>(allocations[1]);757 d_tile_inclusive = reinterpret_cast<T*>(allocations[2]);758 }759 while (0);760 761 return error;762 }763 764 /**765 * @brief Compute device memory needed for tile status766 *767 * @param[in] num_tiles768 * Number of tiles769 *770 * @param[out] temp_storage_bytes771 * Size in bytes of \t d_temp_storage allocation772 */773 __host__ __device__ __forceinline__ static cudaError_t774 AllocationSize(int num_tiles, size_t &temp_storage_bytes)775 {776 // Specify storage allocation requirements777 size_t allocation_sizes[3];778 779 // bytes needed for tile status descriptors780 allocation_sizes[0] = (num_tiles + TILE_STATUS_PADDING) * sizeof(StatusWord);781 782 // bytes needed for partials783 allocation_sizes[1] = (num_tiles + TILE_STATUS_PADDING) * sizeof(Uninitialized<T>);784 785 // bytes needed for inclusives786 allocation_sizes[2] = (num_tiles + TILE_STATUS_PADDING) * sizeof(Uninitialized<T>);787 788 // Set the necessary size of the blob789 void* allocations[3] = {};790 return CubDebug(AliasTemporaries(NULL, temp_storage_bytes, allocations, allocation_sizes));791 }792 793 794 /**795 * Initialize (from device)796 */797 __device__ __forceinline__ void InitializeStatus(int num_tiles)798 {799 int tile_idx = (blockIdx.x * blockDim.x) + threadIdx.x;800 if (tile_idx < num_tiles)801 {802 // Not-yet-set803 d_tile_status[TILE_STATUS_PADDING + tile_idx] = StatusWord(SCAN_TILE_INVALID);804 }805 806 if ((blockIdx.x == 0) && (threadIdx.x < TILE_STATUS_PADDING))807 {808 // Padding809 d_tile_status[threadIdx.x] = StatusWord(SCAN_TILE_OOB);810 }811 }812 813 814 /**815 * Update the specified tile's inclusive value and corresponding status816 */817 __device__ __forceinline__ void SetInclusive(int tile_idx, T tile_inclusive)818 {819 // Update tile inclusive value820 ThreadStore<STORE_CG>(d_tile_inclusive + TILE_STATUS_PADDING + tile_idx, tile_inclusive);821 detail::store_release(d_tile_status + TILE_STATUS_PADDING + tile_idx, StatusWord(SCAN_TILE_INCLUSIVE));822 }823 824 825 /**826 * Update the specified tile's partial value and corresponding status827 */828 __device__ __forceinline__ void SetPartial(int tile_idx, T tile_partial)829 {830 // Update tile partial value831 ThreadStore<STORE_CG>(d_tile_partial + TILE_STATUS_PADDING + tile_idx, tile_partial);832 detail::store_release(d_tile_status + TILE_STATUS_PADDING + tile_idx, StatusWord(SCAN_TILE_PARTIAL));833 }834 835 /**836 * Wait for the corresponding tile to become non-invalid837 */838 template <class DelayT = detail::default_no_delay_t>839 __device__ __forceinline__ void WaitForValid(840 int tile_idx,841 StatusWord &status,842 T &value,843 DelayT delay = {})844 {845 do846 {847 delay();848 status = detail::load_relaxed(d_tile_status + TILE_STATUS_PADDING + tile_idx);849 __threadfence();850 } while (WARP_ANY((status == SCAN_TILE_INVALID), 0xffffffff));851 852 if (status == StatusWord(SCAN_TILE_PARTIAL))853 {854 value = ThreadLoad<LOAD_CG>(d_tile_partial + TILE_STATUS_PADDING + tile_idx);855 }856 else857 {858 value = ThreadLoad<LOAD_CG>(d_tile_inclusive + TILE_STATUS_PADDING + tile_idx);859 }860 }861 862 /**863 * Loads and returns the tile's value. The returned value is undefined if either (a) the tile's status is invalid or864 * (b) there is no memory fence between reading a non-invalid status and the call to LoadValid.865 */866 __device__ __forceinline__ T LoadValid(int tile_idx)867 {868 return d_tile_inclusive[TILE_STATUS_PADDING + tile_idx];869 }870};871 872 873/******************************************************************************874 * ReduceByKey tile status interface types for block-cooperative scans875 ******************************************************************************/876 877/**878 * Tile status interface for reduction by key.879 *880 */881template <882 typename ValueT,883 typename KeyT,884 bool SINGLE_WORD = (Traits<ValueT>::PRIMITIVE) && (sizeof(ValueT) + sizeof(KeyT) < 16)>885struct ReduceByKeyScanTileState;886 887 888/**889 * Tile status interface for reduction by key, specialized for scan status and value types that890 * cannot be combined into one machine word.891 */892template <893 typename ValueT,894 typename KeyT>895struct ReduceByKeyScanTileState<ValueT, KeyT, false> :896 ScanTileState<KeyValuePair<KeyT, ValueT> >897{898 typedef ScanTileState<KeyValuePair<KeyT, ValueT> > SuperClass;899 900 /// Constructor901 __host__ __device__ __forceinline__902 ReduceByKeyScanTileState() : SuperClass() {}903};904 905 906/**907 * Tile status interface for reduction by key, specialized for scan status and value types that908 * can be combined into one machine word that can be read/written coherently in a single access.909 */910template <911 typename ValueT,912 typename KeyT>913struct ReduceByKeyScanTileState<ValueT, KeyT, true>914{915 using KeyValuePairT = KeyValuePair<KeyT, ValueT>;916 917 // Constants918 enum919 {920 PAIR_SIZE = static_cast<int>(sizeof(ValueT) + sizeof(KeyT)),921 TXN_WORD_SIZE = 1 << Log2<PAIR_SIZE + 1>::VALUE,922 STATUS_WORD_SIZE = TXN_WORD_SIZE - PAIR_SIZE,923 924 TILE_STATUS_PADDING = CUB_PTX_WARP_THREADS,925 };926 927 // Status word type928 using StatusWord = cub::detail::conditional_t<929 STATUS_WORD_SIZE == 8,930 unsigned long long,931 cub::detail::conditional_t<932 STATUS_WORD_SIZE == 4,933 unsigned int,934 cub::detail::conditional_t<STATUS_WORD_SIZE == 2, unsigned short, unsigned char>>>;935 936 // Status word type937 using TxnWord = cub::detail::conditional_t<938 TXN_WORD_SIZE == 16,939 ulonglong2,940 cub::detail::conditional_t<TXN_WORD_SIZE == 8, unsigned long long, unsigned int>>;941 942 // Device word type (for when sizeof(ValueT) == sizeof(KeyT))943 struct TileDescriptorBigStatus944 {945 KeyT key;946 ValueT value;947 StatusWord status;948 };949 950 // Device word type (for when sizeof(ValueT) != sizeof(KeyT))951 struct TileDescriptorLittleStatus952 {953 ValueT value;954 StatusWord status;955 KeyT key;956 };957 958 // Device word type959 using TileDescriptor =960 cub::detail::conditional_t<sizeof(ValueT) == sizeof(KeyT),961 TileDescriptorBigStatus,962 TileDescriptorLittleStatus>;963 964 // Device storage965 TxnWord *d_tile_descriptors;966 967 968 /// Constructor969 __host__ __device__ __forceinline__970 ReduceByKeyScanTileState()971 :972 d_tile_descriptors(NULL)973 {}974 975 /**976 * @brief Initializer977 *978 * @param[in] num_tiles979 * Number of tiles980 *981 * @param[in] d_temp_storage982 * Device-accessible allocation of temporary storage. When NULL, the required allocation size983 * is written to \p temp_storage_bytes and no work is done.984 *985 * @param[in] temp_storage_bytes986 * Size in bytes of \t d_temp_storage allocation987 */988 __host__ __device__ __forceinline__ cudaError_t Init(int /*num_tiles*/,989 void *d_temp_storage,990 size_t /*temp_storage_bytes*/)991 {992 d_tile_descriptors = reinterpret_cast<TxnWord *>(d_temp_storage);993 return cudaSuccess;994 }995 996 /**997 * @brief Compute device memory needed for tile status998 *999 * @param[in] num_tiles1000 * Number of tiles1001 *1002 * @param[out] temp_storage_bytes1003 * Size in bytes of \t d_temp_storage allocation1004 */1005 __host__ __device__ __forceinline__ static cudaError_t1006 AllocationSize(int num_tiles, size_t &temp_storage_bytes)1007 {1008 // bytes needed for tile status descriptors1009 temp_storage_bytes = (num_tiles + TILE_STATUS_PADDING) * sizeof(TxnWord);1010 return cudaSuccess;1011 }1012 1013 /**1014 * Initialize (from device)1015 */1016 __device__ __forceinline__ void InitializeStatus(int num_tiles)1017 {1018 int tile_idx = (blockIdx.x * blockDim.x) + threadIdx.x;1019 TxnWord val = TxnWord();1020 TileDescriptor *descriptor = reinterpret_cast<TileDescriptor*>(&val);1021 1022 if (tile_idx < num_tiles)1023 {1024 // Not-yet-set1025 descriptor->status = StatusWord(SCAN_TILE_INVALID);1026 d_tile_descriptors[TILE_STATUS_PADDING + tile_idx] = val;1027 }1028 1029 if ((blockIdx.x == 0) && (threadIdx.x < TILE_STATUS_PADDING))1030 {1031 // Padding1032 descriptor->status = StatusWord(SCAN_TILE_OOB);1033 d_tile_descriptors[threadIdx.x] = val;1034 }1035 }1036 1037 1038 /**1039 * Update the specified tile's inclusive value and corresponding status1040 */1041 __device__ __forceinline__ void SetInclusive(int tile_idx, KeyValuePairT tile_inclusive)1042 {1043 TileDescriptor tile_descriptor;1044 tile_descriptor.status = SCAN_TILE_INCLUSIVE;1045 tile_descriptor.value = tile_inclusive.value;1046 tile_descriptor.key = tile_inclusive.key;1047 1048 TxnWord alias;1049 *reinterpret_cast<TileDescriptor*>(&alias) = tile_descriptor;1050 1051 detail::store_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx, alias);1052 }1053 1054 1055 /**1056 * Update the specified tile's partial value and corresponding status1057 */1058 __device__ __forceinline__ void SetPartial(int tile_idx, KeyValuePairT tile_partial)1059 {1060 TileDescriptor tile_descriptor;1061 tile_descriptor.status = SCAN_TILE_PARTIAL;1062 tile_descriptor.value = tile_partial.value;1063 tile_descriptor.key = tile_partial.key;1064 1065 TxnWord alias;1066 *reinterpret_cast<TileDescriptor*>(&alias) = tile_descriptor;1067 1068 detail::store_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx, alias);1069 }1070 1071 /**1072 * Wait for the corresponding tile to become non-invalid1073 */1074 template <class DelayT = detail::fixed_delay_constructor_t<350, 450>::delay_t>1075 __device__ __forceinline__ void WaitForValid(1076 int tile_idx,1077 StatusWord &status,1078 KeyValuePairT &value,1079 DelayT delay_or_prevent_hoisting = {})1080 {1081// TxnWord alias = ThreadLoad<LOAD_CG>(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx);1082// TileDescriptor tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);1083//1084// while (tile_descriptor.status == SCAN_TILE_INVALID)1085// {1086// __threadfence_block(); // prevent hoisting loads from loop1087//1088// alias = ThreadLoad<LOAD_CG>(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx);1089// tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);1090// }1091//1092// status = tile_descriptor.status;1093// value.value = tile_descriptor.value;1094// value.key = tile_descriptor.key;1095 1096 TileDescriptor tile_descriptor;1097 1098 do1099 {1100 delay_or_prevent_hoisting();1101 TxnWord alias = detail::load_relaxed(d_tile_descriptors + TILE_STATUS_PADDING + tile_idx);1102 tile_descriptor = reinterpret_cast<TileDescriptor&>(alias);1103 1104 } while (WARP_ANY((tile_descriptor.status == SCAN_TILE_INVALID), 0xffffffff));1105 1106 status = tile_descriptor.status;1107 value.value = tile_descriptor.value;1108 value.key = tile_descriptor.key;1109 }1110 1111};1112 1113 1114/******************************************************************************1115 * Prefix call-back operator for coupling local block scan within a1116 * block-cooperative scan1117 ******************************************************************************/1118 1119/**1120 * Stateful block-scan prefix functor. Provides the the running prefix for1121 * the current tile by using the call-back warp to wait on on1122 * aggregates/prefixes from predecessor tiles to become available.1123 *1124 * @tparam DelayConstructorT1125 * Implementation detail, do not specify directly, requirements on the1126 * content of this type are subject to breaking change.1127 */1128template <1129 typename T,1130 typename ScanOpT,1131 typename ScanTileStateT,1132 int LEGACY_PTX_ARCH = 0,1133 typename DelayConstructorT = detail::default_delay_constructor_t<T>>1134struct TilePrefixCallbackOp1135{1136 // Parameterized warp reduce1137 typedef WarpReduce<T, CUB_PTX_WARP_THREADS> WarpReduceT;1138 1139 // Temporary storage type1140 struct _TempStorage1141 {1142 typename WarpReduceT::TempStorage warp_reduce;1143 T exclusive_prefix;1144 T inclusive_prefix;1145 T block_aggregate;1146 };1147 1148 // Alias wrapper allowing temporary storage to be unioned1149 struct TempStorage : Uninitialized<_TempStorage> {};1150 1151 // Type of status word1152 typedef typename ScanTileStateT::StatusWord StatusWord;1153 1154 // Fields1155 _TempStorage &temp_storage; ///< Reference to a warp-reduction instance1156 ScanTileStateT &tile_status; ///< Interface to tile status1157 ScanOpT scan_op; ///< Binary scan operator1158 int tile_idx; ///< The current tile index1159 T exclusive_prefix; ///< Exclusive prefix for the tile1160 T inclusive_prefix; ///< Inclusive prefix for the tile1161 1162 // Constructs prefix functor for a given tile index.1163 // Precondition: thread blocks processing all of the predecessor tiles were scheduled.1164 __device__ __forceinline__ TilePrefixCallbackOp(ScanTileStateT &tile_status,1165 TempStorage &temp_storage,1166 ScanOpT scan_op,1167 int tile_idx)1168 : temp_storage(temp_storage.Alias())1169 , tile_status(tile_status)1170 , scan_op(scan_op)1171 , tile_idx(tile_idx)1172 {}1173 1174 // Computes the tile index and constructs prefix functor with it.1175 // Precondition: thread block per tile assignment.1176 __device__ __forceinline__ TilePrefixCallbackOp(ScanTileStateT &tile_status,1177 TempStorage &temp_storage,1178 ScanOpT scan_op)1179 : TilePrefixCallbackOp(tile_status, temp_storage, scan_op, blockIdx.x)1180 {}1181 1182 /**1183 * @brief Block until all predecessors within the warp-wide window have non-invalid status1184 *1185 * @param predecessor_idx1186 * Preceding tile index to inspect1187 *1188 * @param[out] predecessor_status1189 * Preceding tile status1190 *1191 * @param[out] window_aggregate1192 * Relevant partial reduction from this window of preceding tiles1193 */1194 template <class DelayT = detail::default_delay_t<T>>1195 __device__ __forceinline__ void ProcessWindow(int predecessor_idx,1196 StatusWord &predecessor_status,1197 T &window_aggregate,1198 DelayT delay = {})1199 {1200 T value;