[INFRA] Import NVIDIA/CCCL upstream as optimization reference library
CCCL (CUDA C++ Core Libraries) provides: - CUB: device/block/warp-level GPU primitives (reduce, scan, sort, topk) - Thrust: high-level parallel algorithms (transform_reduce, sort, scan) - libcudacxx: CUDA C++ standard library (atomics, barriers, memory) - cudax: experimental features (memory resources, allocators) - Tuning policies: per-SM hardware-specific algorithm parameters Competition optimization vectors mapped to CCCL: - Output TPS (83% weight): warp_reduce, block_reduce, device_topk - Input TPS (14% weight): device_scan, block_load, prefetch - Cache TPS (3% weight): prefix caching strategy patterns - Memory (0.9 util): pooled/cached/buddy allocators Source: https://github.com/NVIDIA/cccl (shallow clone, HEAD only) License: Apache-2.0
This commit is contained in:
394
cccl_upstream/cub/test/warp/catch2_test_warp_reduce.cu
Normal file
394
cccl_upstream/cub/test/warp/catch2_test_warp_reduce.cu
Normal file
@@ -0,0 +1,394 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
#include <cub/util_macro.cuh>
|
||||
#include <cub/warp/warp_reduce.cuh>
|
||||
|
||||
#include <cuda/cmath>
|
||||
#include <cuda/functional>
|
||||
#include <cuda/iterator>
|
||||
#include <cuda/ptx>
|
||||
#include <cuda/std/__functional/invoke.h>
|
||||
#include <cuda/std/functional>
|
||||
#include <cuda/std/limits>
|
||||
#include <cuda/std/type_traits>
|
||||
|
||||
#include <array>
|
||||
#include <numeric>
|
||||
|
||||
#include <test_util.h>
|
||||
|
||||
#include <c2h/catch2_test_helper.h>
|
||||
#include <c2h/check_results.cuh>
|
||||
#include <c2h/custom_type.h>
|
||||
#include <c2h/operator.cuh>
|
||||
|
||||
/***********************************************************************************************************************
|
||||
* Constants
|
||||
**********************************************************************************************************************/
|
||||
|
||||
inline constexpr int warp_size = 32;
|
||||
inline constexpr auto total_warps = 4u;
|
||||
inline constexpr int num_items_per_thread = 4;
|
||||
|
||||
/***********************************************************************************************************************
|
||||
* Kernel
|
||||
**********************************************************************************************************************/
|
||||
|
||||
template <unsigned LogicalWarpThreads, bool EnableNumItems = false, typename T, typename Output, typename ReductionOp>
|
||||
__device__ void warp_reduce_function(T& thread_data, Output* output, ReductionOp reduction_op, int num_items = 0)
|
||||
{
|
||||
using warp_reduce_t = cub::WarpReduce<Output, LogicalWarpThreads>;
|
||||
using storage_t = typename warp_reduce_t::TempStorage;
|
||||
__shared__ storage_t storage[total_warps];
|
||||
constexpr bool is_power_of_two = cuda::is_power_of_two(LogicalWarpThreads);
|
||||
auto lane = cuda::ptx::get_sreg_laneid();
|
||||
auto logical_warp = is_power_of_two ? threadIdx.x / LogicalWarpThreads : threadIdx.x / warp_size;
|
||||
auto logical_lane = is_power_of_two ? threadIdx.x % LogicalWarpThreads : lane;
|
||||
auto limit = EnableNumItems ? num_items : LogicalWarpThreads;
|
||||
if (!is_power_of_two && lane >= limit)
|
||||
{
|
||||
return;
|
||||
}
|
||||
warp_reduce_t warp_reduce{storage[logical_warp]};
|
||||
using result_t = decltype(reduction_op(warp_reduce, thread_data));
|
||||
result_t result;
|
||||
if constexpr (EnableNumItems)
|
||||
{
|
||||
result = reduction_op(warp_reduce, thread_data, num_items);
|
||||
}
|
||||
else
|
||||
{
|
||||
result = reduction_op(warp_reduce, thread_data);
|
||||
}
|
||||
if (logical_lane == 0)
|
||||
{
|
||||
output[logical_warp] = result;
|
||||
}
|
||||
}
|
||||
|
||||
template <unsigned LogicalWarpThreads, bool EnableNumItems, typename T, typename ReductionOp>
|
||||
__global__ void warp_reduce_kernel(T* input, T* output, ReductionOp reduction_op, int num_items = 0)
|
||||
{
|
||||
auto thread_data = input[threadIdx.x];
|
||||
warp_reduce_function<LogicalWarpThreads, EnableNumItems>(thread_data, output, reduction_op, num_items);
|
||||
}
|
||||
|
||||
template <unsigned LogicalWarpThreads, typename T, typename ReductionOp>
|
||||
__global__ void warp_reduce_multiple_items_kernel(T* input, T* output, ReductionOp reduction_op)
|
||||
{
|
||||
T thread_data[num_items_per_thread];
|
||||
for (int i = 0; i < num_items_per_thread; ++i)
|
||||
{
|
||||
thread_data[i] = input[threadIdx.x * num_items_per_thread + i];
|
||||
}
|
||||
warp_reduce_function<LogicalWarpThreads>(thread_data, output, reduction_op);
|
||||
}
|
||||
|
||||
template <typename Op, typename T>
|
||||
struct warp_reduce_t
|
||||
{
|
||||
template <int LogicalWarpThreads>
|
||||
__device__ auto operator()(cub::WarpReduce<T, LogicalWarpThreads> warp_reduce, T& data) const
|
||||
{
|
||||
return warp_reduce.Reduce(data, Op{});
|
||||
}
|
||||
|
||||
template <int LogicalWarpThreads>
|
||||
__device__ auto operator()(cub::WarpReduce<T, LogicalWarpThreads> warp_reduce, T& data, int num_items) const
|
||||
{
|
||||
return warp_reduce.Reduce(data, Op{}, num_items);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct warp_reduce_t<cuda::std::plus<>, T>
|
||||
{
|
||||
template <int LogicalWarpThreads, typename... TArgs>
|
||||
__device__ auto operator()(cub::WarpReduce<T, LogicalWarpThreads> warp_reduce, TArgs&&... args) const
|
||||
{
|
||||
return warp_reduce.Sum(args...);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct warp_reduce_t<cuda::maximum<>, T>
|
||||
{
|
||||
template <int LogicalWarpThreads, typename... TArgs>
|
||||
__device__ auto operator()(cub::WarpReduce<T, LogicalWarpThreads> warp_reduce, TArgs&&... args) const
|
||||
{
|
||||
return warp_reduce.Max(args...);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct warp_reduce_t<cuda::minimum<>, T>
|
||||
{
|
||||
template <int LogicalWarpThreads, typename... TArgs>
|
||||
__device__ auto operator()(cub::WarpReduce<T, LogicalWarpThreads> warp_reduce, TArgs&&... args) const
|
||||
{
|
||||
return warp_reduce.Min(args...);
|
||||
}
|
||||
};
|
||||
|
||||
template <int LogicalWarpThreads, bool EnableNumItems = false, typename T, typename... TArgs>
|
||||
void warp_reduce_launch(c2h::device_vector<T>& input, c2h::device_vector<T>& output, TArgs... args)
|
||||
{
|
||||
warp_reduce_kernel<LogicalWarpThreads, EnableNumItems><<<1, total_warps * warp_size>>>(
|
||||
thrust::raw_pointer_cast(input.data()), thrust::raw_pointer_cast(output.data()), args...);
|
||||
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
}
|
||||
|
||||
template <int LogicalWarpThreads, typename T, typename... TArgs>
|
||||
void warp_reduce_multiple_items_launch(c2h::device_vector<T>& input, c2h::device_vector<T>& output, TArgs... args)
|
||||
{
|
||||
warp_reduce_multiple_items_kernel<LogicalWarpThreads><<<1, total_warps * warp_size>>>(
|
||||
thrust::raw_pointer_cast(input.data()), thrust::raw_pointer_cast(output.data()), args...);
|
||||
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
}
|
||||
|
||||
/***********************************************************************************************************************
|
||||
* Types
|
||||
**********************************************************************************************************************/
|
||||
|
||||
using custom_t =
|
||||
c2h::custom_type_t<c2h::accumulateable_t, c2h::equal_comparable_t, c2h::lexicographical_less_comparable_t>;
|
||||
|
||||
using full_type_list =
|
||||
c2h::type_list<uint8_t,
|
||||
uint16_t,
|
||||
int32_t,
|
||||
int64_t,
|
||||
custom_t,
|
||||
#if _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4_16a,
|
||||
#else // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4,
|
||||
#endif // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
uchar3,
|
||||
short2>;
|
||||
|
||||
using builtin_type_list = c2h::type_list<uint8_t, uint16_t, int32_t, int64_t>;
|
||||
|
||||
// clang-format off
|
||||
using floating_point_redux_type_list = c2h::type_list<
|
||||
float
|
||||
#if TEST_HALF_T()
|
||||
, __half
|
||||
#endif // TEST_HALF_T()
|
||||
#if TEST_BF_T()
|
||||
, __nv_bfloat16
|
||||
#endif // TEST_BF_T()
|
||||
>;
|
||||
// clang-format on
|
||||
|
||||
using predefined_op_list = c2h::type_list<cuda::std::plus<>, cuda::maximum<>, cuda::minimum<>>;
|
||||
|
||||
using predefined_min_max_op_list = c2h::type_list<cuda::maximum<>, cuda::minimum<>>;
|
||||
|
||||
using logical_warp_threads = c2h::enum_type_list<unsigned, 32, 16, 9, 7, 1>;
|
||||
|
||||
/***********************************************************************************************************************
|
||||
* Reference
|
||||
**********************************************************************************************************************/
|
||||
|
||||
_CCCL_DIAG_PUSH
|
||||
_CCCL_DIAG_SUPPRESS_MSVC(4244) // numeric(33): C: '=': conversion from 'int' to '_Ty', possible loss of data
|
||||
|
||||
template <typename predefined_op, typename T>
|
||||
void compute_host_reference(
|
||||
const c2h::host_vector<T>& h_in,
|
||||
c2h::host_vector<T>& h_out,
|
||||
int logical_warps,
|
||||
int logical_warp_threads,
|
||||
int items_per_logical_warp = 0,
|
||||
int items_per_thread = 1)
|
||||
{
|
||||
const auto identity = identity_v<predefined_op, T>;
|
||||
items_per_logical_warp = items_per_logical_warp == 0 ? logical_warp_threads : items_per_logical_warp;
|
||||
for (unsigned i = 0; i < total_warps; ++i)
|
||||
{
|
||||
for (int j = 0; j < logical_warps; ++j)
|
||||
{
|
||||
auto start =
|
||||
h_in.begin()
|
||||
+ (i * warp_size + j * logical_warp_threads) * items_per_thread; // NOLINT(bugprone-misplaced-widening-cast)
|
||||
auto end = start + static_cast<long>(items_per_logical_warp) * items_per_thread;
|
||||
// NOLINTNEXTLINE(bugprone-misplaced-widening-cast)
|
||||
h_out[i * logical_warps + j] = static_cast<T>(std::accumulate(start, end, identity, predefined_op{}));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_CCCL_DIAG_POP
|
||||
|
||||
std::array<unsigned, 3> get_test_config(unsigned logical_warp_threads, unsigned items_per_thread = 1)
|
||||
{
|
||||
bool is_power_of_two = cuda::is_power_of_two(logical_warp_threads);
|
||||
auto logical_warps = is_power_of_two ? warp_size / logical_warp_threads : 1;
|
||||
auto input_size = total_warps * warp_size * items_per_thread;
|
||||
auto output_size = total_warps * logical_warps;
|
||||
return {input_size, output_size, logical_warps};
|
||||
}
|
||||
|
||||
/***********************************************************************************************************************
|
||||
* Test cases
|
||||
**********************************************************************************************************************/
|
||||
|
||||
C2H_TEST("WarpReduce::Sum, full_type_list", "[reduce][warp][predefined_op][full]", full_type_list, logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<1, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
CAPTURE(c2h::type_name<T>(), c2h::type_name<T>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_launch<logical_warp_threads>(d_in, d_out, warp_reduce_t<cuda::std::plus<>, T>{});
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<cuda::std::plus<>>(h_in, h_out, logical_warps, logical_warp_threads);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduce::Sum/Max/Min, builtin types",
|
||||
"[reduce][warp][predefined_op][full]",
|
||||
builtin_type_list,
|
||||
predefined_op_list,
|
||||
logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
using predefined_op = c2h::get<1, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<2, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
CAPTURE(c2h::type_name<T>(), c2h::type_name<predefined_op>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_launch<logical_warp_threads>(d_in, d_out, warp_reduce_t<predefined_op, T>{});
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<predefined_op>(h_in, h_out, logical_warps, logical_warp_threads);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduce::Max/Min, floating-point redux types",
|
||||
"[reduce][warp][predefined_op][redux]",
|
||||
floating_point_redux_type_list,
|
||||
predefined_min_max_op_list,
|
||||
logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
using predefined_op = c2h::get<1, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<2, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
CAPTURE(c2h::type_name<T>(), c2h::type_name<predefined_op>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_launch<logical_warp_threads>(d_in, d_out, warp_reduce_t<predefined_op, T>{});
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<predefined_op>(h_in, h_out, logical_warps, logical_warp_threads);
|
||||
if constexpr (cuda::std::is_same_v<T, float>)
|
||||
{
|
||||
verify_results_exact(h_out, d_out);
|
||||
}
|
||||
else
|
||||
{
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduce::CustomSum", "[reduce][warp][generic][full]", full_type_list, logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<1, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
CAPTURE(c2h::type_name<T>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(1), d_in);
|
||||
warp_reduce_launch<logical_warp_threads>(d_in, d_out, warp_reduce_t<custom_plus, T>{});
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<cuda::std::plus<>>(h_in, h_out, logical_warps, logical_warp_threads);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
//----------------------------------------------------------------------------------------------------------------------
|
||||
// partial
|
||||
|
||||
C2H_TEST("WarpReduce::Sum/Max/Min Partial",
|
||||
"[reduce][warp][predefined_op][partial]",
|
||||
builtin_type_list,
|
||||
predefined_op_list,
|
||||
logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
using predefined_op = c2h::get<1, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<2, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
const int valid_items = GENERATE_COPY(take(2, random(1u, logical_warp_threads)));
|
||||
CAPTURE(c2h::type_name<T>(), c2h::type_name<predefined_op>(), logical_warp_threads, valid_items);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_launch<logical_warp_threads, true>(d_in, d_out, warp_reduce_t<predefined_op, T>{}, valid_items);
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<predefined_op>(h_in, h_out, logical_warps, logical_warp_threads, valid_items);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduce::Sum", "[reduce][warp][generic][partial]", full_type_list, logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<1, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads);
|
||||
const int valid_items = GENERATE_COPY(take(2, random(1u, logical_warp_threads)));
|
||||
CAPTURE(c2h::type_name<T>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_launch<logical_warp_threads, true>(d_in, d_out, warp_reduce_t<custom_plus, T>{}, valid_items);
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<cuda::std::plus<>>(h_in, h_out, logical_warps, logical_warp_threads, valid_items);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
//----------------------------------------------------------------------------------------------------------------------
|
||||
// multiple items per thread
|
||||
|
||||
C2H_TEST("WarpReduce::Sum/Max/Min Multiple Items Per Thread",
|
||||
"[reduce][warp][predefined_op][full]",
|
||||
builtin_type_list,
|
||||
predefined_op_list,
|
||||
logical_warp_threads)
|
||||
{
|
||||
using T = c2h::get<0, TestType>;
|
||||
using predefined_op = c2h::get<1, TestType>;
|
||||
constexpr auto logical_warp_threads = c2h::get<2, TestType>::value;
|
||||
auto [input_size, output_size, logical_warps] = get_test_config(logical_warp_threads, num_items_per_thread);
|
||||
CAPTURE(c2h::type_name<T>(), c2h::type_name<predefined_op>(), logical_warp_threads);
|
||||
c2h::device_vector<T> d_in(input_size);
|
||||
c2h::device_vector<T> d_out(output_size);
|
||||
c2h::gen(C2H_SEED(10), d_in);
|
||||
warp_reduce_multiple_items_launch<logical_warp_threads>(d_in, d_out, warp_reduce_t<predefined_op, T>{});
|
||||
|
||||
c2h::host_vector<T> h_in = d_in;
|
||||
c2h::host_vector<T> h_out(output_size);
|
||||
compute_host_reference<predefined_op>(h_in, h_out, logical_warps, logical_warp_threads, 0, num_items_per_thread);
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
762
cccl_upstream/cub/test/warp/catch2_test_warp_reduce_batched.cu
Normal file
762
cccl_upstream/cub/test/warp/catch2_test_warp_reduce_batched.cu
Normal file
@@ -0,0 +1,762 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
|
||||
|
||||
#include <cub/util_macro.cuh>
|
||||
#include <cub/warp/warp_reduce.cuh>
|
||||
#include <cub/warp/warp_reduce_batched.cuh>
|
||||
|
||||
#include <cuda/cmath>
|
||||
#include <cuda/functional>
|
||||
#include <cuda/iterator>
|
||||
#include <cuda/ptx>
|
||||
#include <cuda/std/functional>
|
||||
#include <cuda/std/limits>
|
||||
#include <cuda/std/mdspan>
|
||||
#include <cuda/std/span>
|
||||
#include <cuda/std/type_traits>
|
||||
|
||||
#include <numeric>
|
||||
|
||||
#include <test_util.h>
|
||||
|
||||
#include <c2h/catch2_test_helper.h>
|
||||
#include <c2h/check_results.cuh>
|
||||
#include <c2h/custom_type.h>
|
||||
#include <c2h/operator.cuh>
|
||||
|
||||
// %PARAM% TEST_TYPES types 0:1:2
|
||||
|
||||
inline constexpr int warp_size = 32;
|
||||
inline constexpr int block_size = 2 * warp_size;
|
||||
|
||||
// 2D layout: num_batches x batch_size (row-major). Single type with dynamic extents for kernels and host.
|
||||
template <typename T>
|
||||
using input_2d_mdspan_t = cuda::std::mdspan<T, cuda::std::dextents<int, 2>>;
|
||||
|
||||
enum class WarpReduceBatchedMode
|
||||
{
|
||||
SingleOut,
|
||||
ToStriped,
|
||||
ToBlocked
|
||||
};
|
||||
|
||||
template <WarpReduceBatchedMode Mode,
|
||||
bool SyncPhysicalWarp,
|
||||
int Batches,
|
||||
int LogicalWarpThreads,
|
||||
typename T,
|
||||
typename ReductionOp>
|
||||
__global__ void
|
||||
warp_reduce_batched_kernel(input_2d_mdspan_t<T> input_md, cuda::std::span<T> output, ReductionOp reduction_op)
|
||||
{
|
||||
using warp_reduce_batched_t = cub::WarpReduceBatched<T, Batches, LogicalWarpThreads, SyncPhysicalWarp>;
|
||||
__shared__ typename warp_reduce_batched_t::TempStorage temp_storage[block_size / LogicalWarpThreads];
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / LogicalWarpThreads;
|
||||
const int lane_id = tid % LogicalWarpThreads;
|
||||
|
||||
auto inputs = cuda::std::array<T, Batches>{};
|
||||
for (int batch_idx = 0; batch_idx < Batches; ++batch_idx)
|
||||
{
|
||||
inputs[batch_idx] = input_md(logical_warp_id * Batches + batch_idx, lane_id);
|
||||
}
|
||||
|
||||
constexpr int out_per_thread = cuda::ceil_div(Batches, LogicalWarpThreads);
|
||||
auto outputs = cuda::std::array<T, out_per_thread>{};
|
||||
if constexpr (Mode == WarpReduceBatchedMode::SingleOut)
|
||||
{
|
||||
outputs[0] = warp_reduce_batched_t{temp_storage[logical_warp_id]}.Reduce(inputs, reduction_op);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.ReduceToBlocked(inputs, outputs, reduction_op);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToStriped)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.ReduceToStriped(inputs, outputs, reduction_op);
|
||||
}
|
||||
|
||||
for (int idx = 0; idx < out_per_thread; ++idx)
|
||||
{
|
||||
const auto batch_idx =
|
||||
(Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
? (idx + lane_id * out_per_thread)
|
||||
: (idx * LogicalWarpThreads + lane_id);
|
||||
if (batch_idx < Batches)
|
||||
{
|
||||
output[static_cast<std::size_t>(logical_warp_id) * Batches + batch_idx] = outputs[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <WarpReduceBatchedMode Mode, bool SyncPhysicalWarp, int Batches, int LogicalWarpThreads, typename T>
|
||||
__global__ void sum_batched_kernel(input_2d_mdspan_t<T> input_md, cuda::std::span<T> output)
|
||||
{
|
||||
using warp_reduce_batched_t = cub::WarpReduceBatched<T, Batches, LogicalWarpThreads, SyncPhysicalWarp>;
|
||||
__shared__ typename warp_reduce_batched_t::TempStorage temp_storage[block_size / LogicalWarpThreads];
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / LogicalWarpThreads;
|
||||
const int lane_id = tid % LogicalWarpThreads;
|
||||
|
||||
auto inputs = cuda::std::array<T, Batches>{};
|
||||
for (int batch_idx = 0; batch_idx < Batches; ++batch_idx)
|
||||
{
|
||||
inputs[batch_idx] = input_md(logical_warp_id * Batches + batch_idx, lane_id);
|
||||
}
|
||||
|
||||
constexpr int out_per_thread = cuda::ceil_div(Batches, LogicalWarpThreads);
|
||||
auto outputs = cuda::std::array<T, out_per_thread>{};
|
||||
if constexpr (Mode == WarpReduceBatchedMode::SingleOut)
|
||||
{
|
||||
outputs[0] = warp_reduce_batched_t{temp_storage[logical_warp_id]}.Sum(inputs);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.SumToBlocked(inputs, outputs);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToStriped)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.SumToStriped(inputs, outputs);
|
||||
}
|
||||
|
||||
for (int idx = 0; idx < out_per_thread; ++idx)
|
||||
{
|
||||
const auto batch_idx =
|
||||
(Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
? (idx + lane_id * out_per_thread)
|
||||
: (idx * LogicalWarpThreads + lane_id);
|
||||
if (batch_idx < Batches)
|
||||
{
|
||||
output[static_cast<std::size_t>(logical_warp_id) * Batches + batch_idx] = outputs[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <WarpReduceBatchedMode Mode,
|
||||
bool SyncPhysicalWarp,
|
||||
int Batches,
|
||||
int LogicalWarpThreads,
|
||||
typename T,
|
||||
typename ReductionOp>
|
||||
__global__ void
|
||||
warp_reduce_batched_cond_part_kernel(input_2d_mdspan_t<T> input_md, cuda::std::span<T> output, ReductionOp reduction_op)
|
||||
{
|
||||
static_assert(LogicalWarpThreads < warp_size, "Need at least 2 logical warps per physical warp for this test");
|
||||
using warp_reduce_batched_t = cub::WarpReduceBatched<T, Batches, LogicalWarpThreads>;
|
||||
__shared__ typename warp_reduce_batched_t::TempStorage temp_storage[block_size / LogicalWarpThreads];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / LogicalWarpThreads;
|
||||
const int lane_id = tid % LogicalWarpThreads;
|
||||
const bool is_participating = (logical_warp_id % 2 == 0);
|
||||
|
||||
if (is_participating)
|
||||
{
|
||||
const int participant_idx = logical_warp_id / 2;
|
||||
|
||||
auto inputs = cuda::std::array<T, Batches>{};
|
||||
for (int batch_idx = 0; batch_idx < Batches; ++batch_idx)
|
||||
{
|
||||
inputs[batch_idx] = input_md(participant_idx * Batches + batch_idx, lane_id);
|
||||
}
|
||||
|
||||
constexpr int out_per_thread = cuda::ceil_div(Batches, LogicalWarpThreads);
|
||||
auto outputs = cuda::std::array<T, out_per_thread>{};
|
||||
if constexpr (Mode == WarpReduceBatchedMode::SingleOut)
|
||||
{
|
||||
outputs[0] = warp_reduce_batched_t{temp_storage[logical_warp_id]}.Reduce(inputs, reduction_op);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.ReduceToBlocked(inputs, outputs, reduction_op);
|
||||
}
|
||||
if constexpr (Mode == WarpReduceBatchedMode::ToStriped)
|
||||
{
|
||||
warp_reduce_batched_t{temp_storage[logical_warp_id]}.ReduceToStriped(inputs, outputs, reduction_op);
|
||||
}
|
||||
|
||||
for (int idx = 0; idx < out_per_thread; ++idx)
|
||||
{
|
||||
const auto batch_idx =
|
||||
(Mode == WarpReduceBatchedMode::ToBlocked)
|
||||
? (idx + lane_id * out_per_thread)
|
||||
: (idx * LogicalWarpThreads + lane_id);
|
||||
if (batch_idx < Batches)
|
||||
{
|
||||
output[static_cast<std::size_t>(participant_idx) * Batches + batch_idx] = outputs[idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
// When SyncPhysicalWarp is true, the non-participating threads need to exit early to avoid expected deadlocks.
|
||||
if constexpr (!SyncPhysicalWarp)
|
||||
{
|
||||
// Keep non-participating threads from exiting early to check for unexpected deadlocks.
|
||||
__syncwarp();
|
||||
}
|
||||
}
|
||||
|
||||
_CCCL_DIAG_PUSH
|
||||
_CCCL_DIAG_SUPPRESS_MSVC(4244) // numeric(33): C: '=': conversion from 'int' to '_Ty', possible loss of data
|
||||
|
||||
template <typename T, typename ReductionOp>
|
||||
void compute_host_reference(input_2d_mdspan_t<const T> input_md, cuda::std::span<T> output, ReductionOp op)
|
||||
{
|
||||
const auto identity = identity_v<ReductionOp, T>;
|
||||
const int batches = input_md.extent(0);
|
||||
const int batch_size = input_md.extent(1);
|
||||
|
||||
for (int batch_idx = 0; batch_idx < batches; ++batch_idx)
|
||||
{
|
||||
const auto iter = cuda::make_transform_iterator(cuda::make_counting_iterator(0), [input_md, batch_idx](int idx) {
|
||||
return input_md(batch_idx, idx);
|
||||
});
|
||||
output[batch_idx] = static_cast<T>(std::accumulate(iter, iter + batch_size, identity, op));
|
||||
}
|
||||
}
|
||||
|
||||
_CCCL_DIAG_POP
|
||||
|
||||
template <typename T, int N>
|
||||
void gen_bounded_input(c2h::seed_t seed, c2h::device_vector<T>& d_input)
|
||||
{
|
||||
if constexpr (cuda::std::is_floating_point_v<T>)
|
||||
{
|
||||
// Small positive range to minimize floating point error in reductions
|
||||
c2h::gen(seed, d_input, T(0.5), T(1.5));
|
||||
}
|
||||
else if constexpr (cuda::std::is_integral_v<T> && cuda::std::is_signed_v<T>)
|
||||
{
|
||||
// Avoid signed overflow when summing N elements
|
||||
const T gen_max = cuda::std::numeric_limits<T>::max() / static_cast<T>(N);
|
||||
c2h::gen(seed, d_input, -gen_max, gen_max);
|
||||
}
|
||||
else
|
||||
{
|
||||
c2h::gen(seed, d_input);
|
||||
}
|
||||
}
|
||||
|
||||
template <WarpReduceBatchedMode Mode,
|
||||
int Batches,
|
||||
int LogicalWarpThreads,
|
||||
typename T,
|
||||
typename ReductionOp,
|
||||
bool ConvenienceOverload = false,
|
||||
bool CondParticipation = false,
|
||||
bool SyncPhysicalWarp = false>
|
||||
void test_warp_reduce_batched(ReductionOp reduction_op = ReductionOp{})
|
||||
{
|
||||
CAPTURE(c2h::type_name<T>(), Batches, LogicalWarpThreads);
|
||||
|
||||
constexpr int num_logical_warps = block_size / LogicalWarpThreads;
|
||||
constexpr int total_batches = CondParticipation ? num_logical_warps * Batches / 2 : num_logical_warps * Batches;
|
||||
constexpr int total_elements = total_batches * LogicalWarpThreads;
|
||||
|
||||
c2h::device_vector<T> d_input(total_elements);
|
||||
gen_bounded_input<T, LogicalWarpThreads>(C2H_SEED(10), d_input);
|
||||
|
||||
c2h::device_vector<T> d_output(total_batches);
|
||||
|
||||
input_2d_mdspan_t<T> d_input_md(thrust::raw_pointer_cast(d_input.data()), total_batches, LogicalWarpThreads);
|
||||
cuda::std::span<T> d_output_span(thrust::raw_pointer_cast(d_output.data()), total_batches);
|
||||
if constexpr (CondParticipation)
|
||||
{
|
||||
warp_reduce_batched_cond_part_kernel<Mode, SyncPhysicalWarp, Batches, LogicalWarpThreads>
|
||||
<<<1, block_size>>>(d_input_md, d_output_span, reduction_op);
|
||||
}
|
||||
else
|
||||
{
|
||||
if constexpr (ConvenienceOverload && cuda::std::is_same_v<ReductionOp, cuda::std::plus<>>)
|
||||
{
|
||||
sum_batched_kernel<Mode, SyncPhysicalWarp, Batches, LogicalWarpThreads>
|
||||
<<<1, block_size>>>(d_input_md, d_output_span);
|
||||
}
|
||||
else
|
||||
{
|
||||
warp_reduce_batched_kernel<Mode, SyncPhysicalWarp, Batches, LogicalWarpThreads>
|
||||
<<<1, block_size>>>(d_input_md, d_output_span, reduction_op);
|
||||
}
|
||||
}
|
||||
|
||||
cudaError_t err = cudaPeekAtLastError();
|
||||
REQUIRE(err == cudaSuccess);
|
||||
|
||||
err = cudaDeviceSynchronize();
|
||||
REQUIRE(err == cudaSuccess);
|
||||
|
||||
// Host-side: construct mdspans once; pass to reference and verify
|
||||
c2h::host_vector<T> h_input = d_input;
|
||||
c2h::host_vector<T> h_output = d_output;
|
||||
|
||||
c2h::host_vector<T> h_reference(total_batches);
|
||||
input_2d_mdspan_t<const T> h_input_md(h_input.data(), total_batches, LogicalWarpThreads);
|
||||
cuda::std::span<T> h_reference_span(h_reference.data(), total_batches);
|
||||
|
||||
compute_host_reference(h_input_md, h_reference_span, reduction_op);
|
||||
|
||||
verify_results(h_reference, h_output);
|
||||
}
|
||||
|
||||
#if TEST_TYPES == 0
|
||||
using builtin_type_list = c2h::type_list<cuda::std::uint8_t, cuda::std::uint16_t>;
|
||||
#elif TEST_TYPES == 1
|
||||
using builtin_type_list = c2h::type_list<cuda::std::int32_t, cuda::std::int64_t>;
|
||||
#elif TEST_TYPES == 2
|
||||
using builtin_type_list = c2h::type_list<float, double>;
|
||||
#endif
|
||||
|
||||
#if TEST_TYPES == 0
|
||||
using full_type_list = c2h::type_list<cuda::std::uint8_t, cuda::std::uint16_t, uchar3, short2>;
|
||||
#elif TEST_TYPES == 1
|
||||
using full_type_list =
|
||||
c2h::type_list<cuda::std::int32_t,
|
||||
cuda::std::int64_t,
|
||||
# if _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4_16a
|
||||
# else // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4
|
||||
# endif // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
>;
|
||||
#elif TEST_TYPES == 2
|
||||
using custom_t =
|
||||
c2h::custom_type_t<c2h::accumulateable_t, c2h::equal_comparable_t, c2h::lexicographical_less_comparable_t>;
|
||||
|
||||
using full_type_list = c2h::type_list<float, double, custom_t>;
|
||||
#endif // TEST_TYPES
|
||||
|
||||
// (N, M) = (Batches, LogicalWarpThreads) for test parameterization
|
||||
template <int X, int Y>
|
||||
struct int_pair
|
||||
{
|
||||
static constexpr int x = X;
|
||||
static constexpr int y = Y;
|
||||
};
|
||||
|
||||
// N=M configurations (best performance)
|
||||
using equal_nm_configs =
|
||||
c2h::type_list<int_pair<1, 1>, int_pair<2, 2>, int_pair<4, 4>, int_pair<8, 8>, int_pair<16, 16>, int_pair<32, 32>>;
|
||||
|
||||
// N!=M configurations
|
||||
using unequal_nm_configs = c2h::type_list<
|
||||
int_pair<0, 32>,
|
||||
int_pair<1, 32>,
|
||||
int_pair<3, 32>,
|
||||
int_pair<3, 16>,
|
||||
int_pair<5, 16>,
|
||||
int_pair<4, 8>,
|
||||
int_pair<5, 8>,
|
||||
int_pair<6, 8>,
|
||||
int_pair<6, 4>,
|
||||
int_pair<7, 4>,
|
||||
int_pair<8, 4>,
|
||||
int_pair<9, 4>,
|
||||
int_pair<10, 4>,
|
||||
int_pair<1, 2>,
|
||||
int_pair<7, 2>,
|
||||
int_pair<0, 1>,
|
||||
int_pair<2, 1>>;
|
||||
|
||||
// Sub-warp configurations (LogicalWarpThreads < 32, at least 2 logical warps per physical warp)
|
||||
using sub_warp_equal_configs = c2h::type_list<int_pair<2, 2>, int_pair<4, 4>, int_pair<8, 8>, int_pair<16, 16>>;
|
||||
|
||||
using sub_warp_unequal_configs = c2h::type_list<int_pair<3, 16>, int_pair<4, 8>, int_pair<6, 4>, int_pair<1, 2>>;
|
||||
|
||||
using unequal_nm_single_out_configs = c2h::type_list<int_pair<1, 32>, int_pair<3, 16>, int_pair<4, 8>, int_pair<1, 2>>;
|
||||
|
||||
C2H_TEST("WarpReduceBatched::Reduce N=M sum", "[warp][reduce][batched]", full_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::SingleOut, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST(
|
||||
"WarpReduceBatched::Reduce N!=M sum", "[warp][reduce][batched]", builtin_type_list, unequal_nm_single_out_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::SingleOut, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST(
|
||||
"WarpReduceBatched::Reduce max with over-syncing", "[warp][reduce][batched]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::maximum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::SingleOut,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::Reduce min with over-syncing",
|
||||
"[warp][reduce][batched]",
|
||||
builtin_type_list,
|
||||
unequal_nm_single_out_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::minimum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::SingleOut,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::Sum", "[warp][reduce][batched][convenience]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::SingleOut,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped N=M sum", "[warp][reduce][batched]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped N!=M sum", "[warp][reduce][batched]", builtin_type_list, unequal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped max with over-syncing",
|
||||
"[warp][reduce][batched]",
|
||||
builtin_type_list,
|
||||
equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::maximum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped min with over-syncing",
|
||||
"[warp][reduce][batched]",
|
||||
builtin_type_list,
|
||||
unequal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::minimum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::SumToStriped", "[warp][reduce][batched][convenience]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped with conditional participation N=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_equal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped with conditional participation N!=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_unequal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped with conditional participation and over-syncing N=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_equal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped with conditional participation and over-syncing N!=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_unequal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToStriped,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked N=M sum", "[warp][reduce][batched]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked N!=M sum", "[warp][reduce][batched]", builtin_type_list, unequal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked, num_batches, logical_warp_num_threads, value_t, op_t>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked max with over-syncing",
|
||||
"[warp][reduce][batched]",
|
||||
builtin_type_list,
|
||||
equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::maximum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked min with over-syncing",
|
||||
"[warp][reduce][batched]",
|
||||
builtin_type_list,
|
||||
unequal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::minimum<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = false;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::SumToBlocked", "[warp][reduce][batched][convenience]", builtin_type_list, equal_nm_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked with conditional participation N=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_equal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked with conditional participation N!=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_unequal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked with conditional participation and over-syncing N=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_equal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked with conditional participation and over-syncing N!=M sum",
|
||||
"[warp][reduce][batched][lane_mask]",
|
||||
builtin_type_list,
|
||||
sub_warp_unequal_configs)
|
||||
{
|
||||
using value_t = c2h::get<0, TestType>;
|
||||
using op_t = cuda::std::plus<>;
|
||||
constexpr int num_batches = c2h::get<1, TestType>::x;
|
||||
constexpr int logical_warp_num_threads = c2h::get<1, TestType>::y;
|
||||
constexpr bool convenience_overload = false;
|
||||
constexpr bool cond_participation = true;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
test_warp_reduce_batched<WarpReduceBatchedMode::ToBlocked,
|
||||
num_batches,
|
||||
logical_warp_num_threads,
|
||||
value_t,
|
||||
op_t,
|
||||
convenience_overload,
|
||||
cond_participation,
|
||||
sync_physical_warp>();
|
||||
}
|
||||
@@ -0,0 +1,354 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
|
||||
|
||||
#include <cub/warp/warp_reduce_batched.cuh>
|
||||
|
||||
#include <thrust/detail/raw_pointer_cast.h>
|
||||
#include <thrust/device_vector.h>
|
||||
#include <thrust/host_vector.h>
|
||||
|
||||
#include <cuda/__functional/maximum.h>
|
||||
#include <cuda/std/array>
|
||||
#include <cuda/std/span>
|
||||
|
||||
#include <c2h/catch2_test_helper.h>
|
||||
|
||||
__global__ __launch_bounds__(64) void WarpReduceBatchedOverviewKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-overview
|
||||
constexpr int num_batches = 3;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches>;
|
||||
|
||||
// Assume 64 threads per block, so 64 / 32 = 2 logical warps
|
||||
// Each logical warp has its own TempStorage
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[2];
|
||||
|
||||
const int warp_id = static_cast<int>(threadIdx.x) / 32;
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
int thread_data[num_batches];
|
||||
thread_data[0] = tid - 1;
|
||||
thread_data[1] = tid;
|
||||
thread_data[2] = tid + 1;
|
||||
|
||||
int result = WarpReduceBatched{temp_storage[warp_id]}.Reduce(thread_data, cuda::maximum{});
|
||||
// results across threads: [30, 31, 32, ?, ?, ..., ?, 62, 63, 64, ?, ?, ..., ?]
|
||||
// example-end warp-reduce-batched-overview
|
||||
|
||||
const int lane_id = static_cast<int>(threadIdx.x) % 32;
|
||||
if (lane_id < num_batches)
|
||||
{
|
||||
out[warp_id * num_batches + lane_id] = result;
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched overview documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(6);
|
||||
|
||||
WarpReduceBatchedOverviewKernel<<<1, 64>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{30, 31, 32, 62, 63, 64};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ void WarpReduceBatchedReduceApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-reduce
|
||||
// Can't allow for physical warp synchronization since only the first logical warp participates due to the
|
||||
// conditional. The other threads (assuming there are more than 16 threads per block) can't exit early due to the
|
||||
// barrier.
|
||||
constexpr int num_batches = 3;
|
||||
constexpr int logical_warp_threads = 16;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads>;
|
||||
|
||||
// Only the first logical warp participates, so only a single TempStorage is needed
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage;
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
int result{};
|
||||
if (threadIdx.x < logical_warp_threads)
|
||||
{
|
||||
const cuda::std::array<int, num_batches> inputs{tid - 1, tid, tid + 1};
|
||||
result = WarpReduceBatched{temp_storage}.Reduce(inputs, cuda::maximum{});
|
||||
}
|
||||
// results across threads: [14, 15, 16, ?, ?, ..., ?, 0, 0, ..., 0]
|
||||
__syncthreads();
|
||||
// Can reuse TempStorage after the barrier.
|
||||
// example-end warp-reduce-batched-reduce
|
||||
_CCCL_ASSERT(tid < logical_warp_threads || result == 0, "");
|
||||
if (threadIdx.x < num_batches)
|
||||
{
|
||||
out[tid] = result;
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::Reduce documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(3);
|
||||
|
||||
WarpReduceBatchedReduceApiKernel<<<1, 64>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{14, 15, 16};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(8) void WarpReduceBatchedReduceToStripedApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-reduce-to-striped
|
||||
// Can't allow for physical warp synchronization since only every other logical warp participates.
|
||||
// The other threads (assuming there are more than 2 threads per block) can't exit early due to the
|
||||
// barrier.
|
||||
constexpr int num_batches = 3;
|
||||
constexpr int logical_warp_threads = 2;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads>;
|
||||
|
||||
// Assume 8 threads per block, so 8 / 2 = 4 logical warps
|
||||
// Only every other logical warp participates, so only 2 TempStorage are needed
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[2];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / logical_warp_threads;
|
||||
const bool is_participating = logical_warp_id % 2 == 0;
|
||||
const int participant_idx = logical_warp_id / 2;
|
||||
|
||||
constexpr int max_out_per_thread = cuda::ceil_div(num_batches, logical_warp_threads);
|
||||
cuda::std::array<int, max_out_per_thread> results{};
|
||||
if (is_participating)
|
||||
{
|
||||
const cuda::std::array<int, num_batches> inputs{tid - 1, tid, tid + 1};
|
||||
WarpReduceBatched{temp_storage[participant_idx]}.ReduceToStriped(inputs, results, cuda::maximum{});
|
||||
}
|
||||
// results across threads:
|
||||
// [[0, 2], [1, ?], [0, 0], [0, 0], [4, 6], [5, ?], [0, 0], [0, 0]]
|
||||
__syncthreads();
|
||||
// Can reuse TempStorage after the barrier.
|
||||
// example-end warp-reduce-batched-reduce-to-striped
|
||||
|
||||
_CCCL_ASSERT(is_participating || (results[0] == 0 && results[1] == 0), "");
|
||||
if (is_participating)
|
||||
{
|
||||
int const logical_lane_id = tid % logical_warp_threads;
|
||||
for (int i = 0; i < max_out_per_thread; ++i)
|
||||
{
|
||||
// Striped
|
||||
const int batch_idx = i * logical_warp_threads + logical_lane_id;
|
||||
if (batch_idx < num_batches)
|
||||
{
|
||||
out[participant_idx * num_batches + batch_idx] = results[i];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToStriped documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(6);
|
||||
|
||||
WarpReduceBatchedReduceToStripedApiKernel<<<1, 8>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{0, 1, 2, 4, 5, 6};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(8) void WarpReduceBatchedReduceToBlockedApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-reduce-to-blocked
|
||||
// Can't allow for physical warp synchronization since only every other logical warp participates.
|
||||
// The other threads (assuming there are more than 2 threads per block) can't exit early due to the
|
||||
// barrier.
|
||||
constexpr int num_batches = 3;
|
||||
constexpr int logical_warp_threads = 2;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads>;
|
||||
|
||||
// Assume 8 threads per block, so 8 / 2 = 4 logical warps
|
||||
// Only every other logical warp participates, so only 2 TempStorage are needed
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[2];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / logical_warp_threads;
|
||||
const bool is_participating = logical_warp_id % 2 == 0;
|
||||
const int participant_idx = logical_warp_id / 2;
|
||||
|
||||
constexpr int max_out_per_thread = cuda::ceil_div(num_batches, logical_warp_threads);
|
||||
cuda::std::array<int, max_out_per_thread> results{};
|
||||
if (is_participating)
|
||||
{
|
||||
const cuda::std::array<int, num_batches> inputs{tid - 1, tid, tid + 1};
|
||||
|
||||
WarpReduceBatched{temp_storage[participant_idx]}.ReduceToBlocked(inputs, results, cuda::maximum{});
|
||||
}
|
||||
// results across threads:
|
||||
// [[0, 1], [2, ?], [0, 0], [0, 0], [4, 5], [6, ?], [0, 0], [0, 0]]
|
||||
__syncthreads();
|
||||
// Can reuse TempStorage after the barrier.
|
||||
// example-end warp-reduce-batched-reduce-to-blocked
|
||||
|
||||
_CCCL_ASSERT(is_participating || (results[0] == 0 && results[1] == 0), "");
|
||||
if (is_participating)
|
||||
{
|
||||
int const logical_lane_id = tid % logical_warp_threads;
|
||||
for (int i = 0; i < max_out_per_thread; ++i)
|
||||
{
|
||||
// Blocked
|
||||
const int batch_idx = logical_lane_id * max_out_per_thread + i;
|
||||
if (batch_idx < num_batches)
|
||||
{
|
||||
out[participant_idx * num_batches + batch_idx] = results[i];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::ReduceToBlocked documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(6);
|
||||
|
||||
WarpReduceBatchedReduceToBlockedApiKernel<<<1, 8>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{0, 1, 2, 4, 5, 6};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(8) void WarpReduceBatchedSumApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-sum
|
||||
constexpr int num_batches = 3;
|
||||
constexpr int logical_warp_threads = 4;
|
||||
// We can enable physical warp synchronization since all non-exited lanes do participate in the primitive.
|
||||
constexpr bool sync_physical_warp = true;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads, sync_physical_warp>;
|
||||
|
||||
// Assume 8 threads per block, so 8 / 4 = 2 logical warps
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[2];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / logical_warp_threads;
|
||||
|
||||
cuda::std::array<int, num_batches> inputs{tid - 1, tid, tid + 1};
|
||||
int result = WarpReduceBatched{temp_storage[logical_warp_id]}.Sum(inputs);
|
||||
// results across threads:
|
||||
// [2, 6, 10, ?, 18, 22, 26, ?]
|
||||
// example-end warp-reduce-batched-sum
|
||||
const int logical_lane_id = tid % logical_warp_threads;
|
||||
if (logical_lane_id < num_batches)
|
||||
{
|
||||
out[logical_warp_id * num_batches + logical_lane_id] = result;
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::Sum documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(6);
|
||||
|
||||
WarpReduceBatchedSumApiKernel<<<1, 8>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{2, 6, 10, 18, 22, 26};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(8) void WarpReduceBatchedSumToStripedApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-sum-to-striped
|
||||
constexpr int num_batches = 5;
|
||||
constexpr int logical_warp_threads = 2;
|
||||
// We can enable physical warp synchronization since all non-exited lanes do participate in the primitive.
|
||||
constexpr bool sync_physical_warp = true;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads, sync_physical_warp>;
|
||||
|
||||
// Assume 8 threads per block, so 8 / 2 = 4 logical warps
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[4];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / logical_warp_threads;
|
||||
|
||||
cuda::std::array<int, num_batches> inputs{tid - 2, tid - 1, tid, tid + 1, tid + 2};
|
||||
constexpr int max_out_per_thread = cuda::ceil_div(num_batches, logical_warp_threads);
|
||||
// Use a static size span to alias the last 3 elements of inputs for results.
|
||||
cuda::std::span<int, max_out_per_thread> results{cuda::std::end(inputs) - max_out_per_thread, max_out_per_thread};
|
||||
WarpReduceBatched{temp_storage[logical_warp_id]}.SumToStriped(inputs, results);
|
||||
// results across threads:
|
||||
// [[-3, 1, 5], [-1, 3, ?], [1, 5, 9], [3, 7, ?], ..., [9, 13, 17], [11, 15, ?]]
|
||||
// example-end warp-reduce-batched-sum-to-striped
|
||||
|
||||
const int logical_lane_id = tid % logical_warp_threads;
|
||||
for (int i = 0; i < max_out_per_thread; ++i)
|
||||
{
|
||||
// Striped
|
||||
const int batch_idx = i * logical_warp_threads + logical_lane_id;
|
||||
if (batch_idx < num_batches)
|
||||
{
|
||||
out[logical_warp_id * num_batches + batch_idx] = results[i];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::SumToStriped documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(20);
|
||||
|
||||
WarpReduceBatchedSumToStripedApiKernel<<<1, 8>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{-3, -1, 1, 3, 5, 1, 3, 5, 7, 9, 5, 7, 9, 11, 13, 9, 11, 13, 15, 17};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
|
||||
__global__ __launch_bounds__(8) void WarpReduceBatchedSumToBlockedApiKernel(int* out)
|
||||
{
|
||||
// example-begin warp-reduce-batched-sum-to-blocked
|
||||
// We can enable physical warp synchronization since all non-exited lanes do participate in the primitive.
|
||||
constexpr int num_batches = 5;
|
||||
constexpr int logical_warp_threads = 2;
|
||||
constexpr bool sync_physical_warp = true;
|
||||
using WarpReduceBatched = cub::WarpReduceBatched<int, num_batches, logical_warp_threads, sync_physical_warp>;
|
||||
|
||||
// Assume 8 threads per block, so 8 / 2 = 4 logical warps
|
||||
__shared__ typename WarpReduceBatched::TempStorage temp_storage[4];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
const int logical_warp_id = tid / logical_warp_threads;
|
||||
|
||||
cuda::std::array<int, num_batches> inputs{tid - 2, tid - 1, tid, tid + 1, tid + 2};
|
||||
constexpr int max_out_per_thread = cuda::ceil_div(num_batches, logical_warp_threads);
|
||||
// Use a static size span to alias the last 3 elements of inputs for results.
|
||||
cuda::std::span<int, max_out_per_thread> results{cuda::std::end(inputs) - max_out_per_thread, max_out_per_thread};
|
||||
WarpReduceBatched{temp_storage[logical_warp_id]}.SumToBlocked(inputs, results);
|
||||
// results across threads:
|
||||
// [[-3, -1, 1], [3, 5, ?], [1, 3, 5], [7, 9, ?], ..., [9, 11, 13], [15, 17, ?]]
|
||||
// example-end warp-reduce-batched-sum-to-blocked
|
||||
|
||||
const int logical_lane_id = tid % logical_warp_threads;
|
||||
for (int i = 0; i < max_out_per_thread; ++i)
|
||||
{
|
||||
// Blocked
|
||||
const int batch_idx = logical_lane_id * max_out_per_thread + i;
|
||||
if (batch_idx < num_batches)
|
||||
{
|
||||
out[logical_warp_id * num_batches + batch_idx] = results[i];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
C2H_TEST("WarpReduceBatched::SumToBlocked documentation kernel", "[warp][reduce][batched]")
|
||||
{
|
||||
c2h::device_vector<int> d_out(20);
|
||||
|
||||
WarpReduceBatchedSumToBlockedApiKernel<<<1, 8>>>(thrust::raw_pointer_cast(d_out.data()));
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
|
||||
c2h::host_vector<int> expected{-3, -1, 1, 3, 5, 1, 3, 5, 7, 9, 5, 7, 9, 11, 13, 9, 11, 13, 15, 17};
|
||||
REQUIRE(expected == d_out);
|
||||
}
|
||||
350
cccl_upstream/cub/test/warp/catch2_test_warp_segmented_reduce.cu
Normal file
350
cccl_upstream/cub/test/warp/catch2_test_warp_segmented_reduce.cu
Normal file
@@ -0,0 +1,350 @@
|
||||
// SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3
|
||||
|
||||
#include <cub/util_macro.cuh>
|
||||
#include <cub/warp/warp_reduce.cuh>
|
||||
|
||||
#include <cuda/std/functional>
|
||||
#include <cuda/std/limits>
|
||||
#include <cuda/std/type_traits>
|
||||
|
||||
#include <c2h/catch2_test_helper.h>
|
||||
#include <c2h/custom_type.h>
|
||||
|
||||
template <int LOGICAL_WARP_THREADS, int TOTAL_WARPS, typename T, typename ActionT>
|
||||
__global__ void warp_reduce_kernel(T* in, T* out, ActionT action)
|
||||
{
|
||||
using warp_reduce_t = cub::WarpReduce<T, LOGICAL_WARP_THREADS>;
|
||||
using storage_t = typename warp_reduce_t::TempStorage;
|
||||
|
||||
__shared__ storage_t storage[TOTAL_WARPS];
|
||||
|
||||
const int tid = static_cast<int>(threadIdx.x);
|
||||
|
||||
// Get warp index
|
||||
int warp_id = tid / LOGICAL_WARP_THREADS;
|
||||
|
||||
// Load data
|
||||
T thread_data = in[tid];
|
||||
// Instantiate and run warp reduction
|
||||
warp_reduce_t warp_reduce(storage[warp_id]);
|
||||
auto result = action(tid, warp_reduce, thread_data);
|
||||
|
||||
// Write warp aggregate
|
||||
out[tid] = result;
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Delegate wrapper for WarpReduce::TailSegmentedSum
|
||||
*/
|
||||
template <typename T>
|
||||
struct warp_seg_sum_tail_t
|
||||
{
|
||||
uint8_t* d_flags;
|
||||
template <int LOGICAL_WARP_THREADS>
|
||||
__device__ T operator()(int linear_tid, cub::WarpReduce<T, LOGICAL_WARP_THREADS>& warp_reduce, T& thread_data) const
|
||||
{
|
||||
const bool has_agg = (linear_tid % LOGICAL_WARP_THREADS == 0) || ((linear_tid == 0) ? 0 : d_flags[linear_tid - 1]);
|
||||
auto result = warp_reduce.TailSegmentedSum(thread_data, d_flags[linear_tid]);
|
||||
return has_agg ? result : thread_data;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Delegate wrapper for WarpReduce::HeadSegmentedSum
|
||||
*/
|
||||
template <typename T>
|
||||
struct warp_seg_sum_head_t
|
||||
{
|
||||
uint8_t* d_flags;
|
||||
template <int LOGICAL_WARP_THREADS>
|
||||
__device__ T operator()(int linear_tid, cub::WarpReduce<T, LOGICAL_WARP_THREADS>& warp_reduce, T& thread_data) const
|
||||
{
|
||||
const bool has_agg = ((linear_tid % LOGICAL_WARP_THREADS == 0) || d_flags[linear_tid]);
|
||||
auto result = warp_reduce.HeadSegmentedSum(thread_data, d_flags[linear_tid]);
|
||||
return (has_agg) ? result : thread_data;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Delegate wrapper for WarpReduce::TailSegmentedReduce
|
||||
*/
|
||||
template <typename T, typename ReductionOpT>
|
||||
struct warp_seg_reduce_tail_t
|
||||
{
|
||||
uint8_t* d_flags;
|
||||
ReductionOpT reduction_op;
|
||||
template <int LOGICAL_WARP_THREADS>
|
||||
__device__ T operator()(int linear_tid, cub::WarpReduce<T, LOGICAL_WARP_THREADS>& warp_reduce, T& thread_data) const
|
||||
{
|
||||
const bool has_agg = (linear_tid % LOGICAL_WARP_THREADS == 0) || ((linear_tid == 0) ? 0 : d_flags[linear_tid - 1]);
|
||||
auto result = warp_reduce.TailSegmentedReduce(thread_data, d_flags[linear_tid], reduction_op);
|
||||
return has_agg ? result : thread_data;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Delegate wrapper for WarpReduce::HeadSegmentedReduce
|
||||
*/
|
||||
template <typename T, typename ReductionOpT>
|
||||
struct warp_seg_reduce_head_t
|
||||
{
|
||||
uint8_t* d_flags;
|
||||
ReductionOpT reduction_op;
|
||||
template <int LOGICAL_WARP_THREADS>
|
||||
__device__ T operator()(int linear_tid, cub::WarpReduce<T, LOGICAL_WARP_THREADS>& warp_reduce, T& thread_data) const
|
||||
{
|
||||
const bool has_agg = ((linear_tid % LOGICAL_WARP_THREADS == 0) || d_flags[linear_tid]);
|
||||
auto result = warp_reduce.HeadSegmentedReduce(thread_data, d_flags[linear_tid], reduction_op);
|
||||
return (has_agg) ? result : thread_data;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* @brief Dispatch helper function
|
||||
*/
|
||||
template <int LOGICAL_WARP_THREADS, int TOTAL_WARPS, typename T, typename ActionT>
|
||||
void warp_reduce(c2h::device_vector<T>& in, c2h::device_vector<T>& out, ActionT action)
|
||||
{
|
||||
warp_reduce_kernel<LOGICAL_WARP_THREADS, TOTAL_WARPS, T, ActionT><<<1, LOGICAL_WARP_THREADS * TOTAL_WARPS>>>(
|
||||
thrust::raw_pointer_cast(in.data()), thrust::raw_pointer_cast(out.data()), action);
|
||||
|
||||
REQUIRE(cudaSuccess == cudaPeekAtLastError());
|
||||
REQUIRE(cudaSuccess == cudaDeviceSynchronize());
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Compares the results returned from system under test against the expected results.
|
||||
*/
|
||||
template <typename T, cuda::std::enable_if_t<cuda::std::is_floating_point_v<T>, int> = 0>
|
||||
void verify_results(const c2h::host_vector<T>& expected_data, const c2h::device_vector<T>& test_results)
|
||||
{
|
||||
REQUIRE_APPROX_EQ(expected_data, test_results);
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Compares the results returned from system under test against the expected results.
|
||||
*/
|
||||
template <typename T, cuda::std::enable_if_t<!cuda::std::is_floating_point_v<T>, int> = 0>
|
||||
void verify_results(const c2h::host_vector<T>& expected_data, const c2h::device_vector<T>& test_results)
|
||||
{
|
||||
REQUIRE(expected_data == test_results);
|
||||
}
|
||||
|
||||
enum class reduce_mode
|
||||
{
|
||||
all,
|
||||
partial,
|
||||
head_flags,
|
||||
tail_flags,
|
||||
};
|
||||
|
||||
template <typename InputItT, typename FlagInputItT, typename ReductionOp, typename ResultOutItT>
|
||||
void compute_host_reference(
|
||||
reduce_mode mode,
|
||||
InputItT h_in,
|
||||
FlagInputItT h_flags,
|
||||
int warps,
|
||||
int logical_warp_threads,
|
||||
int valid_warp_threads,
|
||||
ReductionOp reduction_op,
|
||||
ResultOutItT h_data_out)
|
||||
{
|
||||
// Accumulate segments (lane 0 of each warp is implicitly a segment head)
|
||||
for (int warp = 0; warp < warps; ++warp)
|
||||
{
|
||||
int warp_offset = warp * logical_warp_threads;
|
||||
int item_offset = warp_offset + valid_warp_threads - 1;
|
||||
|
||||
// Last item in warp
|
||||
auto head_aggregate = h_in[item_offset];
|
||||
auto tail_aggregate = h_in[item_offset];
|
||||
|
||||
if (mode != reduce_mode::tail_flags && h_flags[item_offset])
|
||||
{
|
||||
h_data_out[item_offset] = head_aggregate;
|
||||
}
|
||||
item_offset--;
|
||||
|
||||
// Work backwards
|
||||
while (item_offset >= warp_offset)
|
||||
{
|
||||
if (h_flags[item_offset + 1]) // NOLINT(bugprone-misplaced-widening-cast)
|
||||
{
|
||||
head_aggregate = h_in[item_offset];
|
||||
}
|
||||
else
|
||||
{
|
||||
head_aggregate = reduction_op(head_aggregate, h_in[item_offset]);
|
||||
}
|
||||
|
||||
if (h_flags[item_offset])
|
||||
{
|
||||
if (mode == reduce_mode::head_flags)
|
||||
{
|
||||
h_data_out[item_offset] = head_aggregate;
|
||||
}
|
||||
else if (mode == reduce_mode::tail_flags)
|
||||
{
|
||||
h_data_out[item_offset + 1] = tail_aggregate; // NOLINT(bugprone-misplaced-widening-cast)
|
||||
tail_aggregate = h_in[item_offset];
|
||||
}
|
||||
}
|
||||
else
|
||||
{
|
||||
tail_aggregate = reduction_op(tail_aggregate, h_in[item_offset]);
|
||||
}
|
||||
|
||||
item_offset--;
|
||||
}
|
||||
|
||||
// Record last segment aggregate
|
||||
if (mode == reduce_mode::tail_flags)
|
||||
{
|
||||
h_data_out[warp_offset] = tail_aggregate;
|
||||
}
|
||||
else
|
||||
{
|
||||
h_data_out[warp_offset] = head_aggregate;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// List of types to test
|
||||
using custom_t =
|
||||
c2h::custom_type_t<c2h::accumulateable_t, c2h::equal_comparable_t, c2h::lexicographical_less_comparable_t>;
|
||||
using full_type_list =
|
||||
c2h::type_list<std::uint8_t,
|
||||
std::uint16_t,
|
||||
std::int32_t,
|
||||
std::int64_t,
|
||||
custom_t,
|
||||
#if _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4_16a,
|
||||
#else // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
ulonglong4,
|
||||
#endif // _CCCL_CTK_AT_LEAST(13, 0)
|
||||
uchar3,
|
||||
short2>;
|
||||
|
||||
using builtin_type_list = c2h::type_list<std::uint8_t, std::uint16_t, std::int32_t, std::int64_t>;
|
||||
|
||||
// Logical warp sizes to test
|
||||
using logical_warp_threads = c2h::enum_type_list<int, 32, 16, 9, 7, 1>;
|
||||
|
||||
using segmented_modes = c2h::enum_type_list<reduce_mode, reduce_mode::head_flags, reduce_mode::tail_flags>;
|
||||
|
||||
template <int logical_warp_threads>
|
||||
struct total_warps_t
|
||||
{
|
||||
private:
|
||||
static constexpr int max_warps = 2;
|
||||
static constexpr bool is_arch_warp = (logical_warp_threads == cub::detail::warp_threads);
|
||||
static constexpr bool is_pow_of_two = ((logical_warp_threads & (logical_warp_threads - 1)) == 0);
|
||||
static constexpr int total_warps = (is_arch_warp || is_pow_of_two) ? max_warps : 1;
|
||||
|
||||
public:
|
||||
static constexpr int value()
|
||||
{
|
||||
return total_warps;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename TestType>
|
||||
struct params_t
|
||||
{
|
||||
using type = typename c2h::get<0, TestType>;
|
||||
|
||||
static constexpr int logical_warp_threads = c2h::get<1, TestType>::value;
|
||||
static constexpr int total_warps = total_warps_t<logical_warp_threads>::value();
|
||||
static constexpr int tile_size = total_warps * logical_warp_threads;
|
||||
};
|
||||
|
||||
C2H_TEST("Warp segmented sum works", "[reduce][warp]", full_type_list, logical_warp_threads, segmented_modes)
|
||||
{
|
||||
using params = params_t<TestType>;
|
||||
using type = typename params::type;
|
||||
|
||||
constexpr auto segmented_mod = c2h::get<2, TestType>::value;
|
||||
static_assert(segmented_mod == reduce_mode::tail_flags || segmented_mod == reduce_mode::head_flags,
|
||||
"Segmented tests must either be head or tail flags");
|
||||
using warp_seg_sum_t =
|
||||
cuda::std::_If<(segmented_mod == reduce_mode::tail_flags), warp_seg_sum_tail_t<type>, warp_seg_sum_head_t<type>>;
|
||||
|
||||
// Prepare test data
|
||||
c2h::device_vector<type> d_in(params::tile_size);
|
||||
c2h::device_vector<uint8_t> d_flags(params::tile_size);
|
||||
c2h::device_vector<type> d_out(params::tile_size);
|
||||
constexpr auto valid_items = params::logical_warp_threads;
|
||||
constexpr uint8_t min = 0;
|
||||
constexpr uint8_t max = 2;
|
||||
c2h::gen(C2H_SEED(5), d_in);
|
||||
c2h::gen(C2H_SEED(5), d_flags, min, max);
|
||||
|
||||
// Run test
|
||||
warp_reduce<params::logical_warp_threads, params::total_warps>(
|
||||
d_in, d_out, warp_seg_sum_t{thrust::raw_pointer_cast(d_flags.data())});
|
||||
|
||||
// Prepare verification data
|
||||
c2h::host_vector<type> h_in = d_in;
|
||||
c2h::host_vector<uint8_t> h_flags = d_flags;
|
||||
c2h::host_vector<type> h_out = h_in;
|
||||
compute_host_reference(
|
||||
segmented_mod,
|
||||
h_in,
|
||||
h_flags,
|
||||
params::total_warps,
|
||||
params::logical_warp_threads,
|
||||
valid_items,
|
||||
cuda::std::plus<type>{},
|
||||
h_out.begin());
|
||||
|
||||
// Verify results
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
|
||||
C2H_TEST("Warp segmented reduction works", "[reduce][warp]", builtin_type_list, logical_warp_threads, segmented_modes)
|
||||
{
|
||||
using params = params_t<TestType>;
|
||||
using type = typename params::type;
|
||||
using red_op_t = cuda::minimum<>;
|
||||
|
||||
constexpr auto segmented_mod = c2h::get<2, TestType>::value;
|
||||
static_assert(segmented_mod == reduce_mode::tail_flags || segmented_mod == reduce_mode::head_flags,
|
||||
"Segmented tests must either be head or tail flags");
|
||||
using warp_seg_reduction_t =
|
||||
cuda::std::_If<(segmented_mod == reduce_mode::tail_flags),
|
||||
warp_seg_reduce_tail_t<type, red_op_t>,
|
||||
warp_seg_reduce_head_t<type, red_op_t>>;
|
||||
|
||||
// Prepare test data
|
||||
c2h::device_vector<type> d_in(params::tile_size);
|
||||
c2h::device_vector<uint8_t> d_flags(params::tile_size);
|
||||
c2h::device_vector<type> d_out(params::tile_size);
|
||||
constexpr auto valid_items = params::logical_warp_threads;
|
||||
constexpr uint8_t min = 0;
|
||||
constexpr uint8_t max = 2;
|
||||
c2h::gen(C2H_SEED(5), d_in);
|
||||
c2h::gen(C2H_SEED(5), d_flags, min, max);
|
||||
|
||||
// Run test
|
||||
warp_reduce<params::logical_warp_threads, params::total_warps>(
|
||||
d_in, d_out, warp_seg_reduction_t{thrust::raw_pointer_cast(d_flags.data()), red_op_t{}});
|
||||
|
||||
// Prepare verification data
|
||||
c2h::host_vector<type> h_in = d_in;
|
||||
c2h::host_vector<uint8_t> h_flags = d_flags;
|
||||
c2h::host_vector<type> h_out = h_in;
|
||||
compute_host_reference(
|
||||
segmented_mod,
|
||||
h_in,
|
||||
h_flags,
|
||||
params::total_warps,
|
||||
params::logical_warp_threads,
|
||||
valid_items,
|
||||
red_op_t{},
|
||||
h_out.begin());
|
||||
|
||||
// Verify results
|
||||
verify_results(h_out, d_out);
|
||||
}
|
||||
Reference in New Issue
Block a user