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