#include "wavtool_fft.h" #include "pffft.h" #include WavtoolFFT::WavtoolFFT() { for (auto sz : {64, 96, 128, 160, 192, 256, 384, 480, 512, 640, 768, 800, 1024, 2048, 2400, 4096, 8192, 9216, 16384, 32768, 65536, 131072, 262144}) { pffft_setups[sz] = std::make_pair(pffft_new_setup(sz, PFFFT_REAL), pffft_new_setup(sz, PFFFT_COMPLEX)); } } WavtoolFFT::~WavtoolFFT() { for (const auto &x : pffft_setups) { pffft_destroy_setup(x.second.first); pffft_destroy_setup(x.second.second); } } std::vector WavtoolFFT::get_supported_sizes() { std::vector res; res.reserve(pffft_setups.size()); for (const auto &x : pffft_setups) { res.push_back(x.first); } return res; } int WavtoolFFT::fft_ordered(int size, TransformType type, TransformDirection direction, const float *input, float *output) { auto it = pffft_setups.find(size); if (it == pffft_setups.end()) { return -1; } pffft_transform_ordered( type == TransformType::REAL ? it->second.first : it->second.second, input, output, nullptr, direction == TransformDirection::FORWARD ? PFFFT_FORWARD : PFFFT_BACKWARD); return 0; } int WavtoolFFT::fft_unordered(int size, TransformType type, TransformDirection direction, const float *input, float *output) { auto it = pffft_setups.find(size); if (it == pffft_setups.end()) { return -1; } pffft_transform(type == TransformType::REAL ? it->second.first : it->second.second, input, output, nullptr, direction == TransformDirection::FORWARD ? PFFFT_FORWARD : PFFFT_BACKWARD); return 0; } int WavtoolFFT::fft_convolve_accumulate(int size, TransformType type, float scale, const float *in1, const float *in2, float *output) { auto it = pffft_setups.find(size); if (it == pffft_setups.end()) { return -1; } pffft_zconvolve_accumulate(type == TransformType::REAL ? it->second.first : it->second.second, in1, in2, output, scale); return 0; } int WavtoolFFT::fft_reorder(int size, TransformType type, const float *input, float *output) { auto it = pffft_setups.find(size); if (it == pffft_setups.end()) { return -1; } pffft_zreorder(type == TransformType::REAL ? it->second.first : it->second.second, input, output, PFFFT_FORWARD); return 0; } PFFFT_Setup *WavtoolFFT::get_pffft_setup(int size, TransformType type) { auto it = pffft_setups.find(size); if (it == pffft_setups.end()) { return nullptr; } return type == TransformType::REAL ? it->second.first : it->second.second; } WavtoolFFT wavtool_fft; extern "C" { int fft_get_supported_size_count() { auto x = wavtool_fft.get_supported_sizes(); return x.size(); } void fft_get_supported_sizes(int *output) { auto x = wavtool_fft.get_supported_sizes(); for (int i = 0; i < x.size(); i++) { output[i] = x[i]; } } int fft_ordered(int size, int type, int direction, float *input, float *output) { return wavtool_fft.fft_ordered( size, static_cast(type), static_cast(direction), input, output); } int fft_convolve_real(int size, float *in_fft, float *in_real, float *output) { auto in_res = wavtool_fft.fft_ordered(size, WavtoolFFT::TransformType::REAL, WavtoolFFT::TransformDirection::FORWARD, in_real, output); if (in_res != 0) { return in_res; } float scale = 1.0f / size; for (int i = 0; i < size; i += 2) { auto in_fft_re = in_fft[i]; auto in_fft_im = in_fft[i + 1]; auto in_real_re = output[i]; auto in_real_im = output[i + 1]; output[i] = (in_fft_re * in_real_re - in_fft_im * in_real_im) * scale; output[i + 1] = (in_fft_re * in_real_im + in_fft_im * in_real_re) * scale; } return wavtool_fft.fft_ordered(size, WavtoolFFT::TransformType::REAL, WavtoolFFT::TransformDirection::BACKWARD, output, output); } float *fft_alloc_aligned(int size) { return (float *)pffft_aligned_malloc(size * sizeof(float)); } void fft_free_aligned(float *ptr) { pffft_aligned_free(ptr); } } TEST_CASE("fft tests", "[fft]") { WavtoolFFT f; for (auto size : f.get_supported_sizes()) { // complex { std::vector time(2 * size); std::vector freq(2 * size); for (int i = 0; i < size; i++) { time[i * 2] = std::sin(i); time[i * 2 + 1] = 0.f; } auto res = f.fft_ordered(size, WavtoolFFT::TransformType::COMPLEX, WavtoolFFT::TransformDirection::FORWARD, time.data(), freq.data()); REQUIRE(res == 0); for (int i = 0; i < 2 * size; i++) { time[i] = 0.f; } res = f.fft_ordered(size, WavtoolFFT::TransformType::COMPLEX, WavtoolFFT::TransformDirection::BACKWARD, freq.data(), time.data()); REQUIRE(res == 0); for (int i = 0; i < size; ++i) { const auto delta = std::fabs(time[i * 2] / (float)size - std::sin(i)); if (delta > 1e-5f) { FAIL("FFT error"); } } } // real { std::vector time(size); std::vector freq(size); for (int i = 0; i < size; i++) { time[i] = std::sin(i); } auto res = f.fft_ordered(size, WavtoolFFT::TransformType::REAL, WavtoolFFT::TransformDirection::FORWARD, time.data(), freq.data()); REQUIRE(res == 0); for (int i = 0; i < size; i++) { time[i] = 0.f; } res = f.fft_ordered(size, WavtoolFFT::TransformType::REAL, WavtoolFFT::TransformDirection::BACKWARD, freq.data(), time.data()); REQUIRE(res == 0); for (int i = 0; i < size; ++i) { const auto delta = std::fabs(time[i] / (float)size - std::sin(i)); if (delta > 1e-5f) { FAIL("FFT error"); } } } } }