diff --git a/include/matx/transforms/cub.h b/include/matx/transforms/cub.h index ace85c0de..dca74a043 100644 --- a/include/matx/transforms/cub.h +++ b/include/matx/transforms/cub.h @@ -659,12 +659,30 @@ inline void ExecSort(OutputTensor &a_out, // type of reduction where there's not a single output, since any type of reduction can be generalized // to a segmented type if constexpr (OutputTensor::Rank() > 0) { +#if CUB_MAJOR_VERSION >= 3 && CUB_MINOR_VERSION >= 2 + [[maybe_unused]] cudaError_t err; + if (is_tensor_view_v && a.IsContiguous() && a_out.IsContiguous()) { + const int seg_size = static_cast(TotalSize(a) / TotalSize(out_base)); + err = cub::DeviceSegmentedReduce::Reduce(d_temp, temp_storage_bytes, in_base.Data(), out_base.Data(), static_cast(TotalSize(out_base)), seg_size, cparams_.reduce_op, + cparams_.init, stream); + } + else { + auto ft = [&](auto &&in, auto &&out, auto &&begin, auto &&end) { + return cub::DeviceSegmentedReduce::Reduce(d_temp, temp_storage_bytes, in, out, static_cast(TotalSize(out_base)), begin, end, cparams_.reduce_op, + cparams_.init, stream); + }; + err = ReduceInput(ft, out_base, in_base); + } + + MATX_ASSERT_STR_EXP(err, cudaSuccess, matxCudaError, "Error in cub::DeviceSegmentedReduce::Reduce"); +#else auto ft = [&](auto &&in, auto &&out, auto &&begin, auto &&end) { return cub::DeviceSegmentedReduce::Reduce(d_temp, temp_storage_bytes, in, out, static_cast(TotalSize(out_base)), begin, end, cparams_.reduce_op, cparams_.init, stream); }; [[maybe_unused]] auto rv = ReduceInput(ft, out_base, in_base); MATX_ASSERT_STR_EXP(rv, cudaSuccess, matxCudaError, "Error in cub::DeviceSegmentedReduce::Reduce"); +#endif } else { auto ft = [&](auto &&in, auto &&out, [[maybe_unused]] auto &&unused1, [[maybe_unused]] auto &&unused2) { @@ -710,11 +728,29 @@ inline void ExecSort(OutputTensor &a_out, // type of reduction where there's not a single output, since any type of reduction can be generalized // to a segmented type if constexpr (OutputTensor::Rank() > 0) { + // Check if fixed-size reductions are supported +#if CUB_MAJOR_VERSION >= 3 && CUB_MINOR_VERSION >= 2 + [[maybe_unused]] cudaError_t err; + if (is_tensor_view_v && a.IsContiguous() && a_out.IsContiguous()) { + const int seg_size = static_cast(TotalSize(a) / TotalSize(out_base)); + err = cub::DeviceSegmentedReduce::Sum(d_temp, temp_storage_bytes, in_base.Data(), out_base.Data(), static_cast(TotalSize(out_base)), seg_size, stream); + } + else { + auto ft = [&](auto &&in, auto &&out, auto &&begin, auto &&end) { + return cub::DeviceSegmentedReduce::Sum(d_temp, temp_storage_bytes, in, out, static_cast(TotalSize(out_base)), begin, end, stream); + }; + + err = ReduceInput(ft, out_base, in_base); + } + + MATX_ASSERT_STR_EXP(err, cudaSuccess, matxCudaError, "Error in cub::DeviceSegmentedReduce::Sum"); +#else auto ft = [&](auto &&in, auto &&out, auto &&begin, auto &&end) { - return cub::DeviceSegmentedReduce::Sum(d_temp, temp_storage_bytes, in, out, static_cast(TotalSize(out_base)), begin, end, stream); + return cub::DeviceSegmentedReduce::Sum(d_temp, temp_storage_bytes, in, out, static_cast(TotalSize(out_base)), begin, end, stream); }; [[maybe_unused]] auto rv = ReduceInput(ft, out_base, in_base); MATX_ASSERT_STR_EXP(rv, cudaSuccess, matxCudaError, "Error in cub::DeviceSegmentedReduce::Sum"); +#endif } else { auto ft = [&](auto &&in, auto &&out, [[maybe_unused]] auto &&unused1, [[maybe_unused]] auto &&unused2) {