Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
imklfft_impl.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// This Source Code Form is subject to the terms of the Mozilla
5// Public License v. 2.0. If a copy of the MPL was not distributed
6// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
7// SPDX-FileCopyrightText: The Eigen Authors
8// SPDX-License-Identifier: MPL-2.0
9
10#ifndef EIGEN_FFT_IMKLFFT_IMPL_H
11#define EIGEN_FFT_IMKLFFT_IMPL_H
12
13#include <mkl_dfti.h>
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18#include <complex>
19#include <memory>
20
21namespace Eigen {
22namespace internal {
23namespace imklfft {
24
25#define RUN_OR_ASSERT(EXPR, ERROR_MSG) \
26 { \
27 MKL_LONG status = (EXPR); \
28 eigen_assert(status == DFTI_NO_ERROR && (ERROR_MSG)); \
29 };
30
31inline MKL_Complex16* complex_cast(const std::complex<double>* p) {
32 return const_cast<MKL_Complex16*>(reinterpret_cast<const MKL_Complex16*>(p));
33}
34
35inline MKL_Complex8* complex_cast(const std::complex<float>* p) {
36 return const_cast<MKL_Complex8*>(reinterpret_cast<const MKL_Complex8*>(p));
37}
38
39/*
40 * Parameters:
41 * precision: enum, Precision of the transform: DFTI_SINGLE or DFTI_DOUBLE.
42 * forward_domain: enum, Forward domain of the transform: DFTI_COMPLEX or
43 * DFTI_REAL. dimension: MKL_LONG Dimension of the transform. sizes: MKL_LONG if
44 * dimension = 1. Length of the transform for a one-dimensional transform. sizes:
45 * Array of type MKL_LONG otherwise. Lengths of each dimension for a
46 * multi-dimensional transform.
47 */
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.");
51
52 DFTI_DESCRIPTOR_HANDLE res = nullptr;
53 if (dimension == 1) {
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) {
58 // Set CCE storage
59 RUN_OR_ASSERT(DftiSetValue(handl.get(), DFTI_CONJUGATE_EVEN_STORAGE, DFTI_COMPLEX_COMPLEX),
60 "DftiSetValue failed.")
61 }
62 } else {
63 RUN_OR_ASSERT(DftiCreateDescriptor(&res, precision, DFTI_COMPLEX, dimension, sizes), "DftiCreateDescriptor failed.")
64 handl.reset(res, [](DFTI_DESCRIPTOR_HANDLE handle) { DftiFreeDescriptor(&handle); });
65 }
66
67 RUN_OR_ASSERT(DftiSetValue(handl.get(), DFTI_PLACEMENT, DFTI_NOT_INPLACE), "DftiSetValue failed.")
68 RUN_OR_ASSERT(DftiCommitDescriptor(handl.get()), "DftiCommitDescriptor failed.")
69}
70
71template <typename T>
72struct plan {};
73
74template <>
75struct plan<float> {
76 typedef float scalar_type;
77 typedef MKL_Complex8 complex_type;
78
79 std::shared_ptr<DFTI_DESCRIPTOR> m_plan;
80
81 plan() = default;
82
83 enum DFTI_CONFIG_VALUE precision = DFTI_SINGLE;
84
85 inline void forward(complex_type* dst, complex_type* src, MKL_LONG nfft) {
86 if (m_plan == 0) {
87 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
88 }
89 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
90 }
91
92 inline void inverse(complex_type* dst, complex_type* src, MKL_LONG nfft) {
93 if (m_plan == 0) {
94 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
95 }
96 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
97 }
98
99 inline void forward(complex_type* dst, scalar_type* src, MKL_LONG nfft) {
100 if (m_plan == 0) {
101 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
102 }
103 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
104 }
105
106 inline void inverse(scalar_type* dst, complex_type* src, MKL_LONG nfft) {
107 if (m_plan == 0) {
108 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
109 }
110 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
111 }
112
113 inline void forward2(complex_type* dst, complex_type* src, int n0, int n1) {
114 if (m_plan == 0) {
115 MKL_LONG sizes[2] = {n0, n1};
116 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
117 }
118 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
119 }
120
121 inline void inverse2(complex_type* dst, complex_type* src, int n0, int n1) {
122 if (m_plan == 0) {
123 MKL_LONG sizes[2] = {n0, n1};
124 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
125 }
126 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
127 }
128};
129
130template <>
131struct plan<double> {
132 typedef double scalar_type;
133 typedef MKL_Complex16 complex_type;
134
135 std::shared_ptr<DFTI_DESCRIPTOR> m_plan;
136
137 plan() = default;
138
139 enum DFTI_CONFIG_VALUE precision = DFTI_DOUBLE;
140
141 inline void forward(complex_type* dst, complex_type* src, MKL_LONG nfft) {
142 if (m_plan == 0) {
143 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
144 }
145 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
146 }
147
148 inline void inverse(complex_type* dst, complex_type* src, MKL_LONG nfft) {
149 if (m_plan == 0) {
150 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 1, &nfft);
151 }
152 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
153 }
154
155 inline void forward(complex_type* dst, scalar_type* src, MKL_LONG nfft) {
156 if (m_plan == 0) {
157 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
158 }
159 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
160 }
161
162 inline void inverse(scalar_type* dst, complex_type* src, MKL_LONG nfft) {
163 if (m_plan == 0) {
164 configure_descriptor(m_plan, precision, DFTI_REAL, 1, &nfft);
165 }
166 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
167 }
168
169 inline void forward2(complex_type* dst, complex_type* src, int n0, int n1) {
170 if (m_plan == 0) {
171 MKL_LONG sizes[2] = {n0, n1};
172 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
173 }
174 RUN_OR_ASSERT(DftiComputeForward(m_plan.get(), src, dst), "DftiComputeForward failed.")
175 }
176
177 inline void inverse2(complex_type* dst, complex_type* src, int n0, int n1) {
178 if (m_plan == 0) {
179 MKL_LONG sizes[2] = {n0, n1};
180 configure_descriptor(m_plan, precision, DFTI_COMPLEX, 2, sizes);
181 }
182 RUN_OR_ASSERT(DftiComputeBackward(m_plan.get(), src, dst), "DftiComputeBackward failed.")
183 }
184};
185
186template <typename Scalar_>
187struct imklfft_impl {
188 typedef Scalar_ Scalar;
189 typedef std::complex<Scalar> Complex;
190
191 inline void clear() { m_plans.clear(); }
192
193 // complex-to-complex forward FFT
194 inline void fwd(Complex* dst, const Complex* src, int nfft) {
195 get_plan(nfft, /*real_io=*/false, dst, src).forward(complex_cast(dst), complex_cast(src), nfft);
196 }
197
198 // real-to-complex forward FFT
199 inline void fwd(Complex* dst, const Scalar* src, int nfft) {
200 get_plan(nfft, /*real_io=*/true, dst, src).forward(complex_cast(dst), const_cast<Scalar*>(src), nfft);
201 }
202
203 // 2-d complex-to-complex
204 inline void fwd2(Complex* dst, const Complex* src, int n0, int n1) {
205 get_plan(n0, n1, /*real_io=*/false, dst, src).forward2(complex_cast(dst), complex_cast(src), n0, n1);
206 }
207
208 // inverse complex-to-complex
209 inline void inv(Complex* dst, const Complex* src, int nfft) {
210 get_plan(nfft, /*real_io=*/false, dst, src).inverse(complex_cast(dst), complex_cast(src), nfft);
211 }
212
213 // half-complex to scalar
214 inline void inv(Scalar* dst, const Complex* src, int nfft) {
215 get_plan(nfft, /*real_io=*/true, dst, src).inverse(const_cast<Scalar*>(dst), complex_cast(src), nfft);
216 }
217
218 // 2-d complex-to-complex
219 inline void inv2(Complex* dst, const Complex* src, int n0, int n1) {
220 get_plan(n0, n1, /*real_io=*/false, dst, src).inverse2(complex_cast(dst), complex_cast(src), n0, n1);
221 }
222
223 private:
224 std::map<int64_t, plan<Scalar>> m_plans;
225
226 // Pack (real_io, inplace, aligned) into 3 contiguous low bits of the cache
227 // key. real_io distinguishes plans built with DFTI_REAL from those built
228 // with DFTI_COMPLEX so reusing the same FFT object across real-input and
229 // complex-input transforms doesn't return a cached plan of the wrong domain.
230 // MKL DFTI descriptors are bidirectional, so forward and inverse share
231 // slots intentionally.
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;
236 }
237
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;
240 return m_plans[key];
241 }
242
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;
245 return m_plans[key];
246 }
247};
248
249#undef RUN_OR_ASSERT
250
251} // namespace imklfft
252} // namespace internal
253} // namespace Eigen
254
255#endif // EIGEN_FFT_IMKLFFT_IMPL_H
Namespace containing all symbols from the Eigen library.