From ec91f3bc208969562af92d0e30636a490187a5a5 Mon Sep 17 00:00:00 2001 From: Christoph Ruegg Date: Sat, 15 Oct 2016 16:52:50 +0200 Subject: [PATCH] FFT-MKL: descriptor reuse in multi-threading scenario --- .../Mkl/MklFourierTransformProvider.cs | 89 ++++++++++++------- 1 file changed, 57 insertions(+), 32 deletions(-) diff --git a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs index e6a15c1f..494de40a 100644 --- a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs +++ b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs @@ -30,15 +30,21 @@ using System; using System.Numerics; +using System.Threading; using MathNet.Numerics.Providers.Common.Mkl; namespace MathNet.Numerics.Providers.FourierTransform.Mkl { - public class MklFourierTransformProvider : IFourierTransformProvider + public class MklFourierTransformProvider : IFourierTransformProvider, IDisposable { - IntPtr _currentHandle = IntPtr.Zero; - int _currentLength = -1; - FourierTransformScaling _currentScaling = FourierTransformScaling.NoScaling; + class Kernel + { + public IntPtr Handle; + public int Length; + public FourierTransformScaling Scaling; + } + + Kernel _kernel; /// /// Try to find out whether the provider is available, at least in principle. @@ -70,6 +76,12 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl /// public void FreeBuffers() { + Kernel kernel = Interlocked.Exchange(ref _kernel, null); + if (kernel != null) + { + SafeNativeMethods.x_fft_free(ref kernel.Handle); + } + MklProvider.FreeBuffers(); } @@ -130,46 +142,54 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl 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 _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; - } + SafeNativeMethods.x_fft_free(ref kernel.Handle); + SafeNativeMethods.z_fft_create(out kernel.Handle, length, ForwardScaling(scaling, length), BackwardScaling(scaling, length)); + kernel.Length = length; + kernel.Scaling = scaling; + return kernel; } - 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)); - } - 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; - } + SafeNativeMethods.x_fft_free(ref existing.Handle); } + } + + 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) @@ -213,6 +233,11 @@ namespace MathNet.Numerics.Providers.FourierTransform.Mkl return 1.0; } } + + public void Dispose() + { + FreeBuffers(); + } } }