46 using descr_type = mkl_dft::descriptor<prec, dom>;
50 : descr_(dimensions), queue_ptr_{}
55 void commit(sycl::queue &q)
57 mkl_dft::precision fft_prec = get_precision();
58 if (fft_prec == mkl_dft::precision::DOUBLE &&
59 !q.get_device().has(sycl::aspect::fp64)) {
60 throw py::value_error(
"Descriptor is double precision but the "
61 "device does not support double precision.");
67 py::gil_scoped_release lock{};
71 queue_ptr_ = std::make_unique<sycl::queue>(q);
74 descr_type &get_descriptor() {
return descr_; }
76 const sycl::queue &get_queue()
const
82 throw std::runtime_error(
83 "Attempt to get queue when it is not yet set");
88 template <
typename valT = std::
int64_t>
92 descr_.get_value(mkl_dft::config_param::DIMENSION, &dim);
98 template <
typename valT = std::
int64_t>
99 const valT get_number_of_transforms()
101 valT transforms_count{};
103 descr_.get_value(mkl_dft::config_param::NUMBER_OF_TRANSFORMS,
105 return transforms_count;
108 template <
typename valT = std::
int64_t>
109 void set_number_of_transforms(
const valT &num)
111 descr_.set_value(mkl_dft::config_param::NUMBER_OF_TRANSFORMS, num);
115 template <
typename valT = std::vector<std::
int64_t>>
116 const valT get_fwd_strides()
118 const typename valT::value_type dim = get_dim();
120 valT fwd_strides(dim + 1);
121#if defined(USE_ONEMATH)
123 descr_.get_value(mkl_dft::config_param::FWD_STRIDES,
126 descr_.get_value(mkl_dft::config_param::FWD_STRIDES, &fwd_strides);
131 template <
typename valT = std::vector<std::
int64_t>>
132 void set_fwd_strides(
const valT &strides)
134 const typename valT::value_type dim = get_dim();
136 if (
static_cast<size_t>(dim + 1) != strides.size()) {
137 throw py::value_error(
138 "Strides length does not match descriptor's dimension");
140#if defined(USE_ONEMATH)
142 descr_.set_value(mkl_dft::config_param::FWD_STRIDES, strides.data());
144 descr_.set_value(mkl_dft::config_param::FWD_STRIDES, strides);
149 template <
typename valT = std::vector<std::
int64_t>>
150 const valT get_bwd_strides()
152 const typename valT::value_type dim = get_dim();
154 valT bwd_strides(dim + 1);
155#if defined(USE_ONEMATH)
157 descr_.get_value(mkl_dft::config_param::BWD_STRIDES,
160 descr_.get_value(mkl_dft::config_param::BWD_STRIDES, &bwd_strides);
165 template <
typename valT = std::vector<std::
int64_t>>
166 void set_bwd_strides(
const valT &strides)
168 const typename valT::value_type dim = get_dim();
170 if (
static_cast<size_t>(dim + 1) != strides.size()) {
171 throw py::value_error(
172 "Strides length does not match descriptor's dimension");
174#if defined(USE_ONEMATH)
176 descr_.set_value(mkl_dft::config_param::BWD_STRIDES, strides.data());
178 descr_.set_value(mkl_dft::config_param::BWD_STRIDES, strides);
183 template <
typename valT = std::
int64_t>
184 const valT get_fwd_distance()
188 descr_.get_value(mkl_dft::config_param::FWD_DISTANCE, &dist);
192 template <
typename valT = std::
int64_t>
193 void set_fwd_distance(
const valT &dist)
195 descr_.set_value(mkl_dft::config_param::FWD_DISTANCE, dist);
199 template <
typename valT = std::
int64_t>
200 const valT get_bwd_distance()
204 descr_.get_value(mkl_dft::config_param::BWD_DISTANCE, &dist);
208 template <
typename valT = std::
int64_t>
209 void set_bwd_distance(
const valT &dist)
211 descr_.set_value(mkl_dft::config_param::BWD_DISTANCE, dist);
217 mkl_dft::config_value placement;
218 descr_.get_value(mkl_dft::config_param::PLACEMENT, &placement);
219 return (placement == mkl_dft::config_value::INPLACE);
222 void set_in_place(
const bool &in_place_request)
224 descr_.set_value(mkl_dft::config_param::PLACEMENT,
226 ? mkl_dft::config_value::INPLACE
227 : mkl_dft::config_value::NOT_INPLACE);
231 mkl_dft::precision get_precision()
233 mkl_dft::precision fft_prec;
235 descr_.get_value(mkl_dft::config_param::PRECISION, &fft_prec);
242 mkl_dft::config_value committed;
243 descr_.get_value(mkl_dft::config_param::COMMIT_STATUS, &committed);
244 return (committed == mkl_dft::config_value::COMMITTED);
248 mkl_dft::descriptor<prec, dom> descr_;
249 std::unique_ptr<sycl::queue> queue_ptr_;