Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
single_pass_scan_operators.cuh1296 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 * 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;

Showing the first 1,200 of 1296 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai