// Copyright (C) 2026 Kiyotsugu Arai // SPDX-License-Identifier: LGPL-3.0-or-later // fft_batch.hpp // // Batch FFT — processes multiple same-size signals at once // // Optimization points: // 1. real-pair packing: 2 real signals → 1 complex FFT (stereoFFT) // 2. SIMD batch: process 4 complex FFTs simultaneously with AVX2 (batchFFT) // 3. shared twiddle table: computed only once for N FFTs // 4. L1 cache optimization: 1024 points × 4 signals = 32KB // // Usage: // // FFT a stereo WAV interleave directly (no transpose needed) // sangi::simd_fft::stereoFFT(interleaved, 1024, hannWindow, // specL, specR); // // // FFT N complex signals at once // Complex* signals[8] = { ... }; // sangi::simd_fft::batchFFT(signals, 8, 1024); #ifndef SANGI_FFT_BATCH_HPP #define SANGI_FFT_BATCH_HPP #include #include #include #include #include // ASM kernel declaration (MASM, MSVC x64) #if defined(_MSC_VER) && defined(_M_X64) && defined(SANGI_FFT_HAS_ASM) extern "C" void batch4_butterfly_avx2_float( float* data, int n, int half_len, const float* tw); #define SANGI_BATCH_FFT_HAS_ASM 1 #endif namespace sangi { namespace simd_fft { // ===================================================================== // Internal: batch butterfly (C++ intrinsics fallback) // ===================================================================== namespace detail { // One-stage butterfly over 4-signal interleaved data (AVX2) // data layout: position k → data[k*8 .. k*8+7] // = {s0_re, s0_im, s1_re, s1_im, s2_re, s2_im, s3_re, s3_im} inline void batch4_butterfly_float_cpp( float* data, int n, int half_len, const float* tw) { #if defined(SANGI_FFT_HAS_AVX2) const __m256 sign_mask = _mm256_set_ps( 1.0f, -1.0f, 1.0f, -1.0f, 1.0f, -1.0f, 1.0f, -1.0f); int len = half_len * 2; for (int i = 0; i < n; i += len) { for (int j = 0; j < half_len; ++j) { int k_upper = (i + j) * 8; int k_lower = (i + j + half_len) * 8; __m256 u = _mm256_loadu_ps(&data[k_upper]); __m256 d = _mm256_loadu_ps(&data[k_lower]); // broadcast tw[j] = {re, im} → {re,im,re,im,re,im,re,im} __m256 w = _mm256_castpd_ps( _mm256_broadcast_sd(reinterpret_cast(&tw[j * 2]))); __m256 d_re = _mm256_shuffle_ps(d, d, 0xA0); __m256 d_im = _mm256_shuffle_ps(d, d, 0xF5); __m256 w_flip = _mm256_shuffle_ps(w, w, 0xB1); __m256 p1 = _mm256_mul_ps(d_re, w); __m256 p2 = _mm256_mul_ps(d_im, w_flip); __m256 v = _mm256_add_ps(p1, _mm256_mul_ps(p2, sign_mask)); _mm256_storeu_ps(&data[k_upper], _mm256_add_ps(u, v)); _mm256_storeu_ps(&data[k_lower], _mm256_sub_ps(u, v)); } } #else // scalar fallback int len = half_len * 2; for (int i = 0; i < n; i += len) { for (int j = 0; j < half_len; ++j) { float w_re = tw[j * 2]; float w_im = tw[j * 2 + 1]; for (int s = 0; s < 4; ++s) { int u_idx = (i + j) * 8 + s * 2; int d_idx = (i + j + half_len) * 8 + s * 2; float d_re = data[d_idx]; float d_im = data[d_idx + 1]; float v_re = d_re * w_re - d_im * w_im; float v_im = d_re * w_im + d_im * w_re; float u_re = data[u_idx]; float u_im = data[u_idx + 1]; data[u_idx] = u_re + v_re; data[u_idx + 1] = u_im + v_im; data[d_idx] = u_re - v_re; data[d_idx + 1] = u_im - v_im; } } } #endif } // AoS → interleaved transpose // input: signals[0..3], each n complex // output: interleaved[pos*8 + sig*2 + {re,im}] inline void transpose_to_interleaved( const Complex* const* signals, int count, float* interleaved, int n) { for (int k = 0; k < n; ++k) { int base = k * 8; for (int s = 0; s < count; ++s) { interleaved[base + s * 2] = signals[s][k].real(); interleaved[base + s * 2 + 1] = signals[s][k].imag(); } // zero padding (when count < 4) for (int s = count; s < 4; ++s) { interleaved[base + s * 2] = 0.0f; interleaved[base + s * 2 + 1] = 0.0f; } } } // interleaved → AoS inverse transpose inline void transpose_from_interleaved( const float* interleaved, int n, Complex** signals, int count) { for (int k = 0; k < n; ++k) { int base = k * 8; for (int s = 0; s < count; ++s) { signals[s][k] = Complex( interleaved[base + s * 2], interleaved[base + s * 2 + 1]); } } } // Bit-reversal permutation over interleaved data // Swaps positions for 4 signals at once (each position = 32 bytes) inline void batch_bit_reversal(float* data, int n) { for (int i = 1, j = 0; i < n; ++i) { int bit = n >> 1; while (j & bit) { j ^= bit; bit >>= 1; } j ^= bit; if (i < j) { // swap position i and j (32 bytes each) float* pi = &data[i * 8]; float* pj = &data[j * 8]; #if defined(SANGI_FFT_HAS_AVX2) __m256 a = _mm256_loadu_ps(pi); __m256 b = _mm256_loadu_ps(pj); _mm256_storeu_ps(pi, b); _mm256_storeu_ps(pj, a); #else for (int k = 0; k < 8; ++k) std::swap(pi[k], pj[k]); #endif } } } } // namespace detail // ===================================================================== // batchFFT: process N complex FFTs at once (in-place) // ===================================================================== /// Batch-processes in units of 4 signals. The remainder is handled by the /// ordinary simd_fft::fft(). /// signals: count complex* (each fftSize elements, overwritten in-place) inline void batchFFT(Complex** signals, int count, int fftSize) { if (fftSize <= 1 || count <= 0) return; // twiddle table (shared with the ordinary FFT) thread_local TwiddleTable table; thread_local int cached_n = 0; if (cached_n != fftSize) { table.build(fftSize); cached_n = fftSize; } // allocate the interleave buffer only once (outside the loop) // N=1024 → 32KB, N=4096 → 128KB thread_local std::vector interleaved; if ((int)interleaved.size() < fftSize * 8) interleaved.resize(fftSize * 8); // batch-process 4 signals at a time int batchStart = 0; for (; batchStart + 3 < count; batchStart += 4) { const Complex* src[4] = { signals[batchStart], signals[batchStart + 1], signals[batchStart + 2], signals[batchStart + 3] }; detail::transpose_to_interleaved(src, 4, interleaved.data(), fftSize); // bit-reversal permutation detail::batch_bit_reversal(interleaved.data(), fftSize); // butterflies for all stages int stage = 0; for (int len = 2; len <= fftSize; len <<= 1) { int half = len / 2; const float* tw = reinterpret_cast( table.twiddles[stage].data()); #if defined(SANGI_BATCH_FFT_HAS_ASM) batch4_butterfly_avx2_float(interleaved.data(), fftSize, half, tw); #else detail::batch4_butterfly_float_cpp(interleaved.data(), fftSize, half, tw); #endif ++stage; } // inverse transpose Complex* dst[4] = { signals[batchStart], signals[batchStart + 1], signals[batchStart + 2], signals[batchStart + 3] }; detail::transpose_from_interleaved(interleaved.data(), fftSize, dst, 4); } // remainder uses the ordinary FFT for (int i = batchStart; i < count; ++i) { fft(signals[i], fftSize); } } // ===================================================================== // batchIFFT: process N complex inverse FFTs at once (in-place) // ===================================================================== inline void batchIFFT(Complex** signals, int count, int fftSize) { // conjugate → batchFFT → conjugate + 1/N for (int i = 0; i < count; ++i) { for (int k = 0; k < fftSize; ++k) signals[i][k] = conj(signals[i][k]); } batchFFT(signals, count, fftSize); float inv = 1.0f / static_cast(fftSize); for (int i = 0; i < count; ++i) { for (int k = 0; k < fftSize; ++k) signals[i][k] = conj(signals[i][k]) * inv; } } // ===================================================================== // stereoFFT: stereo interleaved WAV → L/R spectra // ===================================================================== /// interleaved: fftSize*2 floats of {L0,R0,L1,R1,...} /// window: Hann window etc. (fftSize floats), nullptr means no window /// outL, outR: each fftSize/2+1 complex (Hermitian half) inline void stereoFFT( const float* interleaved, int fftSize, const float* window, Complex* outL, Complex* outR) { // apply window treating interleaved as complex std::vector> z(fftSize); if (window) { for (int k = 0; k < fftSize; ++k) { z[k] = Complex( interleaved[k * 2] * window[k], // L * window interleaved[k * 2 + 1] * window[k]); // R * window } } else { for (int k = 0; k < fftSize; ++k) { z[k] = Complex( interleaved[k * 2], interleaved[k * 2 + 1]); } } // a single complex FFT fft(z.data(), fftSize); // L/R separation (using Hermitian symmetry) int half = fftSize / 2; // DC bin (k=0): only Z[0], Z[N] equals Z[0] outL[0] = Complex(z[0].real(), 0.0f); outR[0] = Complex(z[0].imag(), 0.0f); // Nyquist bin (k=N/2) outL[half] = Complex(z[half].real(), 0.0f); outR[half] = Complex(z[half].imag(), 0.0f); // k = 1..N/2-1 for (int k = 1; k < half; ++k) { auto zk = z[k]; auto znk = conj(z[fftSize - k]); // L[k] = (Z[k] + conj(Z[N-k])) / 2 outL[k] = (zk + znk) * 0.5f; // R[k] = (Z[k] - conj(Z[N-k])) / (2i) // = (Z[k] - conj(Z[N-k])) * (-i/2) auto diff = zk - znk; outR[k] = Complex(diff.imag() * 0.5f, -diff.real() * 0.5f); } } // ===================================================================== // stereoIFFT: L/R spectra → 2 real signals (inverse of stereoFFT) // ===================================================================== /// specL, specR: each fftSize/2+1 complex (Hermitian half) /// outL, outR: each fftSize floats (time domain) /// /// Principle: when z[n] = l[n] + i·r[n], then Z[k] = L[k] + i·R[k] /// k=0..N/2 : construct Z[k] directly from specL, specR /// k=N/2+1..N-1: Z[k] = conj(L[N-k]) + i·conj(R[N-k]) /// a single complex IFFT → real=L, imag=R inline void stereoIFFT( const Complex* specL, const Complex* specR, int fftSize, float* outL, float* outR) { int half = fftSize / 2; std::vector> z(fftSize); // k = 0..N/2: Z[k] = L[k] + i·R[k] for (int k = 0; k <= half; ++k) { // (a+bi) + i·(c+di) = (a-d) + (b+c)i z[k] = Complex( specL[k].real() - specR[k].imag(), specL[k].imag() + specR[k].real()); } // k = N/2+1..N-1: Z[k] = conj(L[N-k]) + i·conj(R[N-k]) for (int k = 1; k < half; ++k) { auto Lc = conj(specL[k]); // conj(L[k]) auto Rc = conj(specR[k]); // conj(R[k]) // conj(L[k]) + i·conj(R[k]) = (a+d) + (-b+c)i // where L[k]=a+bi, R[k]=c+di z[fftSize - k] = Complex( Lc.real() - Rc.imag(), Lc.imag() + Rc.real()); } // a single complex IFFT ifft(z.data(), fftSize); // real part = L, imaginary part = R for (int k = 0; k < fftSize; ++k) { outL[k] = z[k].real(); outR[k] = z[k].imag(); } } // ===================================================================== // batchRealFFT: FFT N real signals at once (pair packing + batch) // ===================================================================== /// realSignals: count float* (each fftSize elements) /// outSpectra: count complex* (each fftSize/2+1 elements) /// window: window function (fftSize elements), nullptr means no window inline void batchRealFFT( const float* const* realSignals, int count, int fftSize, Complex** outSpectra, const float* window = nullptr) { int half = fftSize / 2; int pairs = count / 2; // pair 2 signals at a time for a complex FFT std::vector> z(fftSize); std::vector*> complexSignals; for (int p = 0; p < pairs; ++p) { int i0 = p * 2; int i1 = p * 2 + 1; // pack 2 real signals into 1 complex signal if (window) { for (int k = 0; k < fftSize; ++k) { z[k] = Complex( realSignals[i0][k] * window[k], realSignals[i1][k] * window[k]); } } else { for (int k = 0; k < fftSize; ++k) { z[k] = Complex(realSignals[i0][k], realSignals[i1][k]); } } fft(z.data(), fftSize); // separation outSpectra[i0][0] = Complex(z[0].real(), 0.0f); outSpectra[i1][0] = Complex(z[0].imag(), 0.0f); outSpectra[i0][half] = Complex(z[half].real(), 0.0f); outSpectra[i1][half] = Complex(z[half].imag(), 0.0f); for (int k = 1; k < half; ++k) { auto zk = z[k]; auto znk = conj(z[fftSize - k]); outSpectra[i0][k] = (zk + znk) * 0.5f; auto diff = zk - znk; outSpectra[i1][k] = Complex(diff.imag() * 0.5f, -diff.real() * 0.5f); } } // if odd, the last one uses an ordinary real FFT if (count & 1) { int last = count - 1; if (window) { for (int k = 0; k < fftSize; ++k) z[k] = Complex(realSignals[last][k] * window[k], 0.0f); } else { for (int k = 0; k < fftSize; ++k) z[k] = Complex(realSignals[last][k], 0.0f); } fft(z.data(), fftSize); for (int k = 0; k <= half; ++k) outSpectra[last][k] = z[k]; } } } // namespace simd_fft // ===================================================================== // fft_pair / ifft_pair: process 2 real signals simultaneously with a single complex FFT // ===================================================================== // Two-for-one trick: z[k] = x[k] + i·y[k] → FFT(z) → separation // X[k] = (Z[k] + conj(Z[N-k])) / 2 // Y[k] = (Z[k] - conj(Z[N-k])) / (2i) /// @brief Process 2 real signals simultaneously with a single FFT (Two-for-one trick) /// @param x real signal 1 (size fftSize) /// @param y real signal 2 (size fftSize) /// @param fftSize FFT size (power of 2) /// @param outSpecX output spectrum 1 (size fftSize/2+1) /// @param outSpecY output spectrum 2 (size fftSize/2+1) /// @param window window function (fftSize elements, nullptr means no window) template void fft_pair( const Real* x, const Real* y, int fftSize, Complex* outSpecX, Complex* outSpecY, const Real* window = nullptr) { using C = Complex; int half = fftSize / 2; std::vector z(fftSize); if (window) { for (int k = 0; k < fftSize; ++k) z[k] = C(x[k] * window[k], y[k] * window[k]); } else { for (int k = 0; k < fftSize; ++k) z[k] = C(x[k], y[k]); } simd_fft::fft(z.data(), fftSize); // DC, Nyquist outSpecX[0] = C(z[0].real(), Real(0)); outSpecY[0] = C(z[0].imag(), Real(0)); outSpecX[half] = C(z[half].real(), Real(0)); outSpecY[half] = C(z[half].imag(), Real(0)); // k = 1..half-1 for (int k = 1; k < half; ++k) { C zk = z[k]; C znk = conj(z[fftSize - k]); outSpecX[k] = (zk + znk) * Real(0.5); C diff = zk - znk; outSpecY[k] = C(diff.imag() * Real(0.5), -diff.real() * Real(0.5)); } } /// @brief Inverse of fft_pair: recover 2 real signals from 2 spectra /// @param specX spectrum 1 (size fftSize/2+1) /// @param specY spectrum 2 (size fftSize/2+1) /// @param fftSize FFT size (power of 2) /// @param outX output real signal 1 (size fftSize) /// @param outY output real signal 2 (size fftSize) template void ifft_pair( const Complex* specX, const Complex* specY, int fftSize, Real* outX, Real* outY) { using C = Complex; int half = fftSize / 2; // reconstruct Z[k] = X[k] + i·Y[k] std::vector z(fftSize); for (int k = 0; k <= half; ++k) z[k] = C(specX[k].real() - specY[k].imag(), specX[k].imag() + specY[k].real()); // conjugate symmetry: Z[N-k] = conj(Z[k]) for (int k = 1; k < half; ++k) z[fftSize - k] = conj(z[k]); simd_fft::ifft(z.data(), fftSize); for (int k = 0; k < fftSize; ++k) { outX[k] = z[k].real(); outY[k] = z[k].imag(); } } } // namespace sangi #endif // SANGI_FFT_BATCH_HPP