Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
CuFftSupport.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Rasmus Munk Larsen <rmlarsen@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11// cuFFT-specific support types.
12
13#ifndef EIGEN_GPU_CUFFT_SUPPORT_H
14#define EIGEN_GPU_CUFFT_SUPPORT_H
15
16// IWYU pragma: private
17#include "./InternalHeaderCheck.h"
18
19#include "./GpuSupport.h"
20#include <cufft.h>
21
22namespace Eigen {
23namespace gpu {
24namespace internal {
25
26#define EIGEN_CUFFT_CHECK(x) \
27 do { \
28 const cufftResult _r = (x); \
29 if (_r != CUFFT_SUCCESS) \
30 ::Eigen::gpu::internal::gpu_check_failed_code("cuFFT", static_cast<int>(_r), #x, __FILE__, __LINE__); \
31 } while (0)
32
33template <typename Scalar>
34struct cufft_c2c_type;
35
36template <>
37struct cufft_c2c_type<float> {
38 static constexpr cufftType value = CUFFT_C2C;
39};
40template <>
41struct cufft_c2c_type<double> {
42 static constexpr cufftType value = CUFFT_Z2Z;
43};
44
45template <typename Scalar>
46struct cufft_r2c_type;
47
48template <>
49struct cufft_r2c_type<float> {
50 static constexpr cufftType value = CUFFT_R2C;
51};
52template <>
53struct cufft_r2c_type<double> {
54 static constexpr cufftType value = CUFFT_D2Z;
55};
56
57template <typename Scalar>
58struct cufft_c2r_type;
59
60template <>
61struct cufft_c2r_type<float> {
62 static constexpr cufftType value = CUFFT_C2R;
63};
64template <>
65struct cufft_c2r_type<double> {
66 static constexpr cufftType value = CUFFT_Z2D;
67};
68
69// Move-only owner of a cufftHandle. Used as an LruCache Value type so eviction
70// destroys the plan through this destructor, with no callback machinery in the
71// cache itself.
72class CufftPlan {
73 public:
74 CufftPlan() = default;
75 explicit CufftPlan(cufftHandle plan) : plan_(plan), owns_(true) {}
76
77 CufftPlan(const CufftPlan&) = delete;
78 CufftPlan& operator=(const CufftPlan&) = delete;
79
80 CufftPlan(CufftPlan&& o) noexcept : plan_(o.plan_), owns_(o.owns_) { o.owns_ = false; }
81
82 CufftPlan& operator=(CufftPlan&& o) noexcept {
83 if (this != &o) {
84 destroy();
85 plan_ = o.plan_;
86 owns_ = o.owns_;
87 o.owns_ = false;
88 }
89 return *this;
90 }
91
92 ~CufftPlan() { destroy(); }
93
94 cufftHandle get() const { return plan_; }
95
96 private:
97 void destroy() noexcept {
98 if (owns_) (void)cufftDestroy(plan_);
99 owns_ = false;
100 }
101
102 cufftHandle plan_{};
103 bool owns_ = false;
104};
105
106inline cufftResult cufftExecC2C_dispatch(cufftHandle plan, std::complex<float>* in, std::complex<float>* out,
107 int direction) {
108 return cufftExecC2C(plan, reinterpret_cast<cufftComplex*>(in), reinterpret_cast<cufftComplex*>(out), direction);
109}
110inline cufftResult cufftExecC2C_dispatch(cufftHandle plan, std::complex<double>* in, std::complex<double>* out,
111 int direction) {
112 return cufftExecZ2Z(plan, reinterpret_cast<cufftDoubleComplex*>(in), reinterpret_cast<cufftDoubleComplex*>(out),
113 direction);
114}
115
116inline cufftResult cufftExecR2C_dispatch(cufftHandle plan, float* in, std::complex<float>* out) {
117 return cufftExecR2C(plan, in, reinterpret_cast<cufftComplex*>(out));
118}
119inline cufftResult cufftExecR2C_dispatch(cufftHandle plan, double* in, std::complex<double>* out) {
120 return cufftExecD2Z(plan, in, reinterpret_cast<cufftDoubleComplex*>(out));
121}
122
123inline cufftResult cufftExecC2R_dispatch(cufftHandle plan, std::complex<float>* in, float* out) {
124 return cufftExecC2R(plan, reinterpret_cast<cufftComplex*>(in), out);
125}
126inline cufftResult cufftExecC2R_dispatch(cufftHandle plan, std::complex<double>* in, double* out) {
127 return cufftExecZ2D(plan, reinterpret_cast<cufftDoubleComplex*>(in), out);
128}
129
130} // namespace internal
131} // namespace gpu
132} // namespace Eigen
133
134#endif // EIGEN_GPU_CUFFT_SUPPORT_H
Namespace containing all symbols from the Eigen library.