10#ifndef EIGEN_FFT_IMKLFFT_IMPL_H
11#define EIGEN_FFT_IMKLFFT_IMPL_H
16#include "./InternalHeaderCheck.h"
25#define RUN_OR_ASSERT(EXPR, ERROR_MSG) \
27 MKL_LONG status = (EXPR); \
28 eigen_assert(status == DFTI_NO_ERROR && (ERROR_MSG)); \
31inline MKL_Complex16* complex_cast(
const std::complex<double>* p) {
32 return const_cast<MKL_Complex16*
>(
reinterpret_cast<const MKL_Complex16*
>(p));
35inline MKL_Complex8* complex_cast(
const std::complex<float>* p) {
36 return const_cast<MKL_Complex8*
>(
reinterpret_cast<const MKL_Complex8*
>(p));
48inline void configure_descriptor(std::shared_ptr<DFTI_DESCRIPTOR>& handl,
enum DFTI_CONFIG_VALUE precision,
49 enum DFTI_CONFIG_VALUE forward_domain, MKL_LONG dimension, MKL_LONG* sizes) {
50 eigen_assert(dimension == 1 || dimension == 2 &&
"Transformation dimension must be less than 3.");
52 DFTI_DESCRIPTOR_HANDLE res =
nullptr;
54 RUN_OR_ASSERT(DftiCreateDescriptor(&res, precision, forward_domain, dimension, *sizes),
55 "DftiCreateDescriptor failed.")
56 handl.reset(res, [](DFTI_DESCRIPTOR_HANDLE handle) { DftiFreeDescriptor(&handle); });
57 if (forward_domain == DFTI_REAL) {
59 RUN_OR_ASSERT(DftiSetValue(handl.get(), DFTI_CONJUGATE_EVEN_STORAGE, DFTI_COMPLEX_COMPLEX),
60 "DftiSetValue failed.")
63 RUN_OR_ASSERT(DftiCreateDescriptor(&res, precision, DFTI_COMPLEX, dimension, sizes),
"DftiCreateDescriptor failed.")
64 handl.reset(res, [](DFTI_DESCRIPTOR_HANDLE handle) { DftiFreeDescriptor(&handle); });
67 RUN_OR_ASSERT(DftiSetValue(handl.get(), DFTI_PLACEMENT, DFTI_NOT_INPLACE),
"DftiSetValue failed.")
68 RUN_OR_ASSERT(DftiCommitDescriptor(handl.get()),
"DftiCommitDescriptor failed.")
76 typedef float scalar_type;
77 typedef MKL_Complex8 complex_type;
79 std::shared_ptr<DFTI_DESCRIPTOR> m_plan;
83 enum DFTI_CONFIG_VALUE precision = DFTI_SINGLE;
85 inline void forward(complex_type* dst, complex_type* src, MKL_LONG nfft) {
87 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
89 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
92 inline void inverse(complex_type* dst, complex_type* src, MKL_LONG nfft) {
94 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
96 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
99 inline void forward(complex_type* dst, scalar_type* src, MKL_LONG nfft) {
101 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
103 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
106 inline void inverse(scalar_type* dst, complex_type* src, MKL_LONG nfft) {
108 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
110 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
113 inline void forward2(complex_type* dst, complex_type* src,
int n0,
int n1) {
115 MKL_LONG sizes[2] = {n0, n1};
116 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
118 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
121 inline void inverse2(complex_type* dst, complex_type* src,
int n0,
int n1) {
123 MKL_LONG sizes[2] = {n0, n1};
124 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
126 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
132 typedef double scalar_type;
133 typedef MKL_Complex16 complex_type;
135 std::shared_ptr<DFTI_DESCRIPTOR> m_plan;
139 enum DFTI_CONFIG_VALUE precision = DFTI_DOUBLE;
141 inline void forward(complex_type* dst, complex_type* src, MKL_LONG nfft) {
143 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
145 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
148 inline void inverse(complex_type* dst, complex_type* src, MKL_LONG nfft) {
150 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
152 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
155 inline void forward(complex_type* dst, scalar_type* src, MKL_LONG nfft) {
157 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
159 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
162 inline void inverse(scalar_type* dst, complex_type* src, MKL_LONG nfft) {
164 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
166 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
169 inline void forward2(complex_type* dst, complex_type* src,
int n0,
int n1) {
171 MKL_LONG sizes[2] = {n0, n1};
172 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
174 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst),
"DftiComputeForward failed.")
177 inline void inverse2(complex_type* dst, complex_type* src,
int n0,
int n1) {
179 MKL_LONG sizes[2] = {n0, n1};
180 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
182 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst),
"DftiComputeBackward failed.")
186template <
typename Scalar_>
188 typedef Scalar_ Scalar;
189 typedef std::complex<Scalar> Complex;
191 inline void clear() { m_plans.clear(); }
194 inline void fwd(Complex* dst,
const Complex* src,
int nfft) {
195 get_plan(nfft,
false, dst, src).forward(complex_cast(dst), complex_cast(src), nfft);
199 inline void fwd(Complex* dst,
const Scalar* src,
int nfft) {
200 get_plan(nfft,
true, dst, src).forward(complex_cast(dst),
const_cast<Scalar*
>(src), nfft);
204 inline void fwd2(Complex* dst,
const Complex* src,
int n0,
int n1) {
205 get_plan(n0, n1,
false, dst, src).forward2(complex_cast(dst), complex_cast(src), n0, n1);
209 inline void inv(Complex* dst,
const Complex* src,
int nfft) {
210 get_plan(nfft,
false, dst, src).inverse(complex_cast(dst), complex_cast(src), nfft);
214 inline void inv(Scalar* dst,
const Complex* src,
int nfft) {
215 get_plan(nfft,
true, dst, src).inverse(
const_cast<Scalar*
>(dst), complex_cast(src), nfft);
219 inline void inv2(Complex* dst,
const Complex* src,
int n0,
int n1) {
220 get_plan(n0, n1,
false, dst, src).inverse2(complex_cast(dst), complex_cast(src), n0, n1);
224 std::map<int64_t, plan<Scalar>> m_plans;
232 static int64_t plan_flags(
bool real_io,
void* dst,
const void* src) {
233 int inplace = (dst == src) ? 1 : 0;
234 int aligned = ((
reinterpret_cast<size_t>(src) & 15) | (
reinterpret_cast<size_t>(dst) & 15)) == 0 ? 1 : 0;
235 return (
int(real_io) << 2) | (inplace << 1) | aligned;
238 inline plan<Scalar>& get_plan(
int nfft,
bool real_io,
void* dst,
const void* src) {
239 int64_t key = ((nfft << 3) | plan_flags(real_io, dst, src)) << 1;
243 inline plan<Scalar>& get_plan(
int n0,
int n1,
bool real_io,
void* dst,
const void* src) {
244 int64_t key = (((((int64_t)n0) << 31) | (n1 << 3) | plan_flags(real_io, dst, src)) << 1) + 1;
Namespace containing all symbols from the Eigen library.