diff --git a/src/Numerics/Control.cs b/src/Numerics/Control.cs index e7f9fd23..dfd8fe32 100644 --- a/src/Numerics/Control.cs +++ b/src/Numerics/Control.cs @@ -30,6 +30,7 @@ using MathNet.Numerics.Providers.LinearAlgebra; using System; using System.Threading.Tasks; +using MathNet.Numerics.Providers.FourierTransform; namespace MathNet.Numerics { @@ -45,6 +46,7 @@ namespace MathNet.Numerics static int _parallelizeOrder; static int _parallelizeElements; static ILinearAlgebraProvider _linearAlgebraProvider; + static IFourierTransformProvider _fourierTransformProvider; static readonly object _staticLock = new object(); static Control() @@ -66,11 +68,11 @@ namespace MathNet.Numerics TaskScheduler = TaskScheduler.Default; } - private static void InitializeDefaultLinearAlgebraProvider() + private static void InitializeDefaultProviders() { lock (_staticLock) { - if (_linearAlgebraProvider == null) + if (_linearAlgebraProvider == null || _fourierTransformProvider == null) { #if NATIVE try @@ -113,6 +115,7 @@ namespace MathNet.Numerics public static void UseManaged() { LinearAlgebraProvider = new ManagedLinearAlgebraProvider(); + FourierTransformProvider = new ManagedFourierTransformProvider(); } #if NATIVE @@ -124,6 +127,7 @@ namespace MathNet.Numerics public static void UseNativeMKL() { LinearAlgebraProvider = new Providers.LinearAlgebra.Mkl.MklLinearAlgebraProvider(); + FourierTransformProvider = new Providers.FourierTransform.Mkl.MklFourierTransformProvider(); } /// @@ -137,6 +141,7 @@ namespace MathNet.Numerics Providers.LinearAlgebra.Mkl.MklAccuracy accuracy = Providers.LinearAlgebra.Mkl.MklAccuracy.High) { LinearAlgebraProvider = new Providers.LinearAlgebra.Mkl.MklLinearAlgebraProvider(consistency, precision, accuracy); + FourierTransformProvider = new Providers.FourierTransform.Mkl.MklFourierTransformProvider(); } /// @@ -158,6 +163,10 @@ namespace MathNet.Numerics public static void UseNativeCUDA() { LinearAlgebraProvider = new Providers.LinearAlgebra.Cuda.CudaLinearAlgebraProvider(); + if (_fourierTransformProvider == null) + { + FourierTransformProvider = new ManagedFourierTransformProvider(); + } } /// @@ -179,6 +188,10 @@ namespace MathNet.Numerics public static void UseNativeOpenBLAS() { LinearAlgebraProvider = new Providers.LinearAlgebra.OpenBlas.OpenBlasLinearAlgebraProvider(); + if (_fourierTransformProvider == null) + { + FourierTransformProvider = new ManagedFourierTransformProvider(); + } } /// @@ -226,6 +239,7 @@ namespace MathNet.Numerics ThreadSafeRandomNumberGenerators = false; LinearAlgebraProvider.InitializeVerify(); + FourierTransformProvider.InitializeVerify(); } public static void UseMultiThreading() @@ -234,6 +248,7 @@ namespace MathNet.Numerics ThreadSafeRandomNumberGenerators = true; LinearAlgebraProvider.InitializeVerify(); + FourierTransformProvider.InitializeVerify(); } /// @@ -266,7 +281,7 @@ namespace MathNet.Numerics get { if (_linearAlgebraProvider == null) - InitializeDefaultLinearAlgebraProvider(); + InitializeDefaultProviders(); return _linearAlgebraProvider; } @@ -279,6 +294,28 @@ namespace MathNet.Numerics } } + /// + /// Gets or sets the fourier transform provider. Consider to use UseNativeMKL or UseManaged instead. + /// + /// The linear algebra provider. + public static IFourierTransformProvider FourierTransformProvider + { + get + { + if (_fourierTransformProvider == null) + InitializeDefaultProviders(); + + return _fourierTransformProvider; + } + set + { + value.InitializeVerify(); + + // only actually set if verification did not throw + _fourierTransformProvider = value; + } + } + /// /// Gets or sets a value indicating how many parallel worker threads shall be used /// when parallelization is applicable. @@ -293,6 +330,7 @@ namespace MathNet.Numerics // Reinitialize providers: LinearAlgebraProvider.InitializeVerify(); + FourierTransformProvider.InitializeVerify(); } } diff --git a/src/Numerics/Numerics.csproj b/src/Numerics/Numerics.csproj index d224559c..333f4d05 100644 --- a/src/Numerics/Numerics.csproj +++ b/src/Numerics/Numerics.csproj @@ -174,6 +174,7 @@ + diff --git a/src/Numerics/Providers/FourierTransform/IFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/IFourierTransformProvider.cs index 7c6e5c77..28256d78 100644 --- a/src/Numerics/Providers/FourierTransform/IFourierTransformProvider.cs +++ b/src/Numerics/Providers/FourierTransform/IFourierTransformProvider.cs @@ -36,6 +36,11 @@ namespace MathNet.Numerics.Providers.FourierTransform public interface IFourierTransformProvider { + /// + /// Initialize and verify that the provided is indeed available. If not, fall back to alternatives like the managed provider + /// + void InitializeVerify(); + void ForwardInplace(Complex[] complex); void BackwardInplace(Complex[] complex); diff --git a/src/Numerics/Providers/FourierTransform/ManagedFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/ManagedFourierTransformProvider.cs index f4afed72..58c3320e 100644 --- a/src/Numerics/Providers/FourierTransform/ManagedFourierTransformProvider.cs +++ b/src/Numerics/Providers/FourierTransform/ManagedFourierTransformProvider.cs @@ -5,6 +5,10 @@ namespace MathNet.Numerics.Providers.FourierTransform { public class ManagedFourierTransformProvider : IFourierTransformProvider { + public virtual void InitializeVerify() + { + } + public void ForwardInplace(Complex[] complex) { Fourier.BluesteinForward(complex, FourierOptions.Default); diff --git a/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs new file mode 100644 index 00000000..9f99097c --- /dev/null +++ b/src/Numerics/Providers/FourierTransform/Mkl/MklFourierTransformProvider.cs @@ -0,0 +1,14 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Text; + +namespace MathNet.Numerics.Providers.FourierTransform.Mkl +{ + public class MklFourierTransformProvider : ManagedFourierTransformProvider + { + public override void InitializeVerify() + { + } + } +}