Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
GpuFFT.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// GPU FFT via cuFFT: 1D and 2D C2C, R2C, and C2R transforms with plan caching.
12//
13// The stream and cuBLAS handle come from a bound gpu::Context, which defaults to
14// Context::threadLocal() so an FFT shares a stream with the thread's other GPU
15// work; pass an explicit Context to bind elsewhere.
16//
17// Inverse transforms are scaled by 1/n (1D) or 1/(n*m) (2D), so inv(fwd(x)) == x
18// as in Eigen's CPU FFT.
19//
20// Not thread-safe: concurrent fwd/inv calls on one instance race on the cached
21// plans and the bound Context. Use one FFT per thread.
22
23#ifndef EIGEN_GPU_FFT_H
24#define EIGEN_GPU_FFT_H
25
26// IWYU pragma: private
27#include "./InternalHeaderCheck.h"
28
29#include "./CuFftSupport.h"
30#include "./CuBlasSupport.h"
31#include "./GpuContext.h"
32
33namespace Eigen {
34namespace gpu {
35// Default capacity of the per-FFT-instance cuFFT plan cache. Override by
36// passing a capacity to the FFT constructor.
37static constexpr std::size_t kDefaultCufftPlanCacheCapacity = 16;
38
39template <typename Scalar_>
40class FFT {
41 public:
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>;
47
56 explicit FFT(std::size_t plan_cache_capacity = kDefaultCufftPlanCacheCapacity)
57 : ctx_(&Context::threadLocal()), plans_(plan_cache_capacity > 0 ? plan_cache_capacity : 1) {}
58
65 explicit FFT(Context& ctx, std::size_t plan_cache_capacity = kDefaultCufftPlanCacheCapacity)
66 : ctx_(&ctx), plans_(plan_cache_capacity > 0 ? plan_cache_capacity : 1) {}
67
68 // Destructor is implicit: ~LruCache destroys each CufftPlan, which calls
69 // cufftDestroy via CufftPlan's destructor.
70
72 std::size_t plan_cache_capacity() const { return plans_.capacity(); }
73
75 std::size_t plan_cache_size() const { return plans_.size(); }
76
77 FFT(const FFT&) = delete;
78 FFT& operator=(const FFT&) = delete;
79
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);
86
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()));
90
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));
94
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()));
99 return result;
100 }
101
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);
109
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()));
113
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));
117
118 // Scale by 1/n.
119 EIGEN_CUBLAS_CHECK(
120 internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n), static_cast<Complex*>(d_out_.get()), 1));
121
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()));
126 return result;
127 }
128
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);
135
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()));
140
141 cufftHandle plan = get_plan_1d(n, internal::cufft_r2c_type<Scalar>::value);
142 EIGEN_CUFFT_CHECK(
143 internal::cufftExecR2C_dispatch(plan, static_cast<Scalar*>(d_in_.get()), static_cast<Complex*>(d_out_.get())));
144
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()));
149 return result;
150 }
151
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);
162
163 ensure_buffers(n_complex * sizeof(Complex), n * sizeof(Scalar));
164 // cuFFT C2R may overwrite the input, so we copy to d_in_.
165 EIGEN_CUDA_RUNTIME_CHECK(cudaMemcpyAsync(d_in_.get(), input.data(), n_complex * sizeof(Complex),
166 cudaMemcpyHostToDevice, ctx_->stream()));
167
168 cufftHandle plan = get_plan_1d(n, internal::cufft_c2r_type<Scalar>::value);
169 EIGEN_CUFFT_CHECK(
170 internal::cufftExecC2R_dispatch(plan, static_cast<Complex*>(d_in_.get()), static_cast<Scalar*>(d_out_.get())));
171
172 // Scale by 1/n.
173 EIGEN_CUBLAS_CHECK(
174 internal::cublasXscal(ctx_->cublasHandle(), n, Scalar(1) / Scalar(n), static_cast<Scalar*>(d_out_.get()), 1));
175
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()));
180 return result;
181 }
182
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);
191
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()));
195
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));
199
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()));
204 return result;
205 }
206
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);
215
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()));
219
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));
223
224 // Scale by 1/(rows*cols).
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));
228
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()));
233 return result;
234 }
235
236 // DeviceMatrix in / DeviceMatrix out: no host transfer and no host
237 // synchronization. Cross-stream safety follows the DeviceMatrix event
238 // protocol (waitReady before reading, recordReady after enqueuing). 1D
239 // overloads expect column vectors.
240
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);
246 if (n == 0) return;
247 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
248 EIGEN_CUFFT_CHECK(
249 internal::cufftExecC2C_dispatch(plan, const_cast<Complex*>(d_x.data()), d_X.data(), CUFFT_FORWARD));
250 d_X.recordReady(ctx_->stream());
251 }
252
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);
258 if (n == 0) return;
259 cufftHandle plan = get_plan_1d(n, internal::cufft_c2c_type<Scalar>::value);
260 EIGEN_CUFFT_CHECK(
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());
264 }
265
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);
272 if (n == 0) return;
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());
276 }
277
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);
287 if (n == 0) return;
288 // Stage through d_in_: cuFFT C2R may overwrite the input array.
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());
296 }
297
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);
305 EIGEN_CUFFT_CHECK(
306 internal::cufftExecC2C_dispatch(plan, const_cast<Complex*>(d_A.data()), d_B.data(), CUFFT_FORWARD));
307 d_B.recordReady(ctx_->stream());
308 }
309
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);
317 EIGEN_CUFFT_CHECK(
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));
320 EIGEN_CUBLAS_CHECK(
321 internal::cublasXscal(ctx_->cublasHandle(), total_elems, Scalar(1) / Scalar(total_elems), d_B.data(), 1));
322 d_B.recordReady(ctx_->stream());
323 }
324
326 cudaStream_t stream() const { return ctx_->stream(); }
327
329 Context& context() const { return *ctx_; }
330
331 private:
332 Context* 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;
338
339 // Common device-transform prologue: alias check, input/output event waits,
340 // and (destructive) output resize.
341 template <typename InScalar, typename OutScalar>
342 void prepare_out(const DeviceMatrix<InScalar>& in, DeviceMatrix<OutScalar>& out, Index out_rows, Index out_cols) {
343 eigen_assert(
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);
349 }
350
351 // Buffers grow but never shrink. The pre-realloc sync drains the *bound*
352 // Context's stream — including unrelated GEMMs/solves/`device(ctx) = ...`
353 // assignments queued on it — so callers running FFTs alongside other GPU
354 // work on the same Context should size up front (call fwd/inv with the
355 // largest expected n once) to avoid mid-pipeline stalls.
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;
361 }
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;
366 }
367 }
368
369 // Plan key encoding: rank (1 bit) | type (4 bits) | dims.
370 // cufftType uses 7 bits; the top 3 (precision discriminator) are redundant
371 // since Scalar fixes precision per FFT instance, so mask to 4 bits — without
372 // it, e.g. plan_key_1d(5, C2C) and plan_key_1d(7, C2C) collide.
373 static constexpr int64_t kTypeMask = 0xF;
374 static constexpr int kCols2DBits = 30; // bits 5..34
375 static constexpr int kRows2DBits = 29; // bits 35..63
376 static int64_t plan_key_1d(int n, cufftType type) { return (int64_t(n) << 5) | (int64_t(type & kTypeMask) << 1) | 0; }
377
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;
384 }
385
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();
389
390 cufftHandle plan;
391 EIGEN_CUFFT_CHECK(cufftPlan1d(&plan, n, type, /*batch=*/1));
392 EIGEN_CUFFT_CHECK(cufftSetStream(plan, ctx_->stream()));
393 return plans_.insert(key, internal::CufftPlan(plan))->get();
394 }
395
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();
399
400 // cuFFT uses row-major (C order) for 2D: first dim = rows, second = cols.
401 // Eigen matrices are column-major, so we pass (cols, rows) to cuFFT
402 // to get the correct 2D transform.
403 cufftHandle plan;
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();
407 }
408};
409} // namespace gpu
410} // namespace Eigen
411
412#endif // EIGEN_GPU_FFT_H
Namespace containing all symbols from the Eigen library.