[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:
EngineX CI
2026-07-30 09:35:51 +00:00
parent b4d01f481e
commit 56fd68e7dd
8871 changed files with 1454674 additions and 0 deletions

View 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);
}

View 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>();
}

View File

@@ -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);
}

View 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);
}