Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 26 additions & 15 deletions src/dft/backends/rocfft/commit.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -181,6 +181,12 @@ class rocfft_commit final : public dft::detail::commit_impl<prec, dom> {
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<std::size_t, max_supported_dims> lengths;
// rocfft does dimensions in the reverse order to oneMath
std::copy(config_values.dimensions.crbegin(), config_values.dimensions.crend(),
Expand Down Expand Up @@ -236,8 +242,8 @@ class rocfft_commit final : public dft::detail::commit_impl<prec, dom> {

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)
Expand Down Expand Up @@ -300,13 +306,15 @@ class rocfft_commit final : public dft::detail::commit_impl<prec, dom> {
std::unique_ptr<rocfft_plan_description_t, decltype(description_destroy)>
description_destroyer_bwd(plan_desc_bwd, description_destroy);

std::array<std::size_t, 3> 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<std::size_t, max_supported_dims> 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<std::size_t, 3> stride_b_indices{ 0, 1, 2 };
std::sort(&stride_b_indices[0], &stride_b_indices[dimensions],
std::array<std::size_t, max_supported_dims> 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];
});
Expand All @@ -321,7 +329,7 @@ class rocfft_commit final : public dft::detail::commit_impl<prec, dom> {
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]]));
};
Expand All @@ -335,16 +343,19 @@ class rocfft_commit final : public dft::detail::commit_impl<prec, dom> {
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) {
Expand Down
4 changes: 3 additions & 1 deletion tests/unit_tests/dft/include/test_common.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -284,7 +284,9 @@ bool check_equal_strided(const vec1& v, const vec2& v_ref, std::vector<int64_t>
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<decltype(v[0])>;
std::int64_t size0 = sizes[0];
Expand Down
Loading