From 31106ef3b07dd1de8720445beafaf0a7cab699d2 Mon Sep 17 00:00:00 2001 From: Christoph Ruegg Date: Wed, 26 Oct 2016 09:43:16 +0200 Subject: [PATCH] FFT-MKL: support for 2D FFT --- src/NativeProviders/MKL/fft.cpp | 24 +++++++++++++++++++ .../Providers/Common/Mkl/SafeNativeMethods.cs | 6 +++++ 2 files changed, 30 insertions(+) diff --git a/src/NativeProviders/MKL/fft.cpp b/src/NativeProviders/MKL/fft.cpp index cd20a16b..a16af563 100644 --- a/src/NativeProviders/MKL/fft.cpp +++ b/src/NativeProviders/MKL/fft.cpp @@ -22,6 +22,20 @@ inline MKL_LONG fft_create_1d(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_LONG n, return status; } +template +inline MKL_LONG fft_create_2d(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_LONG m, const MKL_LONG n, const Precision forward_scale, const Precision backward_scale, const DFTI_CONFIG_VALUE precision, const DFTI_CONFIG_VALUE domain) +{ + MKL_LONG sizes[2]; + sizes[0] = m; + sizes[1] = n; + MKL_LONG status = DftiCreateDescriptor(handle, precision, domain, 2, sizes); + 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 status; +} + template inline MKL_LONG fft_compute(const DFTI_DESCRIPTOR_HANDLE handle, Data x[], FFT fft) { @@ -45,6 +59,16 @@ extern "C" { return fft_create_1d(handle, n, forward_scale, backward_scale, DFTI_SINGLE, DFTI_COMPLEX); } + DLLEXPORT MKL_LONG z_fft_create_2d(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_LONG m, const MKL_LONG n, const double forward_scale, const double backward_scale) + { + return fft_create_2d(handle, m, n, forward_scale, backward_scale, DFTI_DOUBLE, DFTI_COMPLEX); + } + + DLLEXPORT MKL_LONG c_fft_create_2d(DFTI_DESCRIPTOR_HANDLE* handle, const MKL_LONG m, const MKL_LONG n, const float forward_scale, const float backward_scale) + { + return fft_create_2d(handle, m, n, forward_scale, backward_scale, DFTI_SINGLE, DFTI_COMPLEX); + } + DLLEXPORT MKL_LONG z_fft_forward(const DFTI_DESCRIPTOR_HANDLE handle, MKL_Complex16 x[]) { return fft_compute(handle, x, DftiComputeForward); diff --git a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs index c336e8f6..1aa93443 100644 --- a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs +++ b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs @@ -392,6 +392,12 @@ namespace MathNet.Numerics.Providers.Common.Mkl [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] internal static extern int c_fft_create([Out] out IntPtr handle, int n, float forward_scale, float backward_scale); + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int z_fft_create_2d([Out] out IntPtr handle, int m, int n, double forward_scale, double backward_scale); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int c_fft_create_2d([Out] out IntPtr handle, int m, int n, float forward_scale, float backward_scale); + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] internal static extern int z_fft_forward([In] IntPtr handle, [In, Out] Complex[] x);