23#ifndef EIGEN_GPU_FFT_H
24#define EIGEN_GPU_FFT_H
27#include "./InternalHeaderCheck.h"
29#include "./CuFftSupport.h"
30#include "./CuBlasSupport.h"
31#include "./GpuContext.h"
37static constexpr std::size_t kDefaultCufftPlanCacheCapacity = 16;
39template <
typename Scalar_>
42 using Scalar = Scalar_;
43 using Complex = std::complex<Scalar>;
44 using ComplexVector = Matrix<Complex, Dynamic, 1>;
45 using RealVector = Matrix<Scalar, Dynamic, 1>;
46 using ComplexMatrix = Matrix<Complex, Dynamic, Dynamic, ColMajor>;
56 explicit FFT(std::size_t plan_cache_capacity = kDefaultCufftPlanCacheCapacity)
57 : ctx_(&Context::threadLocal()), plans_(plan_cache_capacity > 0 ? plan_cache_capacity : 1) {}
65 explicit FFT(Context& ctx, std::size_t plan_cache_capacity = kDefaultCufftPlanCacheCapacity)
66 : ctx_(&ctx), plans_(plan_cache_capacity > 0 ? plan_cache_capacity : 1) {}
72 std::size_t plan_cache_capacity()
const {
return plans_.capacity(); }
75 std::size_t plan_cache_size()
const {
return plans_.size(); }
77 FFT(
const FFT&) =
delete;
78 FFT& operator=(
const FFT&) =
delete;
81 template <typename Derived, std::enable_if_t<NumTraits<typename Derived::Scalar>::IsComplex>* =
nullptr>
82 ComplexVector fwd(
const MatrixBase<Derived>& x) {
83 const ComplexVector input(x.derived());
84 const int n =
static_cast<int>(input.size());
85 if (n == 0)
return ComplexVector(0);
87 ensure_buffers(n *
sizeof(Complex), n *
sizeof(Complex));
88 EIGEN_CUDA_RUNTIME_CHECK(
89 cudaMemcpyAsync(d_in_.get(), input.data(), n *
sizeof(Complex), cudaMemcpyHostToDevice, ctx_->stream()));
91 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
92 EIGEN_CUFFT_CHECK(internal::cufftExecC2C_dispatch(plan,
static_cast<Complex*
>(d_in_.get()),
93 static_cast<Complex*
>(d_out_.get()), CUFFT_FORWARD));
95 ComplexVector result(n);
96 EIGEN_CUDA_RUNTIME_CHECK(
97 cudaMemcpyAsync(result.data(), d_out_.get(), n *
sizeof(Complex), cudaMemcpyDeviceToHost, ctx_->stream()));
98 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
103 template <
typename Derived>
104 ComplexVector inv(
const MatrixBase<Derived>& X) {
105 static_assert(NumTraits<typename Derived::Scalar>::IsComplex,
"inv() requires complex input");
106 const ComplexVector input(X.derived());
107 const int n =
static_cast<int>(input.size());
108 if (n == 0)
return ComplexVector(0);
110 ensure_buffers(n *
sizeof(Complex), n *
sizeof(Complex));
111 EIGEN_CUDA_RUNTIME_CHECK(
112 cudaMemcpyAsync(d_in_.get(), input.data(), n *
sizeof(Complex), cudaMemcpyHostToDevice, ctx_->stream()));
114 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
115 EIGEN_CUFFT_CHECK(internal::cufftExecC2C_dispatch(plan,
static_cast<Complex*
>(d_in_.get()),
116 static_cast<Complex*
>(d_out_.get()), CUFFT_INVERSE));
120 internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n),
static_cast<Complex*
>(d_out_.get()), 1));
122 ComplexVector result(n);
123 EIGEN_CUDA_RUNTIME_CHECK(
124 cudaMemcpyAsync(result.data(), d_out_.get(), n *
sizeof(Complex), cudaMemcpyDeviceToHost, ctx_->stream()));
125 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
130 template <typename Derived, std::enable_if_t<!NumTraits<typename Derived::Scalar>::IsComplex>* =
nullptr>
131 ComplexVector fwd(
const MatrixBase<Derived>& x) {
132 const RealVector input(x.derived());
133 const int n =
static_cast<int>(input.size());
134 if (n == 0)
return ComplexVector(0);
136 const int n_complex = n / 2 + 1;
137 ensure_buffers(n *
sizeof(Scalar), n_complex *
sizeof(Complex));
138 EIGEN_CUDA_RUNTIME_CHECK(
139 cudaMemcpyAsync(d_in_.get(), input.data(), n *
sizeof(Scalar), cudaMemcpyHostToDevice, ctx_->stream()));
141 cufftHandle plan = get_plan_1d(n, internal::cufft_r2c_type<Scalar>::value);
143 internal::cufftExecR2C_dispatch(plan,
static_cast<Scalar*
>(d_in_.get()),
static_cast<Complex*
>(d_out_.get())));
145 ComplexVector result(n_complex);
146 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(result.data(), d_out_.get(), n_complex *
sizeof(Complex),
147 cudaMemcpyDeviceToHost, ctx_->stream()));
148 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
154 template <
typename Derived>
155 RealVector invReal(
const MatrixBase<Derived>& X, Index nfft) {
156 static_assert(NumTraits<typename Derived::Scalar>::IsComplex,
"invReal() requires complex input");
157 const ComplexVector input(X.derived());
158 const int n =
static_cast<int>(nfft);
159 const int n_complex = n / 2 + 1;
160 eigen_assert(input.size() == n_complex);
161 if (n == 0)
return RealVector(0);
163 ensure_buffers(n_complex *
sizeof(Complex), n *
sizeof(Scalar));
165 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_in_.get(), input.data(), n_complex *
sizeof(Complex),
166 cudaMemcpyHostToDevice, ctx_->stream()));
168 cufftHandle plan = get_plan_1d(n, internal::cufft_c2r_type<Scalar>::value);
170 internal::cufftExecC2R_dispatch(plan,
static_cast<Complex*
>(d_in_.get()),
static_cast<Scalar*
>(d_out_.get())));
174 internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n),
static_cast<Scalar*
>(d_out_.get()), 1));
176 RealVector result(n);
177 EIGEN_CUDA_RUNTIME_CHECK(
178 cudaMemcpyAsync(result.data(), d_out_.get(), n *
sizeof(Scalar), cudaMemcpyDeviceToHost, ctx_->stream()));
179 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
184 template <
typename Derived>
185 ComplexMatrix fwd2(
const MatrixBase<Derived>& A) {
186 static_assert(NumTraits<typename Derived::Scalar>::IsComplex,
"fwd2() requires complex input");
187 const ComplexMatrix input(A.derived());
188 const int rows =
static_cast<int>(input.rows());
189 const int cols =
static_cast<int>(input.cols());
190 if (rows == 0 || cols == 0)
return ComplexMatrix(rows, cols);
192 const size_t total =
static_cast<size_t>(rows) *
static_cast<size_t>(cols) *
sizeof(Complex);
193 ensure_buffers(total, total);
194 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_in_.get(), input.data(), total, cudaMemcpyHostToDevice, ctx_->stream()));
196 cufftHandle plan = get_plan_2d(rows, cols, internal::cufft_c2c_type<Scalar>::value);
197 EIGEN_CUFFT_CHECK(internal::cufftExecC2C_dispatch(plan,
static_cast<Complex*
>(d_in_.get()),
198 static_cast<Complex*
>(d_out_.get()), CUFFT_FORWARD));
200 ComplexMatrix result(rows, cols);
201 EIGEN_CUDA_RUNTIME_CHECK(
202 cudaMemcpyAsync(result.data(), d_out_.get(), total, cudaMemcpyDeviceToHost, ctx_->stream()));
203 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
208 template <
typename Derived>
209 ComplexMatrix inv2(
const MatrixBase<Derived>& A) {
210 static_assert(NumTraits<typename Derived::Scalar>::IsComplex,
"inv2() requires complex input");
211 const ComplexMatrix input(A.derived());
212 const int rows =
static_cast<int>(input.rows());
213 const int cols =
static_cast<int>(input.cols());
214 if (rows == 0 || cols == 0)
return ComplexMatrix(rows, cols);
216 const size_t total =
static_cast<size_t>(rows) *
static_cast<size_t>(cols) *
sizeof(Complex);
217 ensure_buffers(total, total);
218 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_in_.get(), input.data(), total, cudaMemcpyHostToDevice, ctx_->stream()));
220 cufftHandle plan = get_plan_2d(rows, cols, internal::cufft_c2c_type<Scalar>::value);
221 EIGEN_CUFFT_CHECK(internal::cufftExecC2C_dispatch(plan,
static_cast<Complex*
>(d_in_.get()),
222 static_cast<Complex*
>(d_out_.get()), CUFFT_INVERSE));
225 const int total_elems = rows * cols;
226 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx_->cublasHandle(), total_elems, Scalar(1) / Scalar(total_elems),
227 static_cast<Complex*
>(d_out_.get()), 1));
229 ComplexMatrix result(rows, cols);
230 EIGEN_CUDA_RUNTIME_CHECK(
231 cudaMemcpyAsync(result.data(), d_out_.get(), total, cudaMemcpyDeviceToHost, ctx_->stream()));
232 EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
242 void fwd(
const DeviceMatrix<Complex>& d_x, DeviceMatrix<Complex>& d_X) {
243 eigen_assert(d_x.cols() <= 1 &&
"device fwd(): expected a column vector");
244 const int n = internal::to_blas_int(d_x.rows());
245 prepare_out(d_x, d_X, n, 1);
247 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
249 internal::cufftExecC2C_dispatch(plan,
const_cast<Complex*
>(d_x.data()), d_X.data(), CUFFT_FORWARD));
250 d_X.recordReady(ctx_->stream());
254 void inv(
const DeviceMatrix<Complex>& d_X, DeviceMatrix<Complex>& d_x) {
255 eigen_assert(d_X.cols() <= 1 &&
"device inv(): expected a column vector");
256 const int n = internal::to_blas_int(d_X.rows());
257 prepare_out(d_X, d_x, n, 1);
259 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
261 internal::cufftExecC2C_dispatch(plan,
const_cast<Complex*
>(d_X.data()), d_x.data(), CUFFT_INVERSE));
262 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n), d_x.data(), 1));
263 d_x.recordReady(ctx_->stream());
267 void fwd(
const DeviceMatrix<Scalar>& d_x, DeviceMatrix<Complex>& d_X) {
268 eigen_assert(d_x.cols() <= 1 &&
"device fwd(): expected a column vector");
269 const int n = internal::to_blas_int(d_x.rows());
270 const int n_complex = n / 2 + 1;
271 prepare_out(d_x, d_X, n == 0 ? 0 : n_complex, 1);
273 cufftHandle plan = get_plan_1d(n, internal::cufft_r2c_type<Scalar>::value);
274 EIGEN_CUFFT_CHECK(internal::cufftExecR2C_dispatch(plan,
const_cast<Scalar*
>(d_x.data()), d_X.data()));
275 d_X.recordReady(ctx_->stream());
281 void invReal(
const DeviceMatrix<Complex>& d_X, DeviceMatrix<Scalar>& d_x, Index nfft) {
282 eigen_assert(d_X.cols() <= 1 &&
"device invReal(): expected a column vector");
283 const int n = internal::to_blas_int(nfft);
284 const int n_complex = n / 2 + 1;
285 eigen_assert(n == 0 || d_X.rows() == n_complex);
286 prepare_out(d_X, d_x, n, 1);
289 ensure_buffers(
static_cast<size_t>(n_complex) *
sizeof(Complex), 0);
290 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_in_.get(), d_X.data(), n_complex *
sizeof(Complex),
291 cudaMemcpyDeviceToDevice, ctx_->stream()));
292 cufftHandle plan = get_plan_1d(n, internal::cufft_c2r_type<Scalar>::value);
293 EIGEN_CUFFT_CHECK(internal::cufftExecC2R_dispatch(plan,
static_cast<Complex*
>(d_in_.get()), d_x.data()));
294 EIGEN_CUBLAS_CHECK(internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n), d_x.data(), 1));
295 d_x.recordReady(ctx_->stream());
299 void fwd2(
const DeviceMatrix<Complex>& d_A, DeviceMatrix<Complex>& d_B) {
300 const int rows = internal::to_blas_int(d_A.rows());
301 const int cols = internal::to_blas_int(d_A.cols());
302 prepare_out(d_A, d_B, rows, cols);
303 if (rows == 0 || cols == 0)
return;
304 cufftHandle plan = get_plan_2d(rows, cols, internal::cufft_c2c_type<Scalar>::value);
306 internal::cufftExecC2C_dispatch(plan,
const_cast<Complex*
>(d_A.data()), d_B.data(), CUFFT_FORWARD));
307 d_B.recordReady(ctx_->stream());
311 void inv2(
const DeviceMatrix<Complex>& d_A, DeviceMatrix<Complex>& d_B) {
312 const int rows = internal::to_blas_int(d_A.rows());
313 const int cols = internal::to_blas_int(d_A.cols());
314 prepare_out(d_A, d_B, rows, cols);
315 if (rows == 0 || cols == 0)
return;
316 cufftHandle plan = get_plan_2d(rows, cols, internal::cufft_c2c_type<Scalar>::value);
318 internal::cufftExecC2C_dispatch(plan,
const_cast<Complex*
>(d_A.data()), d_B.data(), CUFFT_INVERSE));
319 const int total_elems = internal::to_blas_int(
static_cast<int64_t
>(rows) *
static_cast<int64_t
>(cols));
321 internal::cublasXscal(ctx_->cublasHandle(), total_elems, Scalar(1) / Scalar(total_elems), d_B.data(), 1));
322 d_B.recordReady(ctx_->stream());
326 cudaStream_t stream()
const {
return ctx_->stream(); }
329 Context& context()
const {
return *ctx_; }
333 Eigen::internal::LruCache<int64_t, internal::CufftPlan> plans_;
334 internal::DeviceBuffer d_in_;
335 internal::DeviceBuffer d_out_;
336 size_t d_in_size_ = 0;
337 size_t d_out_size_ = 0;
341 template <
typename InScalar,
typename OutScalar>
342 void prepare_out(
const DeviceMatrix<InScalar>& in, DeviceMatrix<OutScalar>& out, Index out_rows, Index out_cols) {
344 (in.data() ==
nullptr ||
static_cast<const void*
>(in.data()) !=
static_cast<const void*
>(out.data())) &&
345 "device FFT: output must not alias input");
346 in.waitReady(ctx_->stream());
347 if (!out.empty()) out.waitReady(ctx_->stream());
348 out.resize(out_rows, out_cols);
356 void ensure_buffers(
size_t in_bytes,
size_t out_bytes) {
357 if (in_bytes > d_in_size_) {
358 if (d_in_) EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
359 d_in_ = internal::DeviceBuffer(in_bytes);
360 d_in_size_ = in_bytes;
362 if (out_bytes > d_out_size_) {
363 if (d_out_) EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(ctx_->stream()));
364 d_out_ = internal::DeviceBuffer(out_bytes);
365 d_out_size_ = out_bytes;
373 static constexpr int64_t kTypeMask = 0xF;
374 static constexpr int kCols2DBits = 30;
375 static constexpr int kRows2DBits = 29;
376 static int64_t plan_key_1d(
int n, cufftType type) {
return (int64_t(n) << 5) | (int64_t(type & kTypeMask) << 1) | 0; }
378 static int64_t plan_key_2d(
int rows,
int cols, cufftType type) {
379 eigen_assert(rows >= 0 && int64_t(rows) < (int64_t(1) << kRows2DBits) &&
380 "FFT plan rows exceed plan-key bit budget");
381 eigen_assert(cols >= 0 && int64_t(cols) < (int64_t(1) << kCols2DBits) &&
382 "FFT plan cols exceed plan-key bit budget");
383 return (int64_t(rows) << 35) | (int64_t(cols) << 5) | (int64_t(type & kTypeMask) << 1) | 1;
386 cufftHandle get_plan_1d(
int n, cufftType type) {
387 const int64_t key = plan_key_1d(n, type);
388 if (internal::CufftPlan* hit = plans_.find(key))
return hit->get();
391 EIGEN_CUFFT_CHECK(cufftPlan1d(&plan, n, type, 1));
392 EIGEN_CUFFT_CHECK(cufftSetStream(plan, ctx_->stream()));
393 return plans_.insert(key, internal::CufftPlan(plan))->get();
396 cufftHandle get_plan_2d(
int rows,
int cols, cufftType type) {
397 const int64_t key = plan_key_2d(rows, cols, type);
398 if (internal::CufftPlan* hit = plans_.find(key))
return hit->get();
404 EIGEN_CUFFT_CHECK(cufftPlan2d(&plan, cols, rows, type));
405 EIGEN_CUFFT_CHECK(cufftSetStream(plan, ctx_->stream()));
406 return plans_.insert(key, internal::CufftPlan(plan))->get();
Namespace containing all symbols from the Eigen library.