diff --git a/src/NativeProviders/MKL/fft.cpp b/src/NativeProviders/MKL/fft.cpp index 8fd0c091..c35985b6 100644 --- a/src/NativeProviders/MKL/fft.cpp +++ b/src/NativeProviders/MKL/fft.cpp @@ -7,27 +7,48 @@ #include "mkl_service.h" #include "mkl_dfti.h" +template +inline MKL_LONG fft_inplace(MKL_LONG n, Data x[], DFTI_CONFIG_VALUE precision, DFTI_CONFIG_VALUE domain, FFT fft) +{ + MKL_LONG status = 0; + DFTI_DESCRIPTOR_HANDLE descriptor = 0; + status = DftiCreateDescriptor(&descriptor, precision, domain, 1, n); + if (0 != status) goto failed; + + status = DftiCommitDescriptor(descriptor); + if (0 != status) goto failed; + + status = fft(descriptor, x); + if (0 != status) goto failed; + +cleanup: + DftiFreeDescriptor(&descriptor); + return status; + +failed: + status = 1; + goto cleanup; +} + extern "C" { DLLEXPORT MKL_LONG z_fft_forward_inplace(MKL_LONG n, MKL_Complex16 x[]) { - MKL_LONG status = 0; - DFTI_DESCRIPTOR_HANDLE hand = 0; - status = DftiCreateDescriptor(&hand, DFTI_DOUBLE, DFTI_COMPLEX, 1, n); - if (0 != status) goto failed; - - status = DftiCommitDescriptor(hand); - if (0 != status) goto failed; + return fft_inplace(n, x, DFTI_DOUBLE, DFTI_COMPLEX, DftiComputeForward); + } - status = DftiComputeForward(hand, x); - if (0 != status) goto failed; + DLLEXPORT MKL_LONG c_fft_forward_inplace(MKL_LONG n, MKL_Complex8 x[]) + { + return fft_inplace(n, x, DFTI_SINGLE, DFTI_COMPLEX, DftiComputeForward); + } - cleanup: - DftiFreeDescriptor(&hand); - return status; + DLLEXPORT MKL_LONG z_fft_backward_inplace(MKL_LONG n, MKL_Complex16 x[]) + { + return fft_inplace(n, x, DFTI_DOUBLE, DFTI_COMPLEX, DftiComputeBackward); + } - failed: - status = 1; - goto cleanup; + DLLEXPORT MKL_LONG c_fft_backward_inplace(MKL_LONG n, MKL_Complex8 x[]) + { + return fft_inplace(n, x, DFTI_SINGLE, DFTI_COMPLEX, DftiComputeBackward); } } diff --git a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs index 1c472ffe..49f57311 100644 --- a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs +++ b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs @@ -30,15 +30,36 @@ using System.Numerics; namespace MathNet.Numerics.Providers.FourierTransform.Mkl { - public class MklFourierTransformProvider : ManagedFourierTransformProvider + public class MklFourierTransformProvider : IFourierTransformProvider { - public override void InitializeVerify() + public void InitializeVerify() { } - public override void ForwardInplace(Complex[] complex) + public void ForwardInplace(Complex[] complex) { SafeNativeMethods.z_fft_forward_inplace(complex.Length, complex); } + + public void BackwardInplace(Complex[] complex) + { + SafeNativeMethods.z_fft_backward_inplace(complex.Length, complex); + } + + public Complex[] Forward(Complex[] complexTimeSpace) + { + Complex[] work = new Complex[complexTimeSpace.Length]; + complexTimeSpace.Copy(work); + ForwardInplace(work); + return work; + } + + public Complex[] Backward(Complex[] complexFrequenceSpace) + { + Complex[] work = new Complex[complexFrequenceSpace.Length]; + complexFrequenceSpace.Copy(work); + BackwardInplace(work); + return work; + } } } diff --git a/src/Numerics/Providers/FourierTransform/Mkl/SafeNativeMethods.cs b/src/Numerics/Providers/FourierTransform/Mkl/SafeNativeMethods.cs index 60078c92..27cfcb5a 100644 --- a/src/Numerics/Providers/FourierTransform/Mkl/SafeNativeMethods.cs +++ b/src/Numerics/Providers/FourierTransform/Mkl/SafeNativeMethods.cs @@ -84,6 +84,15 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] internal static extern long z_fft_forward_inplace(long n, [In, Out] Complex[] x); + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long c_fft_forward_inplace(long n, [In, Out] Complex32[] x); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long z_fft_backward_inplace(long n, [In, Out] Complex[] x); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long c_fft_backward_inplace(long n, [In, Out] Complex32[] x); + #endregion FFT // ReSharper restore InconsistentNaming