Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
agent_batch_memcpy.cuh1217 linesDownload Raw Back to agent
1/******************************************************************************2 * Copyright (c) 2011-2022, NVIDIA CORPORATION. All rights reserved.3 *4 * Redistribution and use in source and binary forms, with or without5 * modification, are permitted provided that the following conditions are met:6 *     * Redistributions of source code must retain the above copyright7 *       notice, this list of conditions and the following disclaimer.8 *     * Redistributions in binary form must reproduce the above copyright9 *       notice, this list of conditions and the following disclaimer in the10 *       documentation and/or other materials provided with the distribution.11 *     * Neither the name of the NVIDIA CORPORATION nor the12 *       names of its contributors may be used to endorse or promote products13 *       derived from this software without specific prior written permission.14 *15 * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND16 * ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED17 * WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE18 * DISCLAIMED. IN NO EVENT SHALL NVIDIA CORPORATION BE LIABLE FOR ANY19 * DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES20 * (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;21 * LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND22 * ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT23 * (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS24 * SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.25 *26 ******************************************************************************/27 28/**29 * \file30 * cub::AgentBatchMemcpy implements device-wide copying of a batch of device-accessible31 * source-buffers to device-accessible destination-buffers.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_exchange.cuh>48#include <cub/block/block_load.cuh>49#include <cub/block/block_run_length_decode.cuh>50#include <cub/block/block_scan.cuh>51#include <cub/block/block_store.cuh>52#include <cub/util_ptx.cuh>53#include <cub/util_type.cuh>54 55#include <cuda/std/type_traits>56 57#include <cstdint>58 59CUB_NAMESPACE_BEGIN60 61namespace detail62{63template <bool PTR_IS_FOUR_BYTE_ALIGNED>64__forceinline__ __device__ void LoadVectorAndFunnelShiftR(uint32_t const *aligned_ptr,65                                                          uint32_t bit_shift,66                                                          uint4 &data_out)67{68  data_out = {aligned_ptr[0], aligned_ptr[1], aligned_ptr[2], aligned_ptr[3]};69 70  if (!PTR_IS_FOUR_BYTE_ALIGNED)71  {72    uint32_t tail = aligned_ptr[4];73    data_out.x    = __funnelshift_r(data_out.x, data_out.y, bit_shift);74    data_out.y    = __funnelshift_r(data_out.y, data_out.z, bit_shift);75    data_out.z    = __funnelshift_r(data_out.z, data_out.w, bit_shift);76    data_out.w    = __funnelshift_r(data_out.w, tail, bit_shift);77  }78}79 80template <bool PTR_IS_FOUR_BYTE_ALIGNED>81__forceinline__ __device__ void LoadVectorAndFunnelShiftR(uint32_t const *aligned_ptr,82                                                          uint32_t bit_shift,83                                                          uint2 &data_out)84{85  data_out = {aligned_ptr[0], aligned_ptr[1]};86 87  if (!PTR_IS_FOUR_BYTE_ALIGNED)88  {89    uint32_t tail = aligned_ptr[2];90    data_out.x    = __funnelshift_r(data_out.x, data_out.y, bit_shift);91    data_out.y    = __funnelshift_r(data_out.y, tail, bit_shift);92  }93}94 95template <bool PTR_IS_FOUR_BYTE_ALIGNED>96__forceinline__ __device__ void LoadVectorAndFunnelShiftR(uint32_t const *aligned_ptr,97                                                          uint32_t bit_shift,98                                                          uint32_t &data_out)99{100  data_out = aligned_ptr[0];101 102  if (!PTR_IS_FOUR_BYTE_ALIGNED)103  {104    uint32_t tail = aligned_ptr[1];105    data_out      = __funnelshift_r(data_out, tail, bit_shift);106  }107}108 109/**110 * @brief Loads data from \p ptr into \p data_out without requiring \p ptr to be aligned.111 * @note If \p ptr isn't aligned to four bytes, the bytes from the last four-byte aligned address up112 * to \p ptr are loaded too (but dropped) and, hence, need to be device-accessible. Similarly, if113 * \p ptr isn't aligned to four bytes, the bytes from `(ptr + sizeof(VectorT))` up to the following114 * four-byte aligned address are loaded too (but dropped), and, hence, need to be device-accessible.115 *116 * @tparam VectorT The vector type used for vectorized stores (i.e., one of uint4, uint2, uint32_t)117 * @param ptr The pointer from which the data is supposed to be loaded118 * @param data_out The vector type that stores the data loaded from \p ptr119 */120template <typename VectorT>121__forceinline__ __device__ void LoadVector(const char *ptr, VectorT &data_out)122{123  const uint32_t offset            = reinterpret_cast<std::uintptr_t>(ptr) % 4U;124  const uint32_t *aligned_ptr      = reinterpret_cast<uint32_t const *>(ptr - offset);125  constexpr uint32_t bits_per_byte = 8U;126  const uint32_t bit_shift         = offset * bits_per_byte;127 128  // If `ptr` is aligned to four bytes, we can perform a simple uint32_t-aliased load129  if (offset == 0)130  {131    LoadVectorAndFunnelShiftR<true>(aligned_ptr, bit_shift, data_out);132  }133  // Otherwise, we need to load extra bytes and perform funnel-shifting134  else135  {136    LoadVectorAndFunnelShiftR<false>(aligned_ptr, bit_shift, data_out);137  }138}139 140/**141 * @brief Helper data structure to hold information on the byte range for which we can safely142 * perform vectorized copies.143 *144 * @tparam VectorT The vector type used for vectorized stores (i.e., one of uint4, uint2, uint32_t)145 */146template <typename VectorT>147struct PointerRange148{149  VectorT *out_begin;150  VectorT *out_end;151  const char *in_begin;152  const char *in_end;153};154 155/**156 * @brief Both `out_start_aligned` and `out_end_aligned` are indices into `out_ptr`.157 * `out_start_aligned` is the first VectorT-aligned memory location after `out_ptr + 3`.158 * `out_end_aligned` is the last VectorT-aligned memory location before `out_end - 4`, where out_end159 * corresponds to one past the last byte to be copied. Bytes between `[out_start_aligned,160 * out_end_aligned)` will be copied using VectorT. `out_ptr + 3` and `out_end - 4` are used instead161 * of `out_ptr` and `out_end` to avoid `LoadVector` reading beyond data boundaries.162 *163 * @tparam VectorT The vector type used for vectorized stores (i.e., one of uint4, uint2, uint32_t)164 * @tparam ByteOffsetT Type used to index the bytes within the buffers165 * @param in_begin Pointer to the beginning of the byte range that shall be copied166 * @param out_begin Pointer to the beginning of the byte range that shall be copied167 * @param num_bytes Number of bytes that shall be copied168 * @return The byte range that can safely be copied using vectorized stores of type VectorT169 */170template <typename VectorT, typename ByteOffsetT>171__device__ __forceinline__ PointerRange<VectorT> GetAlignedPtrs(const void *in_begin,172                                                                void *out_begin,173                                                                ByteOffsetT num_bytes)174{175  // Data type size used for vectorized stores176  constexpr size_t out_datatype_size = sizeof(VectorT);177  // Data type size used for type-aliased loads178  constexpr size_t in_datatype_size = sizeof(uint32_t);179 180  // char-aliased ptrs to simplify pointer arithmetic181  char *out_ptr      = reinterpret_cast<char *>(out_begin);182  const char *in_ptr = reinterpret_cast<const char *>(in_begin);183 184  // Number of bytes between the first VectorT-aligned address at or before out_begin and out_begin185  const uint32_t alignment_offset = reinterpret_cast<std::uintptr_t>(out_ptr) % out_datatype_size;186 187  // The first VectorT-aligned address before (or at) out_begin188  char *out_chars_aligned = reinterpret_cast<char *>(out_ptr - alignment_offset);189 190  // The number of extra bytes preceding `in_ptr` that are loaded but dropped191  uint32_t in_extra_bytes = reinterpret_cast<std::uintptr_t>(in_ptr) % in_datatype_size;192 193  // The offset required by `LoadVector`:194  // If the input pointer is not aligned, we load data from the last aligned address preceding the195  // pointer. That is, loading up to (in_datatype_size-1) bytes before `in_ptr`196  uint32_t in_offset_req = in_extra_bytes;197 198  // Bytes after `out_chars_aligned` to the first VectorT-aligned address at or after `out_begin`199  uint32_t out_start_aligned =200    CUB_QUOTIENT_CEILING(in_offset_req + alignment_offset, out_datatype_size) * out_datatype_size;201 202  // Compute the beginning of the aligned ranges (output and input pointers)203  VectorT *out_aligned_begin   = reinterpret_cast<VectorT *>(out_chars_aligned + out_start_aligned);204  const char *in_aligned_begin = in_ptr + (reinterpret_cast<char *>(out_aligned_begin) - out_ptr);205 206  // If the aligned range is not aligned for the input pointer, we load up to (in_datatype_size-1)207  // bytes after the last byte that is copied. That is, we always load four bytes up to the next208  // aligned input address at a time. E.g., if the last byte loaded is one byte past the last209  // aligned address we'll also load the three bytes after that byte.210  uint32_t in_extra_bytes_from_aligned =211    (reinterpret_cast<std::uintptr_t>(in_aligned_begin) % in_datatype_size);212  uint32_t in_end_padding_req = (in_datatype_size - in_extra_bytes_from_aligned) % in_datatype_size;213 214  // Bytes after `out_chars_aligned` to the last VectorT-aligned215  // address at (or before) `out_begin` + `num_bytes`216  uint32_t out_end_aligned{};217  if (in_end_padding_req + alignment_offset > num_bytes)218  {219    out_end_aligned = out_start_aligned;220  }221  else222  {223    out_end_aligned = (num_bytes - in_end_padding_req + alignment_offset) / out_datatype_size *224                      out_datatype_size;225  }226 227  VectorT *out_aligned_end   = reinterpret_cast<VectorT *>(out_chars_aligned + out_end_aligned);228  const char *in_aligned_end = in_ptr + (reinterpret_cast<char *>(out_aligned_end) - out_ptr);229 230  return {out_aligned_begin, out_aligned_end, in_aligned_begin, in_aligned_end};231}232 233/**234 * @brief Cooperatively copies \p num_bytes from \p src to \p dest using vectorized stores of type235 * \p VectorT for addresses within [dest, dest + num_bytes) that are aligned to \p VectorT. A236 * byte-wise copy is used for byte-ranges that are not aligned to \p VectorT.237 *238 * @tparam LOGICAL_WARP_SIZE The number of threads cooperaing to copy the data; all threads within239 * [0,  `LOGICAL_WARP_SIZE`) must invoke this method with the same arguments240 * @tparam VectorT The vector type used for vectorized stores (i.e., one of uint4, uint2, uint32_t)241 * @tparam ByteOffsetT Type used to index the bytes within the buffers242 * @param thread_rank The thread rank within the group that cooperates to copy the data must be243 * within [0, `LOGICAL_WARP_SIZE`)244 * @param dest Pointer to the memory location to copy to245 * @param num_bytes Number of bytes to copy246 * @param src Pointer to the memory location to copy from247 */248template <int LOGICAL_WARP_SIZE, typename VectorT, typename ByteOffsetT>249__device__ __forceinline__ void250VectorizedCopy(int32_t thread_rank, void *dest, ByteOffsetT num_bytes, const void *src)251{252  char *out_ptr      = reinterpret_cast<char *>(dest);253  const char *in_ptr = reinterpret_cast<const char *>(src);254 255  // Gets the byte range that can safely be copied using vectorized stores of type VectorT256  auto aligned_range = GetAlignedPtrs<VectorT>(src, dest, num_bytes);257 258  // If byte range for which we can use vectorized copies is empty -> use byte-wise copies259  if (aligned_range.out_end <= aligned_range.out_begin)260  {261    for (ByteOffsetT ichar = thread_rank; ichar < num_bytes; ichar += LOGICAL_WARP_SIZE)262    {263      out_ptr[ichar] = in_ptr[ichar];264    }265  }266  else267  {268    // Copy bytes in range `[dest, aligned_range.out_begin)`269    out_ptr += thread_rank;270    in_ptr += thread_rank;271    while (out_ptr < reinterpret_cast<char *>(aligned_range.out_begin))272    {273      *out_ptr = *in_ptr;274      out_ptr += LOGICAL_WARP_SIZE;275      in_ptr += LOGICAL_WARP_SIZE;276    }277 278    // Copy bytes in range `[aligned_range.out_begin, aligned_range.out_end)`279    VectorT *aligned_range_begin = aligned_range.out_begin + thread_rank;280    const char *in_aligned_begin = aligned_range.in_begin + thread_rank * sizeof(VectorT);281    while (aligned_range_begin < aligned_range.out_end)282    {283      VectorT data_in;284      LoadVector(in_aligned_begin, data_in);285      *aligned_range_begin = data_in;286      in_aligned_begin += sizeof(VectorT) * LOGICAL_WARP_SIZE;287      aligned_range_begin += LOGICAL_WARP_SIZE;288    }289 290    // Copy bytes in range `[aligned_range.out_end, dest + num_bytes)`.291    out_ptr = reinterpret_cast<char *>(aligned_range.out_end) + thread_rank;292    in_ptr  = aligned_range.in_end + thread_rank;293    while (out_ptr < reinterpret_cast<char *>(dest) + num_bytes)294    {295      *out_ptr = *in_ptr;296      out_ptr += LOGICAL_WARP_SIZE;297      in_ptr += LOGICAL_WARP_SIZE;298    }299  }300}301 302template <bool IsMemcpy,303          uint32_t LOGICAL_WARP_SIZE,304          typename InputBufferT,305          typename OutputBufferT,306          typename OffsetT,307          typename ::cuda::std::enable_if<IsMemcpy, int>::type = 0>308__device__ __forceinline__ void copy_items(InputBufferT input_buffer,309                                           OutputBufferT output_buffer,310                                           OffsetT num_bytes,311                                           OffsetT offset = 0)312{313  VectorizedCopy<LOGICAL_WARP_SIZE, uint4>(threadIdx.x % LOGICAL_WARP_SIZE,314                                           &reinterpret_cast<char *>(output_buffer)[offset],315                                           num_bytes,316                                           &reinterpret_cast<const char *>(input_buffer)[offset]);317}318 319template <bool IsMemcpy,320          uint32_t LOGICAL_WARP_SIZE,321          typename InputBufferT,322          typename OutputBufferT,323          typename OffsetT,324          typename ::cuda::std::enable_if<!IsMemcpy, int>::type = 0>325__device__ __forceinline__ void copy_items(InputBufferT input_buffer,326                                           OutputBufferT output_buffer,327                                           OffsetT num_items,328                                           OffsetT offset = 0)329{330  output_buffer += offset;331  input_buffer += offset;332  for (OffsetT i = threadIdx.x % LOGICAL_WARP_SIZE; i < num_items; i += LOGICAL_WARP_SIZE)333  {334    *(output_buffer + i) = *(input_buffer + i);335  }336}337 338template <bool IsMemcpy,339          typename AliasT,340          typename InputIt,341          typename OffsetT,342          typename ::cuda::std::enable_if<IsMemcpy, int>::type = 0>343__device__ __forceinline__ AliasT read_item(InputIt buffer_src, OffsetT offset)344{345  return *(reinterpret_cast<const AliasT *>(buffer_src) + offset);346}347 348template <bool IsMemcpy,349          typename AliasT,350          typename InputIt,351          typename OffsetT,352          typename ::cuda::std::enable_if<!IsMemcpy, int>::type = 0>353__device__ __forceinline__ AliasT read_item(InputIt buffer_src, OffsetT offset)354{355  return *(buffer_src + offset);356}357 358template <bool IsMemcpy,359          typename AliasT,360          typename OutputIt,361          typename OffsetT,362          typename ::cuda::std::enable_if<IsMemcpy, int>::type = 0>363__device__ __forceinline__ void write_item(OutputIt buffer_dst, OffsetT offset, AliasT value)364{365  *(reinterpret_cast<AliasT *>(buffer_dst) + offset) = value;366}367 368template <bool IsMemcpy,369          typename AliasT,370          typename OutputIt,371          typename OffsetT,372          typename ::cuda::std::enable_if<!IsMemcpy, int>::type = 0>373__device__ __forceinline__ void write_item(OutputIt buffer_dst, OffsetT offset, AliasT value)374{375  *(buffer_dst + offset) = value;376}377 378/**379 * @brief A helper class that allows threads to maintain multiple counters, where the counter that380 * shall be incremented can be addressed dynamically without incurring register spillage.381 *382 * @tparam NUM_ITEMS The number of counters to allocate383 * @tparam MAX_ITEM_VALUE The maximum count that must be supported.384 * @tparam PREFER_POW2_BITS Whether the number of bits to dedicate to each counter should be a385 * power-of-two. If enabled, this allows replacing integer multiplication with a bit-shift in386 * exchange for higher register pressure.387 * @tparam BackingUnitT The data type that is used to provide the bits of all the counters that388 * shall be allocated.389 */390template <uint32_t NUM_ITEMS,391          uint32_t MAX_ITEM_VALUE,392          bool PREFER_POW2_BITS,393          typename BackingUnitT = uint32_t>394class BitPackedCounter395{396private:397  /// The minimum number of bits required to represent all values from [0, MAX_ITEM_VALUE]398  static constexpr uint32_t MIN_BITS_PER_ITEM =399    (MAX_ITEM_VALUE == 0U) ? 1U : cub::Log2<static_cast<int32_t>(MAX_ITEM_VALUE + 1U)>::VALUE;400 401  /// The number of bits allocated for each item. For pre-Volta, we prefer a power-of-2 here to402  /// have the compiler replace costly integer multiplication with bit-shifting.403  static constexpr uint32_t BITS_PER_ITEM =404    PREFER_POW2_BITS ? (0x01ULL << (cub::Log2<static_cast<int32_t>(MIN_BITS_PER_ITEM)>::VALUE))405                     : MIN_BITS_PER_ITEM;406 407  /// The number of bits that each backing data type can store408  static constexpr uint32_t NUM_BITS_PER_UNIT = sizeof(BackingUnitT) * 8;409 410  /// The number of items that each backing data type can store411  static constexpr uint32_t ITEMS_PER_UNIT = NUM_BITS_PER_UNIT / BITS_PER_ITEM;412 413  /// The number of bits the backing data type is actually making use of414  static constexpr uint32_t USED_BITS_PER_UNIT = ITEMS_PER_UNIT * BITS_PER_ITEM;415 416  /// The number of backing data types required to store the given number of items417  static constexpr uint32_t NUM_TOTAL_UNITS = CUB_QUOTIENT_CEILING(NUM_ITEMS, ITEMS_PER_UNIT);418 419  /// This is the net number of bit-storage provided by each unit (remainder bits are unused)420  static constexpr uint32_t UNIT_MASK = (USED_BITS_PER_UNIT >= (8U * sizeof(uint32_t)))421                                          ? 0xFFFFFFFF422                                          : (0x01U << USED_BITS_PER_UNIT) - 1;423  /// This is the bit-mask for each item424  static constexpr uint32_t ITEM_MASK = (BITS_PER_ITEM >= (8U * sizeof(uint32_t)))425                                          ? 0xFFFFFFFF426                                          : (0x01U << BITS_PER_ITEM) - 1;427 428  //------------------------------------------------------------------------------429  // ACCESSORS430  //------------------------------------------------------------------------------431public:432  __device__ __forceinline__ uint32_t Get(uint32_t index) const433  {434    const uint32_t target_offset = index * BITS_PER_ITEM;435    uint32_t val                 = 0;436 437#pragma unroll438    for (uint32_t i = 0; i < NUM_TOTAL_UNITS; ++i)439    {440      // In case the bit-offset of the counter at <index> is larger than the bit range of the441      // current unit, the bit_shift amount will be larger than the bits provided by this unit. As442      // C++'s bit-shift has undefined behaviour if the bits being shifted exceed the operand width,443      // we use the PTX instruction `shr` to make sure behaviour is well-defined.444      // Negative bit-shift amounts wrap around in unsigned integer math and are ultimately clamped.445      const uint32_t bit_shift = target_offset - i * USED_BITS_PER_UNIT;446      val |= detail::LogicShiftRight(data[i], bit_shift) & ITEM_MASK;447    }448    return val;449  }450 451  __device__ __forceinline__ void Add(uint32_t index, uint32_t value)452  {453    const uint32_t target_offset = index * BITS_PER_ITEM;454 455#pragma unroll456    for (uint32_t i = 0; i < NUM_TOTAL_UNITS; ++i)457    {458      // In case the bit-offset of the counter at <index> is larger than the bit range of the459      // current unit, the bit_shift amount will be larger than the bits provided by this unit. As460      // C++'s bit-shift has undefined behaviour if the bits being shifted exceed the operand width,461      // we use the PTX instruction `shl` to make sure behaviour is well-defined.462      // Negative bit-shift amounts wrap around in unsigned integer math and are ultimately clamped.463      const uint32_t bit_shift = target_offset - i * USED_BITS_PER_UNIT;464      data[i] += detail::LogicShiftLeft(value, bit_shift) & UNIT_MASK;465    }466  }467 468  __device__ BitPackedCounter operator+(const BitPackedCounter &rhs) const469  {470    BitPackedCounter result;471#pragma unroll472    for (uint32_t i = 0; i < NUM_TOTAL_UNITS; ++i)473    {474      result.data[i] = data[i] + rhs.data[i];475    }476    return result;477  }478 479  //------------------------------------------------------------------------------480  // MEMBER VARIABLES481  //------------------------------------------------------------------------------482private:483  BackingUnitT data[NUM_TOTAL_UNITS] = {};484};485 486/**487 * Parameterizable tuning policy type for AgentBatchMemcpy488 */489template <uint32_t _BLOCK_THREADS,490          uint32_t _BUFFERS_PER_THREAD,491          uint32_t _TLEV_BYTES_PER_THREAD,492          bool _PREFER_POW2_BITS,493          uint32_t _BLOCK_LEVEL_TILE_SIZE,494          uint32_t _WARP_LEVEL_THRESHOLD,495          uint32_t _BLOCK_LEVEL_THRESHOLD,496          class BuffDelayConstructor,497          class BlockDelayConstructor>498struct AgentBatchMemcpyPolicy499{500  /// Threads per thread block501  static constexpr uint32_t BLOCK_THREADS = _BLOCK_THREADS;502  /// Items per thread (per tile of input)503  static constexpr uint32_t BUFFERS_PER_THREAD = _BUFFERS_PER_THREAD;504  /// The number of bytes that each thread will work on with each iteration of reading in bytes505  /// from one or more506  // source-buffers and writing them out to the respective destination-buffers.507  static constexpr uint32_t TLEV_BYTES_PER_THREAD = _TLEV_BYTES_PER_THREAD;508  /// Whether the BitPackedCounter should prefer allocating a power-of-2 number of bits per509  /// counter510  static constexpr uint32_t PREFER_POW2_BITS = _PREFER_POW2_BITS;511  /// BLEV tile size granularity512  static constexpr uint32_t BLOCK_LEVEL_TILE_SIZE = _BLOCK_LEVEL_TILE_SIZE;513 514  static constexpr uint32_t WARP_LEVEL_THRESHOLD  = _WARP_LEVEL_THRESHOLD;515  static constexpr uint32_t BLOCK_LEVEL_THRESHOLD = _BLOCK_LEVEL_THRESHOLD;516 517  using buff_delay_constructor = BuffDelayConstructor;518  using block_delay_constructor = BlockDelayConstructor;519};520 521template <typename AgentMemcpySmallBuffersPolicyT,522          typename InputBufferIt,523          typename OutputBufferIt,524          typename BufferSizeIteratorT,525          typename BufferOffsetT,526          typename BlevBufferSrcsOutItT,527          typename BlevBufferDstsOutItT,528          typename BlevBufferSizesOutItT,529          typename BlevBufferTileOffsetsOutItT,530          typename BlockOffsetT,531          typename BLevBufferOffsetTileState,532          typename BLevBlockOffsetTileState,533          bool IsMemcpy>534class AgentBatchMemcpy535{536private:537  //---------------------------------------------------------------------538  // CONFIGS / CONSTANTS539  //---------------------------------------------------------------------540  // Tuning policy-based configurations541  static constexpr uint32_t BLOCK_THREADS      = AgentMemcpySmallBuffersPolicyT::BLOCK_THREADS;542  static constexpr uint32_t BUFFERS_PER_THREAD = AgentMemcpySmallBuffersPolicyT::BUFFERS_PER_THREAD;543  static constexpr uint32_t TLEV_BYTES_PER_THREAD =544    AgentMemcpySmallBuffersPolicyT::TLEV_BYTES_PER_THREAD;545  static constexpr bool PREFER_POW2_BITS = AgentMemcpySmallBuffersPolicyT::PREFER_POW2_BITS;546  static constexpr uint32_t BLOCK_LEVEL_TILE_SIZE =547    AgentMemcpySmallBuffersPolicyT::BLOCK_LEVEL_TILE_SIZE;548 549  // Derived configs550  static constexpr uint32_t BUFFERS_PER_BLOCK       = BUFFERS_PER_THREAD * BLOCK_THREADS;551  static constexpr uint32_t TLEV_BUFFERS_PER_THREAD = BUFFERS_PER_THREAD;552  static constexpr uint32_t BLEV_BUFFERS_PER_THREAD = BUFFERS_PER_THREAD;553 554  static constexpr uint32_t WARP_LEVEL_THRESHOLD =555    AgentMemcpySmallBuffersPolicyT::WARP_LEVEL_THRESHOLD;556 557  static constexpr uint32_t BLOCK_LEVEL_THRESHOLD =558    AgentMemcpySmallBuffersPolicyT::BLOCK_LEVEL_THRESHOLD;559 560  static constexpr uint32_t BUFFER_STABLE_PARTITION = false;561 562  // Constants563  enum : uint32_t564  {565    TLEV_SIZE_CLASS = 0,566    WLEV_SIZE_CLASS,567    BLEV_SIZE_CLASS,568    NUM_SIZE_CLASSES,569  };570 571  //---------------------------------------------------------------------572  // TYPE DECLARATIONS573  //---------------------------------------------------------------------574  /// Internal load/store type. For byte-wise memcpy, a single-byte type575  using AliasT = typename ::cuda::std::conditional<576    IsMemcpy,577    std::iterator_traits<char *>,578    std::iterator_traits<cub::detail::value_t<InputBufferIt>>>::type::value_type;579 580  /// Types of the input and output buffers581  using InputBufferT  = cub::detail::value_t<InputBufferIt>;582  using OutputBufferT = cub::detail::value_t<OutputBufferIt>;583 584  /// Type that has to be sufficiently large to hold any of the buffers' sizes.585  /// The BufferSizeIteratorT's value type must be convertible to this type.586  using BufferSizeT = cub::detail::value_t<BufferSizeIteratorT>;587 588  /// Type used to index into the tile of buffers that this thread block is assigned to.589  using BlockBufferOffsetT = uint16_t;590 591  /// Internal type used to index into the bytes of and represent size of a TLEV buffer592  using TLevBufferSizeT = uint16_t;593 594  /**595   * @brief Helper struct to simplify BlockExchange within a single four-byte word596   */597  struct ZippedTLevByteAssignment598  {599    // The buffer id within this tile600    BlockBufferOffsetT tile_buffer_id;601 602    // Byte-offset within that buffer603    TLevBufferSizeT buffer_byte_offset;604  };605 606  /**607   * POD to keep track of <buffer_id, buffer_size> pairs after having partitioned this tile's608   * buffers by their size.609   */610  struct BufferTuple611  {612    // Size is only valid (and relevant) for buffers that are use thread-level collaboration613    TLevBufferSizeT size;614 615    // The buffer id relativ to this tile (i.e., the buffer id within this tile)616    BlockBufferOffsetT buffer_id;617  };618 619  // Load buffers in a striped arrangement if we do not want to performa a stable partitioning into620  // small, medium, and large buffers, otherwise load them in a blocked arrangement621  using BufferLoadT =622    BlockLoad<BufferSizeT,623              static_cast<int32_t>(BLOCK_THREADS),624              static_cast<int32_t>(BUFFERS_PER_THREAD),625              BUFFER_STABLE_PARTITION ? BLOCK_LOAD_WARP_TRANSPOSE : BLOCK_LOAD_STRIPED>;626 627  // A vectorized counter that will count the number of buffers that fall into each of the628  // size-classes. Where the size class representes the collaboration level that is required to629  // process a buffer. The collaboration level being either:630  //-> (1) TLEV (thread-level collaboration), requiring one or multiple threads but not a FULL warp631  // to collaborate632  //-> (2) WLEV (warp-level collaboration), requiring a full warp to collaborate on a buffer633  //-> (3) BLEV (block-level collaboration), requiring one or multiple thread blocks to collaborate634  // on a buffer */635  using VectorizedSizeClassCounterT =636    BitPackedCounter<NUM_SIZE_CLASSES, BUFFERS_PER_BLOCK, PREFER_POW2_BITS>;637 638  // Block-level scan used to compute the write offsets639  using BlockSizeClassScanT =640    cub::BlockScan<VectorizedSizeClassCounterT, static_cast<int32_t>(BLOCK_THREADS)>;641 642  //643  using BlockBLevTileCountScanT = cub::BlockScan<BlockOffsetT, static_cast<int32_t>(BLOCK_THREADS)>;644 645  // Block-level run-length decode algorithm to evenly distribute work of all buffers requiring646  // thread-level collaboration647  using BlockRunLengthDecodeT =648    cub::BlockRunLengthDecode<BlockBufferOffsetT,649                              static_cast<int32_t>(BLOCK_THREADS),650                              static_cast<int32_t>(TLEV_BUFFERS_PER_THREAD),651                              static_cast<int32_t>(TLEV_BYTES_PER_THREAD)>;652 653  using BlockExchangeTLevT = cub::BlockExchange<ZippedTLevByteAssignment,654                                                static_cast<int32_t>(BLOCK_THREADS),655                                                static_cast<int32_t>(TLEV_BYTES_PER_THREAD)>;656 657  using BLevBuffScanPrefixCallbackOpT =658    TilePrefixCallbackOp<BufferOffsetT,659                         Sum,660                         BLevBufferOffsetTileState,661                         0,662                         typename AgentMemcpySmallBuffersPolicyT::buff_delay_constructor>;663 664  using BLevBlockScanPrefixCallbackOpT =665    TilePrefixCallbackOp<BlockOffsetT,666                         Sum,667                         BLevBlockOffsetTileState,668                         0,669                         typename AgentMemcpySmallBuffersPolicyT::block_delay_constructor>;670 671  //-----------------------------------------------------------------------------672  // SHARED MEMORY DECLARATIONS673  //-----------------------------------------------------------------------------674  struct _TempStorage675  {676    union677    {678      typename BufferLoadT::TempStorage load_storage;679 680      // Stage 1: histogram over the size classes in preparation for partitioning buffers by size681      typename BlockSizeClassScanT::TempStorage size_scan_storage;682 683      // Stage 2: Communicate the number ofer buffers requiring block-level collaboration684      typename BLevBuffScanPrefixCallbackOpT::TempStorage buffer_scan_callback;685 686      // Stage 3; batch memcpy buffers that require only thread-level collaboration687      struct688      {689        BufferTuple buffers_by_size_class[BUFFERS_PER_BLOCK];690 691        // Stage 3.1: Write buffers requiring block-level collaboration to queue692        union693        {694          struct695          {696            typename BLevBlockScanPrefixCallbackOpT::TempStorage block_scan_callback;697            typename BlockBLevTileCountScanT::TempStorage block_scan_storage;698          } blev;699 700          // Stage 3.3: run-length decode & block exchange for tlev701          // rld_state needs to be persistent across loop iterations (RunLengthDecode calls) and,702          // hence, cannot alias block_exchange_storage703          struct704          {705            typename BlockRunLengthDecodeT::TempStorage rld_state;706            typename BlockExchangeTLevT::TempStorage block_exchange_storage;707          } tlev;708        };709      } staged;710    };711    BufferOffsetT blev_buffer_offset;712  };713 714  //-----------------------------------------------------------------------------715  // PUBLIC TYPE MEMBERS716  //-----------------------------------------------------------------------------717public:718  struct TempStorage : Uninitialized<_TempStorage>719  {};720 721  //-----------------------------------------------------------------------------722  // PRIVATE MEMBER FUNCTIONS723  //-----------------------------------------------------------------------------724private:725  /// Shared storage reference726  _TempStorage &temp_storage;727 728  /**729   * @brief Loads this tile's buffers' sizes, without any guards (i.e., out-of-bounds checks)730   */731  __device__ __forceinline__ void732  LoadBufferSizesFullTile(BufferSizeIteratorT tile_buffer_sizes_it,733                          BufferSizeT (&buffer_sizes)[BUFFERS_PER_THREAD])734  {735    BufferLoadT(temp_storage.load_storage).Load(tile_buffer_sizes_it, buffer_sizes);736  }737 738  /**739   * @brief Loads this tile's buffers' sizes, making sure to read at most \p num_valid items.740   */741  __device__ __forceinline__ void742  LoadBufferSizesPartialTile(BufferSizeIteratorT tile_buffer_sizes_it,743                             BufferSizeT (&buffer_sizes)[BUFFERS_PER_THREAD],744                             BufferOffsetT num_valid)745  {746    // Out-of-bounds buffer items are initialized to '0', so those buffers will simply be ignored747    // later on748    constexpr BufferSizeT OOB_DEFAULT_BUFFER_SIZE = 0U;749 750    BufferLoadT(temp_storage.load_storage)751      .Load(tile_buffer_sizes_it, buffer_sizes, num_valid, OOB_DEFAULT_BUFFER_SIZE);752  }753 754  /**755   * @brief Computes the histogram over the number of buffers belonging to each of the three756   * size-classes (TLEV, WLEV, BLEV).757   */758  __device__ __forceinline__ VectorizedSizeClassCounterT759  GetBufferSizeClassHistogram(const BufferSizeT (&buffer_sizes)[BUFFERS_PER_THREAD])760  {761    VectorizedSizeClassCounterT vectorized_counters{};762#pragma unroll763    for (uint32_t i = 0; i < BUFFERS_PER_THREAD; i++)764    {765      // Whether to increment ANY of the buffer size classes at all766      const uint32_t increment = buffer_sizes[i] > 0 ? 1U : 0U;767      // Identify the buffer's size class768      uint32_t buffer_size_class = 0;769      buffer_size_class += buffer_sizes[i] > WARP_LEVEL_THRESHOLD ? 1U : 0U;770      buffer_size_class += buffer_sizes[i] > BLOCK_LEVEL_THRESHOLD ? 1U : 0U;771 772      // Increment the count of the respective size class773      vectorized_counters.Add(buffer_size_class, increment);774    }775    return vectorized_counters;776  }777 778  /**779   * @brief Scatters the buffers into the respective buffer's size-class partition.780   */781  __device__ __forceinline__ void782  PartitionBuffersBySize(const BufferSizeT (&buffer_sizes)[BUFFERS_PER_THREAD],783                         VectorizedSizeClassCounterT &vectorized_offsets,784                         BufferTuple (&buffers_by_size_class)[BUFFERS_PER_BLOCK])785  {786    // If we intend to perform a stable partitioning, the thread's buffer are in a blocked787    // arrangement, otherwise they are in a striped arrangement788    BlockBufferOffsetT buffer_id = BUFFER_STABLE_PARTITION ? (BUFFERS_PER_THREAD * threadIdx.x)789                                                           : (threadIdx.x);790    constexpr BlockBufferOffsetT BUFFER_STRIDE = BUFFER_STABLE_PARTITION791                                                   ? static_cast<BlockBufferOffsetT>(1)792                                                   : static_cast<BlockBufferOffsetT>(BLOCK_THREADS);793 794#pragma unroll795    for (uint32_t i = 0; i < BUFFERS_PER_THREAD; i++)796    {797      if (buffer_sizes[i] > 0)798      {799        uint32_t buffer_size_class = 0;800        buffer_size_class += buffer_sizes[i] > WARP_LEVEL_THRESHOLD ? 1U : 0U;801        buffer_size_class += buffer_sizes[i] > BLOCK_LEVEL_THRESHOLD ? 1U : 0U;802        const uint32_t write_offset         = vectorized_offsets.Get(buffer_size_class);803        buffers_by_size_class[write_offset] = {static_cast<TLevBufferSizeT>(buffer_sizes[i]),804                                               buffer_id};805        vectorized_offsets.Add(buffer_size_class, 1U);806      }807      buffer_id += BUFFER_STRIDE;808    }809  }810 811  /**812   * @brief Read in all the buffers that require block-level collaboration and put them to a queue813   * that will get picked up in a separate, subsequent kernel.814   */815  __device__ __forceinline__ void EnqueueBLEVBuffers(BufferTuple *buffers_by_size_class,816                                                     InputBufferIt tile_buffer_srcs,817                                                     OutputBufferIt tile_buffer_dsts,818                                                     BufferSizeIteratorT tile_buffer_sizes,819                                                     BlockBufferOffsetT num_blev_buffers,820                                                     BufferOffsetT tile_buffer_offset,821                                                     BufferOffsetT tile_id)822  {823    BlockOffsetT block_offset[BLEV_BUFFERS_PER_THREAD];824    // Read in the BLEV buffer partition (i.e., the buffers that require block-level collaboration)825    uint32_t blev_buffer_offset = threadIdx.x * BLEV_BUFFERS_PER_THREAD;826#pragma unroll827    for (uint32_t i = 0; i < BLEV_BUFFERS_PER_THREAD; i++)828    {829      if (blev_buffer_offset < num_blev_buffers)830      {831        BlockBufferOffsetT tile_buffer_id = buffers_by_size_class[blev_buffer_offset].buffer_id;832        block_offset[i]                   = CUB_QUOTIENT_CEILING(tile_buffer_sizes[tile_buffer_id],833                                               BLOCK_LEVEL_TILE_SIZE);834      }835      else836      {837        // Out-of-bounds buffers are assigned a tile count of '0'838        block_offset[i] = 0U;839      }840      blev_buffer_offset++;841    }842 843    if (tile_id == 0)844    {845      BlockOffsetT block_aggregate;846      BlockBLevTileCountScanT(temp_storage.staged.blev.block_scan_storage)847        .ExclusiveSum(block_offset, block_offset, block_aggregate);848      if (threadIdx.x == 0)849      {850        blev_block_scan_state.SetInclusive(0, block_aggregate);851      }852    }853    else854    {855      BLevBlockScanPrefixCallbackOpT blev_tile_prefix_op(856        blev_block_scan_state,857        temp_storage.staged.blev.block_scan_callback,858        Sum(),859        tile_id);860      BlockBLevTileCountScanT(temp_storage.staged.blev.block_scan_storage)861        .ExclusiveSum(block_offset, block_offset, blev_tile_prefix_op);862    }863    CTA_SYNC();864 865    // Read in the BLEV buffer partition (i.e., the buffers that require block-level collaboration)866    blev_buffer_offset = threadIdx.x * BLEV_BUFFERS_PER_THREAD;867#pragma unroll868    for (uint32_t i = 0; i < BLEV_BUFFERS_PER_THREAD; i++)869    {870      if (blev_buffer_offset < num_blev_buffers)871      {872        BlockBufferOffsetT tile_buffer_id = buffers_by_size_class[blev_buffer_offset].buffer_id;873        blev_buffer_srcs[tile_buffer_offset + blev_buffer_offset] =874          tile_buffer_srcs[tile_buffer_id];875        blev_buffer_dsts[tile_buffer_offset + blev_buffer_offset] =876          tile_buffer_dsts[tile_buffer_id];877        blev_buffer_sizes[tile_buffer_offset + blev_buffer_offset] =878          tile_buffer_sizes[tile_buffer_id];879        blev_buffer_tile_offsets[tile_buffer_offset + blev_buffer_offset] = block_offset[i];880        blev_buffer_offset++;881      }882    }883  }884 885  /**886   * @brief Read in all the buffers of this tile that require warp-level collaboration and copy887   * their bytes to the corresponding destination buffer888   */889  __device__ __forceinline__ void BatchMemcpyWLEVBuffers(BufferTuple *buffers_by_size_class,890                                                         InputBufferIt tile_buffer_srcs,891                                                         OutputBufferIt tile_buffer_dsts,892                                                         BufferSizeIteratorT tile_buffer_sizes,893                                                         BlockBufferOffsetT num_wlev_buffers)894  {895    const int32_t warp_id              = threadIdx.x / CUB_PTX_WARP_THREADS;896    constexpr uint32_t WARPS_PER_BLOCK = BLOCK_THREADS / CUB_PTX_WARP_THREADS;897 898    for (BlockBufferOffsetT buffer_offset = warp_id; buffer_offset < num_wlev_buffers;899         buffer_offset += WARPS_PER_BLOCK)900    {901      const auto buffer_id = buffers_by_size_class[buffer_offset].buffer_id;902      copy_items<IsMemcpy, CUB_PTX_WARP_THREADS, InputBufferT, OutputBufferT, BufferSizeT>(903        tile_buffer_srcs[buffer_id],904        tile_buffer_dsts[buffer_id],905        tile_buffer_sizes[buffer_id]);906    }907  }908 909  /**910   * @brief Read in all the buffers of this tile that require thread-level collaboration and copy911   * their bytes to the corresponding destination buffer912   */913  __device__ __forceinline__ void BatchMemcpyTLEVBuffers(BufferTuple *buffers_by_size_class,914                                                         InputBufferIt tile_buffer_srcs,915                                                         OutputBufferIt tile_buffer_dsts,916                                                         BlockBufferOffsetT num_tlev_buffers)917  {918    // Read in the buffers' ids that require thread-level collaboration (where buffer id is the919    // buffer within this tile)920    BlockBufferOffsetT tlev_buffer_ids[TLEV_BUFFERS_PER_THREAD];921    TLevBufferSizeT tlev_buffer_sizes[TLEV_BUFFERS_PER_THREAD];922    // Currently we do not go over the TLEV buffers in multiple iterations, so we need to make sure923    // we are able to be covered for the case that all our buffers are TLEV buffers924    static_assert(TLEV_BUFFERS_PER_THREAD >= BUFFERS_PER_THREAD,925                  "Unsupported confiugraiton: The number of 'thread-level buffers' must be at "926                  "least as large as the number of overall buffers being processed by each "927                  "thread.");928 929    // Read in the TLEV buffer partition (i.e., the buffers that require thread-level collaboration)930    uint32_t tlev_buffer_offset = threadIdx.x * TLEV_BUFFERS_PER_THREAD;931 932    // Pre-populate the buffer sizes to 0 (i.e. zero-padding towards the end) to ensure933    // out-of-bounds TLEV buffers will not be considered934#pragma unroll935    for (uint32_t i = 0; i < TLEV_BUFFERS_PER_THREAD; i++)936    {937      tlev_buffer_sizes[i] = 0;938    }939 940    // Assign TLEV buffers in a blocked arrangement (each thread is assigned consecutive TLEV941    // buffers)942#pragma unroll943    for (uint32_t i = 0; i < TLEV_BUFFERS_PER_THREAD; i++)944    {945      if (tlev_buffer_offset < num_tlev_buffers)946      {947        tlev_buffer_ids[i]   = buffers_by_size_class[tlev_buffer_offset].buffer_id;948        tlev_buffer_sizes[i] = buffers_by_size_class[tlev_buffer_offset].size;949      }950      tlev_buffer_offset++;951    }952 953    // Evenly distribute all the bytes that have to be copied from all the buffers that require954    // thread-level collaboration using BlockRunLengthDecode955    uint32_t num_total_tlev_bytes = 0U;956    BlockRunLengthDecodeT block_run_length_decode(temp_storage.staged.tlev.rld_state,957                                                  tlev_buffer_ids,958                                                  tlev_buffer_sizes,959                                                  num_total_tlev_bytes);960 961    // Run-length decode the buffers' sizes into a window buffer of limited size. This is repeated962    // until we were able to cover all the bytes of TLEV buffers963    uint32_t decoded_window_offset = 0U;964    while (decoded_window_offset < num_total_tlev_bytes)965    {966      BlockBufferOffsetT buffer_id[TLEV_BYTES_PER_THREAD];967      TLevBufferSizeT buffer_byte_offset[TLEV_BYTES_PER_THREAD];968 969      // Now we have a balanced assignment: buffer_id[i] will hold the tile's buffer id and970      // buffer_byte_offset[i] that buffer's byte that this thread supposed to copy971      block_run_length_decode.RunLengthDecode(buffer_id, buffer_byte_offset, decoded_window_offset);972 973      // Zip from SoA to AoS974      ZippedTLevByteAssignment zipped_byte_assignment[TLEV_BYTES_PER_THREAD];975#pragma unroll976      for (int32_t i = 0; i < TLEV_BYTES_PER_THREAD; i++)977      {978        zipped_byte_assignment[i] = {buffer_id[i], buffer_byte_offset[i]};979      }980 981      // Exchange from blocked to striped arrangement for coalesced memory reads and writes982      BlockExchangeTLevT(temp_storage.staged.tlev.block_exchange_storage)983        .BlockedToStriped(zipped_byte_assignment, zipped_byte_assignment);984 985      // Read in the bytes that this thread is assigned to986      constexpr uint32_t WINDOW_SIZE = (TLEV_BYTES_PER_THREAD * BLOCK_THREADS);987      const bool is_full_window      = decoded_window_offset + WINDOW_SIZE < num_total_tlev_bytes;988      if (is_full_window)989      {990        uint32_t absolute_tlev_byte_offset = decoded_window_offset + threadIdx.x;991        AliasT src_byte[TLEV_BYTES_PER_THREAD];992#pragma unroll993        for (int32_t i = 0; i < TLEV_BYTES_PER_THREAD; i++)994        {995          src_byte[i] = read_item<IsMemcpy, AliasT, InputBufferT>(996            tile_buffer_srcs[zipped_byte_assignment[i].tile_buffer_id],997            zipped_byte_assignment[i].buffer_byte_offset);998          absolute_tlev_byte_offset += BLOCK_THREADS;999        }1000#pragma unroll1001        for (int32_t i = 0; i < TLEV_BYTES_PER_THREAD; i++)1002        {1003          write_item<IsMemcpy, AliasT, OutputBufferT>(1004            tile_buffer_dsts[zipped_byte_assignment[i].tile_buffer_id],1005            zipped_byte_assignment[i].buffer_byte_offset,1006            src_byte[i]);1007        }1008      }1009      else1010      {1011        uint32_t absolute_tlev_byte_offset = decoded_window_offset + threadIdx.x;1012#pragma unroll1013        for (int32_t i = 0; i < TLEV_BYTES_PER_THREAD; i++)1014        {1015          if (absolute_tlev_byte_offset < num_total_tlev_bytes)1016          {1017            const AliasT src_byte = read_item<IsMemcpy, AliasT, InputBufferT>(1018              tile_buffer_srcs[zipped_byte_assignment[i].tile_buffer_id],1019              zipped_byte_assignment[i].buffer_byte_offset);1020            write_item<IsMemcpy, AliasT, OutputBufferT>(1021              tile_buffer_dsts[zipped_byte_assignment[i].tile_buffer_id],1022              zipped_byte_assignment[i].buffer_byte_offset,1023              src_byte);1024          }1025          absolute_tlev_byte_offset += BLOCK_THREADS;1026        }1027      }1028 1029      decoded_window_offset += WINDOW_SIZE;1030 1031      // Ensure all threads finished collaborative BlockExchange so temporary storage can be reused1032      // with next iteration1033      CTA_SYNC();1034    }1035  }1036 1037  //-----------------------------------------------------------------------------1038  // PUBLIC MEMBER FUNCTIONS1039  //-----------------------------------------------------------------------------1040public:1041  __device__ __forceinline__ void ConsumeTile(BufferOffsetT tile_id)1042  {1043    // Offset into this tile's buffers1044    BufferOffsetT buffer_offset = tile_id * BUFFERS_PER_BLOCK;1045 1046    // Indicates whether all of this tiles items are within bounds1047    bool is_full_tile = buffer_offset + BUFFERS_PER_BLOCK < num_buffers;1048 1049    // Load the buffer sizes of this tile's buffers1050    BufferSizeIteratorT tile_buffer_sizes_it = buffer_sizes_it + buffer_offset;1051    BufferSizeT buffer_sizes[BUFFERS_PER_THREAD];1052    if (is_full_tile)1053    {1054      LoadBufferSizesFullTile(tile_buffer_sizes_it, buffer_sizes);1055    }1056    else1057    {1058      LoadBufferSizesPartialTile(tile_buffer_sizes_it, buffer_sizes, num_buffers - buffer_offset);1059    }1060 1061    // Ensure we can repurpose the BlockLoad's temporary storage1062    CTA_SYNC();1063 1064    // Count how many buffers fall into each size-class1065    VectorizedSizeClassCounterT size_class_histogram = GetBufferSizeClassHistogram(buffer_sizes);1066 1067    // Compute the prefix sum over the histogram1068    VectorizedSizeClassCounterT size_class_agg = {};1069    BlockSizeClassScanT(temp_storage.size_scan_storage)1070      .ExclusiveSum(size_class_histogram, size_class_histogram, size_class_agg);1071 1072    // Ensure we can repurpose the scan's temporary storage for scattering the buffer ids1073    CTA_SYNC();1074 1075    // Factor in the per-size-class counts / offsets1076    // That is, WLEV buffer offset has to be offset by the TLEV buffer count and BLEV buffer offset1077    // has to be offset by the TLEV+WLEV buffer count1078    uint32_t buffer_count = 0U;1079    for (uint32_t i = 0; i < NUM_SIZE_CLASSES; i++)1080    {1081      size_class_histogram.Add(i, buffer_count);1082      buffer_count += size_class_agg.Get(i);1083    }1084 1085    // Signal the number of BLEV buffers we're planning to write out1086    BufferOffsetT buffer_exclusive_prefix = 0;1087    if (tile_id == 0)1088    {1089      if (threadIdx.x == 0)1090      {1091        blev_buffer_scan_state.SetInclusive(tile_id, size_class_agg.Get(BLEV_SIZE_CLASS));1092      }1093      buffer_exclusive_prefix = 0;1094    }1095    else1096    {1097      BLevBuffScanPrefixCallbackOpT blev_buffer_prefix_op(blev_buffer_scan_state,1098                                                          temp_storage.buffer_scan_callback,1099                                                          Sum(),1100                                                          tile_id);1101 1102      // Signal our partial prefix and wait for the inclusive prefix of previous tiles1103      if (threadIdx.x < CUB_PTX_WARP_THREADS)1104      {1105        buffer_exclusive_prefix = blev_buffer_prefix_op(size_class_agg.Get(BLEV_SIZE_CLASS));1106      }1107    }1108    if (threadIdx.x == 0)1109    {1110      temp_storage.blev_buffer_offset = buffer_exclusive_prefix;1111    }1112 1113    // Ensure the prefix callback has finished using its temporary storage and that it can be reused1114    // in the next stage1115    CTA_SYNC();1116 1117    // Scatter the buffers into one of the three partitions (TLEV, WLEV, BLEV) depending on their1118    // size1119    PartitionBuffersBySize(buffer_sizes,1120                           size_class_histogram,1121                           temp_storage.staged.buffers_by_size_class);1122 1123    // Ensure all buffers have been partitioned by their size class AND1124    // ensure that blev_buffer_offset has been written to shared memory1125    CTA_SYNC();1126 1127    // TODO: think about prefetching tile_buffer_{srcs,dsts} into shmem1128    InputBufferIt tile_buffer_srcs        = input_buffer_it + buffer_offset;1129    OutputBufferIt tile_buffer_dsts       = output_buffer_it + buffer_offset;1130    BufferSizeIteratorT tile_buffer_sizes = buffer_sizes_it + buffer_offset;1131 1132    // Copy block-level buffers1133    EnqueueBLEVBuffers(1134      &temp_storage.staged.buffers_by_size_class[size_class_agg.Get(TLEV_SIZE_CLASS) +1135                                                 size_class_agg.Get(WLEV_SIZE_CLASS)],1136      tile_buffer_srcs,1137      tile_buffer_dsts,1138      tile_buffer_sizes,1139      size_class_agg.Get(BLEV_SIZE_CLASS),1140      temp_storage.blev_buffer_offset,1141      tile_id);1142 1143    // Ensure we can repurpose the temporary storage required by EnqueueBLEVBuffers1144    CTA_SYNC();1145 1146    // Copy warp-level buffers1147    BatchMemcpyWLEVBuffers(1148      &temp_storage.staged.buffers_by_size_class[size_class_agg.Get(TLEV_SIZE_CLASS)],1149      tile_buffer_srcs,1150      tile_buffer_dsts,1151      tile_buffer_sizes,1152      size_class_agg.Get(WLEV_SIZE_CLASS));1153 1154    // Perform batch memcpy for all the buffers that require thread-level collaboration1155    uint32_t num_tlev_buffers = size_class_agg.Get(TLEV_SIZE_CLASS);1156    BatchMemcpyTLEVBuffers(temp_storage.staged.buffers_by_size_class,1157                           tile_buffer_srcs,1158                           tile_buffer_dsts,1159                           num_tlev_buffers);1160  }1161 1162  //-----------------------------------------------------------------------------1163  // CONSTRUCTOR1164  //-----------------------------------------------------------------------------1165  __device__ __forceinline__ AgentBatchMemcpy(TempStorage &temp_storage,1166                                              InputBufferIt input_buffer_it,1167                                              OutputBufferIt output_buffer_it,1168                                              BufferSizeIteratorT buffer_sizes_it,1169                                              BufferOffsetT num_buffers,1170                                              BlevBufferSrcsOutItT blev_buffer_srcs,1171                                              BlevBufferDstsOutItT blev_buffer_dsts,1172                                              BlevBufferSizesOutItT blev_buffer_sizes,1173                                              BlevBufferTileOffsetsOutItT blev_buffer_tile_offsets,1174                                              BLevBufferOffsetTileState blev_buffer_scan_state,1175                                              BLevBlockOffsetTileState blev_block_scan_state)1176      : temp_storage(temp_storage.Alias())1177      , input_buffer_it(input_buffer_it)1178      , output_buffer_it(output_buffer_it)1179      , buffer_sizes_it(buffer_sizes_it)1180      , num_buffers(num_buffers)1181      , blev_buffer_srcs(blev_buffer_srcs)1182      , blev_buffer_dsts(blev_buffer_dsts)1183      , blev_buffer_sizes(blev_buffer_sizes)1184      , blev_buffer_tile_offsets(blev_buffer_tile_offsets)1185      , blev_buffer_scan_state(blev_buffer_scan_state)1186      , blev_block_scan_state(blev_block_scan_state)1187  {}1188 1189private:1190  // Iterator providing the pointers to the source memory buffers1191  InputBufferIt input_buffer_it;1192  // Iterator providing the pointers to the destination memory buffers1193  OutputBufferIt output_buffer_it;1194  // Iterator providing the number of bytes to be copied for each pair of buffers1195  BufferSizeIteratorT buffer_sizes_it;1196  // The total number of buffer pairs1197  BufferOffsetT num_buffers;1198  // Output iterator to which the source pointers of the BLEV buffers are written1199  BlevBufferSrcsOutItT blev_buffer_srcs;1200  // Output iterator to which the destination pointers of the BLEV buffers are written

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

codekingpro/portable-devtools · Team Ai