147TEST_CASE(
"2D c2c FFT with cuFFT, float",
"[math::ft]" )
149 if( !mxlibTest::cudaDeviceAvailable() )
151 WARN(
"CUDA runtime is available but no CUDA device is present" );
157 SECTION(
"out-of-place, forward, default constructed, raw interface" )
159 mx::math::ft::fftT<std::complex<float>, std::complex<float>, 2, 1> fft;
161 fft.plan( 128, 128 );
167 mx::cuda::cudaPtr<std::complex<float>> devIn, devOut;
168 devIn.upload( in.data(), in.rows(), in.cols() );
169 devOut.resize( out.rows(), out.cols() );
171 cufftResult rv = fft( devOut.data(), devIn.data() );
173 REQUIRE( rv == CUFFT_SUCCESS );
175 devOut.download( out.data() );
177 float sin = in.abs2().sum() * ( out.rows() * out.cols() );
178 float sout = out.abs2().sum();
181 REQUIRE_THAT( sin, Catch::Matchers::WithinAbs( sout, ( sin ) * ( 1e-3 ) ) );
184 SECTION(
"out-of-place, forward, default constructed, eigen interface" )
186 mx::math::ft::fftT<std::complex<float>, std::complex<float>, 2, 1> fft;
188 fft.plan( 128, 128 );
194 mx::cuda::cudaPtr<std::complex<float>> devIn, devOut;
195 devIn.upload( in.data(), in.rows(), in.cols() );
196 devOut.resize( out.rows(), out.cols() );
198 cufftResult rv = fft( devOut, devIn );
200 REQUIRE( rv == CUFFT_SUCCESS );
202 devOut.download( out.data() );
204 float sin = in.abs2().sum();
205 float sout = out.abs2().sum() / ( out.rows() * out.cols() );
208 REQUIRE_THAT( sin, Catch::Matchers::WithinAbs( sout, ( sin ) * ( 1e-3 ) ) );
211 SECTION(
"out-of-place, forward, plan constructor" )
213 mx::math::ft::fftT<std::complex<float>, std::complex<float>, 2, 1> fft( 128, 128 );
219 mx::cuda::cudaPtr<std::complex<float>> devIn, devOut;
220 devIn.upload( in.data(), in.rows(), in.cols() );
221 devOut.resize( out.rows(), out.cols() );
223 cufftResult rv = fft( devOut, devIn );
225 REQUIRE( rv == CUFFT_SUCCESS );
227 devOut.download( out.data() );
229 float sin = in.abs2().sum();
230 float sout = out.abs2().sum() / ( in.rows() * out.cols() );
233 REQUIRE_THAT( sin, Catch::Matchers::WithinAbs( sout, ( sin ) * ( 1e-3 ) ) );
236 SECTION(
"out-of-place, backward" )
244 mx::cuda::cudaPtr<std::complex<float>> devIn, devOut;
245 devIn.upload( in.data(), in.rows(), in.cols() );
246 devOut.resize( out.rows(), out.cols() );
248 cufftResult rv = fft( devOut, devIn );
250 REQUIRE( rv == CUFFT_SUCCESS );
252 devOut.download( out.data() );
254 float sin = in.abs2().sum() / ( in.rows() * in.cols() );
255 float sout = out.abs2().sum();
258 REQUIRE_THAT( sin, Catch::Matchers::WithinAbs( sout, ( sin ) * ( 1e-3 ) ) );