diff --git a/src/dft/backends/rocfft/commit.cpp b/src/dft/backends/rocfft/commit.cpp index 47f12f336..3740ebd85 100644 --- a/src/dft/backends/rocfft/commit.cpp +++ b/src/dft/backends/rocfft/commit.cpp @@ -181,6 +181,12 @@ class rocfft_commit final : public dft::detail::commit_impl { const std::size_t dimensions = config_values.dimensions.size(); constexpr std::size_t max_supported_dims = 3; + if (dimensions > max_supported_dims) { + throw oneapi::math::unimplemented( + "DFT", __FUNCTION__, + "rocfft only supports up to " + std::to_string(max_supported_dims) + + " dimensions, but " + std::to_string(dimensions) + " were given."); + } std::array lengths; // rocfft does dimensions in the reverse order to oneMath std::copy(config_values.dimensions.crbegin(), config_values.dimensions.crend(), @@ -236,8 +242,8 @@ class rocfft_commit final : public dft::detail::commit_impl { auto func = __FUNCTION__; auto check_strides = [&](const auto& strides) { - for (int i = 1; i <= dimensions; i++) { - for (int j = 1; j <= dimensions; j++) { + for (std::size_t i = 1; i <= dimensions; i++) { + for (std::size_t j = 1; j <= dimensions; j++) { std::int64_t cplx_dim = config_values.dimensions[j - 1]; std::int64_t real_dim = (dom == dft::domain::REAL && j == dimensions) ? (cplx_dim / 2 + 1) @@ -300,13 +306,15 @@ class rocfft_commit final : public dft::detail::commit_impl { std::unique_ptr description_destroyer_bwd(plan_desc_bwd, description_destroy); - std::array stride_a_indices{ 0, 1, 2 }; - std::sort(&stride_a_indices[0], &stride_a_indices[dimensions], + // Note: index with iterators rather than &array[dimensions], which is an + // out-of-range subscript when dimensions == max_supported_dims. + std::array stride_a_indices{ 0, 1, 2 }; + std::sort(stride_a_indices.begin(), stride_a_indices.begin() + dimensions, [&](std::size_t a, std::size_t b) { return stride_vecs.vec_a[a] < stride_vecs.vec_a[b]; }); - std::array stride_b_indices{ 0, 1, 2 }; - std::sort(&stride_b_indices[0], &stride_b_indices[dimensions], + std::array stride_b_indices{ 0, 1, 2 }; + std::sort(stride_b_indices.begin(), stride_b_indices.begin() + dimensions, [&](std::size_t a, std::size_t b) { return stride_vecs.vec_b[a] < stride_vecs.vec_b[b]; }); @@ -321,7 +329,7 @@ class rocfft_commit final : public dft::detail::commit_impl { auto are_strides_smaller_than_lengths = [=](auto& svec, auto& sindices, auto& domain_lengths) { return dimensions == 1 || - (domain_lengths[sindices[0]] <= svec[sindices[1]] && + (svec[sindices[0]] * domain_lengths[sindices[0]] <= svec[sindices[1]] && (dimensions == 2 || svec[sindices[1]] * domain_lengths[sindices[1]] <= svec[sindices[2]])); }; @@ -335,16 +343,19 @@ class rocfft_commit final : public dft::detail::commit_impl { const bool vec_b_valid_as_bwd_domain = are_strides_smaller_than_lengths(stride_vecs.vec_b, stride_b_indices, lengths_cplx); - // Test if the stride vector being used as the fwd/bwd domain for each direction has valid strides for that use. - bool valid_forward = (stride_vecs.fwd_in == stride_vecs.vec_a && - vec_a_valid_as_fwd_domain && vec_b_valid_as_bwd_domain) || - (vec_b_valid_as_fwd_domain && vec_a_valid_as_bwd_domain); - bool valid_backward = (stride_vecs.bwd_in == stride_vecs.vec_a && - vec_a_valid_as_bwd_domain && vec_b_valid_as_fwd_domain) || - (vec_b_valid_as_bwd_domain && vec_a_valid_as_fwd_domain); + // Test if the stride vector being used as the fwd/bwd domain for each direction has + // valid strides for that use. The forward direction reads forward-domain data through + // vec_a (fwd_in) and writes backward-domain data through vec_b (fwd_out). + bool valid_forward = vec_a_valid_as_fwd_domain && vec_b_valid_as_bwd_domain; + // With FWD/BWD_STRIDES each vector describes one domain, so the backward direction has + // the same requirements. With INPUT/OUTPUT_STRIDES the domains swap: vec_a (bwd_in) + // describes backward-domain data and vec_b (bwd_out) forward-domain data. + bool valid_backward = stride_api_choice == dft::detail::stride_api::FB_STRIDES + ? valid_forward + : (vec_a_valid_as_bwd_domain && vec_b_valid_as_fwd_domain); if (!valid_forward && !valid_backward) { - throw math::exception("dft/backends/cufft", __FUNCTION__, "Invalid strides."); + throw math::exception("dft/backends/rocfft", __FUNCTION__, "Invalid strides."); } if (valid_forward) { diff --git a/tests/unit_tests/dft/include/test_common.hpp b/tests/unit_tests/dft/include/test_common.hpp index 5b1647e94..46b97f24f 100644 --- a/tests/unit_tests/dft/include/test_common.hpp +++ b/tests/unit_tests/dft/include/test_common.hpp @@ -284,7 +284,9 @@ bool check_equal_strided(const vec1& v, const vec2& v_ref, std::vector else { strides_arr = get_default_strides(sizes); } - strides = { &strides_arr[0], &strides_arr[sizes.size() + 1] }; + // Use data() rather than &strides_arr[sizes.size() + 1], which is an out-of-range + // subscript for 3 dimensions. + strides = { strides_arr.data(), strides_arr.data() + sizes.size() + 1 }; } using T = std::decay_t; std::int64_t size0 = sizes[0];