Browse Source

FFT-MKL: descriptor reuse in multi-threading scenario

benchmark-la
Christoph Ruegg 10 years ago
parent
commit
ec91f3bc20
  1. 89
      src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs

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

@ -30,15 +30,21 @@
using System; using System;
using System.Numerics; using System.Numerics;
using System.Threading;
using MathNet.Numerics.Providers.Common.Mkl; using MathNet.Numerics.Providers.Common.Mkl;
namespace MathNet.Numerics.Providers.FourierTransform.Mkl namespace MathNet.Numerics.Providers.FourierTransform.Mkl
{ {
public class MklFourierTransformProvider : IFourierTransformProvider public class MklFourierTransformProvider : IFourierTransformProvider, IDisposable
{ {
IntPtr _currentHandle = IntPtr.Zero; class Kernel
int _currentLength = -1; {
FourierTransformScaling _currentScaling = FourierTransformScaling.NoScaling; public IntPtr Handle;
public int Length;
public FourierTransformScaling Scaling;
}
Kernel _kernel;
/// <summary> /// <summary>
/// Try to find out whether the provider is available, at least in principle. /// Try to find out whether the provider is available, at least in principle.
@ -70,6 +76,12 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl
/// </summary> /// </summary>
public void FreeBuffers() public void FreeBuffers()
{ {
Kernel kernel = Interlocked.Exchange(ref _kernel, null);
if (kernel != null)
{
SafeNativeMethods.x_fft_free(ref kernel.Handle);
}
MklProvider.FreeBuffers(); MklProvider.FreeBuffers();
} }
@ -130,46 +142,54 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl
return MklProvider.Describe(); return MklProvider.Describe();
} }
public void ForwardInplace(Complex[] complex, FourierTransformScaling scaling) Kernel Configure(int length, FourierTransformScaling scaling)
{ {
if (_currentHandle == IntPtr.Zero) Kernel kernel = Interlocked.Exchange(ref _kernel, null);
if (kernel == null)
{ {
SafeNativeMethods.z_fft_create(out _currentHandle, complex.Length, ForwardScaling(scaling, complex.Length), BackwardScaling(scaling, complex.Length)); kernel = new Kernel
{
Length = length,
Scaling = scaling
};
SafeNativeMethods.z_fft_create(out kernel.Handle, length, ForwardScaling(scaling, length), BackwardScaling(scaling, length));
return kernel;
} }
else
if (kernel.Length != length || kernel.Scaling != scaling)
{ {
if (complex.Length != _currentLength || scaling != _currentScaling) SafeNativeMethods.x_fft_free(ref kernel.Handle);
{ SafeNativeMethods.z_fft_create(out kernel.Handle, length, ForwardScaling(scaling, length), BackwardScaling(scaling, length));
SafeNativeMethods.x_fft_free(ref _currentHandle); kernel.Length = length;
_currentHandle = IntPtr.Zero; kernel.Scaling = scaling;
SafeNativeMethods.z_fft_create(out _currentHandle, complex.Length, ForwardScaling(scaling, complex.Length), BackwardScaling(scaling, complex.Length)); return kernel;
_currentLength = complex.Length;
_currentScaling = scaling;
}
} }
SafeNativeMethods.z_fft_forward_inplace(_currentHandle, complex); return kernel;
} }
public void BackwardInplace(Complex[] complex, FourierTransformScaling scaling) void Release(Kernel kernel)
{ {
if (_currentHandle == IntPtr.Zero) Kernel existing = Interlocked.Exchange(ref _kernel, kernel);
if (existing != null)
{ {
SafeNativeMethods.z_fft_create(out _currentHandle, complex.Length, ForwardScaling(scaling, complex.Length), BackwardScaling(scaling, complex.Length)); SafeNativeMethods.x_fft_free(ref existing.Handle);
}
else
{
if (complex.Length != _currentLength || scaling != _currentScaling)
{
SafeNativeMethods.x_fft_free(ref _currentHandle);
_currentHandle = IntPtr.Zero;
SafeNativeMethods.z_fft_create(out _currentHandle, complex.Length, ForwardScaling(scaling, complex.Length), BackwardScaling(scaling, complex.Length));
_currentLength = complex.Length;
_currentScaling = scaling;
}
} }
}
public void ForwardInplace(Complex[] complex, FourierTransformScaling scaling)
{
Kernel kernel = Configure(complex.Length, scaling);
SafeNativeMethods.z_fft_forward_inplace(kernel.Handle, complex);
Release(kernel);
}
SafeNativeMethods.z_fft_backward_inplace(_currentHandle, complex); public void BackwardInplace(Complex[] complex, FourierTransformScaling scaling)
{
Kernel kernel = Configure(complex.Length, scaling);
SafeNativeMethods.z_fft_backward_inplace(kernel.Handle, complex);
Release(kernel);
} }
public Complex[] Forward(Complex[] complexTimeSpace, FourierTransformScaling scaling) public Complex[] Forward(Complex[] complexTimeSpace, FourierTransformScaling scaling)
@ -213,6 +233,11 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl
return 1.0; return 1.0;
} }
} }
public void Dispose()
{
FreeBuffers();
}
} }
} }

Loading…
Cancel
Save