diff --git a/src/NativeProviders/MKL/fft.cpp b/src/NativeProviders/MKL/fft.cpp index a3598c89..76afb112 100644 --- a/src/NativeProviders/MKL/fft.cpp +++ b/src/NativeProviders/MKL/fft.cpp @@ -6,49 +6,64 @@ #include #include "mkl_dfti.h" -template -inline MKL_INT64 fft_1d_inplace(const MKL_INT64 n, Data x[], const Precision forward_scale, const Precision backward_scale, const DFTI_CONFIG_VALUE precision, const DFTI_CONFIG_VALUE domain, FFT fft) +inline MKL_INT64 fft_free(DFTI_DESCRIPTOR_HANDLE* handle) { - MKL_LONG status; - DFTI_DESCRIPTOR_HANDLE descriptor = nullptr; - status = DftiCreateDescriptor(&descriptor, precision, domain, 1, static_cast(n)); - if (0 != status) goto cleanup; - - status = DftiSetValue(descriptor, DFTI_FORWARD_SCALE, forward_scale); - if (0 != status) goto cleanup; - - status = DftiSetValue(descriptor, DFTI_BACKWARD_SCALE, backward_scale); - if (0 != status) goto cleanup; - - status = DftiCommitDescriptor(descriptor); - if (0 != status) goto cleanup; + MKL_LONG status = DftiFreeDescriptor(handle); + return static_cast(status); +} - status = fft(descriptor, x); +template +inline MKL_INT64 fft_1d_create(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_INT64 n, const Precision forward_scale, const Precision backward_scale, const DFTI_CONFIG_VALUE precision, const DFTI_CONFIG_VALUE domain) +{ + MKL_LONG status = DftiCreateDescriptor(handle, precision, domain, 1, static_cast(n)); + DFTI_DESCRIPTOR_HANDLE descriptor = *handle; + if (0 == status) status = DftiSetValue(descriptor, DFTI_FORWARD_SCALE, forward_scale); + if (0 == status) status = DftiSetValue(descriptor, DFTI_BACKWARD_SCALE, backward_scale); + if (0 == status) status = DftiCommitDescriptor(descriptor); + return static_cast(status); +} -cleanup: - DftiFreeDescriptor(&descriptor); +template +inline MKL_INT64 fft_1d_inplace(const DFTI_DESCRIPTOR_HANDLE handle, Data x[], FFT fft) +{ + MKL_LONG status = fft(handle, x); return static_cast(status); } extern "C" { - DLLEXPORT MKL_INT64 z_fft_forward_inplace(const MKL_INT64 n, const double scaling, MKL_Complex16 x[]) + DLLEXPORT MKL_INT64 x_fft_free(DFTI_DESCRIPTOR_HANDLE* handle) + { + return fft_free(handle); + } + + DLLEXPORT MKL_INT64 z_fft_create(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_INT64 n, const double forward_scale, const double backward_scale) + { + return fft_1d_create(handle, n, forward_scale, backward_scale, DFTI_DOUBLE, DFTI_COMPLEX); + } + + DLLEXPORT MKL_INT64 c_fft_create(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_INT64 n, const float forward_scale, const float backward_scale) + { + return fft_1d_create(handle, n, forward_scale, backward_scale, DFTI_SINGLE, DFTI_COMPLEX); + } + + DLLEXPORT MKL_INT64 z_fft_forward_inplace(const DFTI_DESCRIPTOR_HANDLE handle, MKL_Complex16 x[]) { - return fft_1d_inplace(n, x, scaling, 1.0, DFTI_DOUBLE, DFTI_COMPLEX, DftiComputeForward); + return fft_1d_inplace(handle, x, DftiComputeForward); } - DLLEXPORT MKL_INT64 c_fft_forward_inplace(const MKL_INT64 n, const float scaling, MKL_Complex8 x[]) + DLLEXPORT MKL_INT64 c_fft_forward_inplace(const DFTI_DESCRIPTOR_HANDLE handle, MKL_Complex8 x[]) { - return fft_1d_inplace(n, x, scaling, 1.0f, DFTI_SINGLE, DFTI_COMPLEX, DftiComputeForward); + return fft_1d_inplace(handle, x, DftiComputeForward); } - DLLEXPORT MKL_INT64 z_fft_backward_inplace(const MKL_INT64 n, const double scaling, MKL_Complex16 x[]) + DLLEXPORT MKL_INT64 z_fft_backward_inplace(const DFTI_DESCRIPTOR_HANDLE handle, MKL_Complex16 x[]) { - return fft_1d_inplace(n, x, 1.0, scaling, DFTI_DOUBLE, DFTI_COMPLEX, DftiComputeBackward); + return fft_1d_inplace(handle, x, DftiComputeBackward); } - DLLEXPORT MKL_INT64 c_fft_backward_inplace(const MKL_INT64 n, const float scaling, MKL_Complex8 x[]) + DLLEXPORT MKL_INT64 c_fft_backward_inplace(const DFTI_DESCRIPTOR_HANDLE handle, MKL_Complex8 x[]) { - return fft_1d_inplace(n, x, 1.0f, scaling, DFTI_SINGLE, DFTI_COMPLEX, DftiComputeBackward); + return fft_1d_inplace(handle, x, DftiComputeBackward); } } diff --git a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs index 287b5c83..170c74fc 100644 --- a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs +++ b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs @@ -28,6 +28,7 @@ #if NATIVE +using System; using System.Numerics; using System.Runtime.InteropServices; using System.Security; @@ -371,16 +372,25 @@ namespace MathNet.Numerics.Providers.Common.Mkl #region FFT [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] - internal static extern long z_fft_forward_inplace(long n, double scaling, [In, Out] Complex[] x); + internal static extern long x_fft_free([In] ref IntPtr handle); [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] - internal static extern long c_fft_forward_inplace(long n, float scaling, [In, Out] Complex32[] x); + internal static extern long z_fft_create([Out] out IntPtr handle, long n, double forward_scale, double backward_scale); [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] - internal static extern long z_fft_backward_inplace(long n, double scaling, [In, Out] Complex[] x); + internal static extern long c_fft_create([Out] out IntPtr handle, long n, float forward_scale, float backward_scale); [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] - internal static extern long c_fft_backward_inplace(long n, float scaling, [In, Out] Complex32[] x); + internal static extern long z_fft_forward_inplace([In] IntPtr handle, [In, Out] Complex[] x); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long c_fft_forward_inplace([In] IntPtr handle, [In, Out] Complex32[] x); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long z_fft_backward_inplace([In] IntPtr handle, [In, Out] Complex[] x); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern long c_fft_backward_inplace([In] IntPtr handle, [In, Out] Complex32[] x); #endregion FFT diff --git a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs index eb7b413d..68023cd2 100644 --- a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs +++ b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs @@ -128,12 +128,18 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl public void ForwardInplace(Complex[] complex, FourierTransformScaling scaling) { - SafeNativeMethods.z_fft_forward_inplace(complex.Length, ForwardScaling(scaling, complex.Length), complex); + IntPtr handle; + SafeNativeMethods.z_fft_create(out handle, complex.Length, ForwardScaling(scaling, complex.Length), 1.0); + SafeNativeMethods.z_fft_forward_inplace(handle, complex); + SafeNativeMethods.x_fft_free(ref handle); } public void BackwardInplace(Complex[] complex, FourierTransformScaling scaling) { - SafeNativeMethods.z_fft_backward_inplace(complex.Length, BackwardScaling(scaling, complex.Length), complex); + IntPtr handle; + SafeNativeMethods.z_fft_create(out handle, complex.Length, 1.0, BackwardScaling(scaling, complex.Length)); + SafeNativeMethods.z_fft_backward_inplace(handle, complex); + SafeNativeMethods.x_fft_free(ref handle); } public Complex[] Forward(Complex[] complexTimeSpace, FourierTransformScaling scaling)