Browse Source

FFT-MKL: wip towards reuseable descriptors

benchmark-la
Christoph Ruegg 10 years ago
parent
commit
608fcf4d91
  1. 67
      src/NativeProviders/MKL/fft.cpp
  2. 18
      src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs
  3. 10
      src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs

67
src/NativeProviders/MKL/fft.cpp

@ -6,49 +6,64 @@
#include <float.h> #include <float.h>
#include "mkl_dfti.h" #include "mkl_dfti.h"
template<typename Data, typename Precision, typename FFT> inline MKL_INT64 fft_free(DFTI_DESCRIPTOR_HANDLE* handle)
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)
{ {
MKL_LONG status; MKL_LONG status = DftiFreeDescriptor(handle);
DFTI_DESCRIPTOR_HANDLE descriptor = nullptr; return static_cast<MKL_INT64>(status);
status = DftiCreateDescriptor(&descriptor, precision, domain, 1, static_cast<MKL_LONG>(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;
status = fft(descriptor, x); template<typename Precision>
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<MKL_LONG>(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<MKL_INT64>(status);
}
cleanup: template<typename Data, typename FFT>
DftiFreeDescriptor(&descriptor); inline MKL_INT64 fft_1d_inplace(const DFTI_DESCRIPTOR_HANDLE handle, Data x[], FFT fft)
{
MKL_LONG status = fft(handle, x);
return static_cast<MKL_INT64>(status); return static_cast<MKL_INT64>(status);
} }
extern "C" { 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);
} }
} }

18
src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs

@ -28,6 +28,7 @@
#if NATIVE #if NATIVE
using System;
using System.Numerics; using System.Numerics;
using System.Runtime.InteropServices; using System.Runtime.InteropServices;
using System.Security; using System.Security;
@ -371,16 +372,25 @@ namespace MathNet.Numerics.Providers.Common.Mkl
#region FFT #region FFT
[DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] [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)] [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)] [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)] [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 #endregion FFT

10
src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs

@ -128,12 +128,18 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl
public void ForwardInplace(Complex[] complex, FourierTransformScaling scaling) 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) 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) public Complex[] Forward(Complex[] complexTimeSpace, FourierTransformScaling scaling)

Loading…
Cancel
Save