From b47d2191d1e4a7771548071901c9561a52fb9e89 Mon Sep 17 00:00:00 2001 From: Christoph Ruegg Date: Wed, 21 Feb 2018 09:37:55 +0100 Subject: [PATCH] Native Providers: provide access to features like memory management, skip reloading if already loaded --- .../Providers/Common/Cuda/CudaProvider.cs | 22 ++++- .../Providers/Common/Mkl/MklProvider.cs | 95 ++++++++++++++++--- .../Common/OpenBlas/OpenBlasProvider.cs | 33 ++++++- 3 files changed, 131 insertions(+), 19 deletions(-) diff --git a/src/Numerics/Providers/Common/Cuda/CudaProvider.cs b/src/Numerics/Providers/Common/Cuda/CudaProvider.cs index 4f3c1a52..5f76bec2 100644 --- a/src/Numerics/Providers/Common/Cuda/CudaProvider.cs +++ b/src/Numerics/Providers/Common/Cuda/CudaProvider.cs @@ -34,15 +34,21 @@ using System.Collections.Generic; namespace MathNet.Numerics.Providers.Common.Cuda { - internal static class CudaProvider + public static class CudaProvider { static int _nativeRevision; static bool _nativeX86; static bool _nativeX64; static bool _nativeIA64; + static bool _loaded; internal static bool IsAvailable(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return true; + } + try { if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) @@ -63,6 +69,11 @@ namespace MathNet.Numerics.Providers.Common.Cuda internal static void Load(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return; + } + int a, b; try { @@ -93,10 +104,17 @@ namespace MathNet.Numerics.Providers.Common.Cuda { throw new NotSupportedException("Cuda Native Provider too old. Consider upgrading to a newer version."); } + + _loaded = true; } - internal static string Describe() + public static string Describe() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + var parts = new List(); if (_nativeX86) parts.Add("x86"); if (_nativeX64) parts.Add("x64"); diff --git a/src/Numerics/Providers/Common/Mkl/MklProvider.cs b/src/Numerics/Providers/Common/Mkl/MklProvider.cs index 2dec55cf..935ad81c 100644 --- a/src/Numerics/Providers/Common/Mkl/MklProvider.cs +++ b/src/Numerics/Providers/Common/Mkl/MklProvider.cs @@ -3,7 +3,7 @@ // http://numerics.mathdotnet.com // http://github.com/mathnet/mathnet-numerics // -// Copyright (c) 2009-2016 Math.NET +// Copyright (c) 2009-2018 Math.NET // // Permission is hereby granted, free of charge, to any person // obtaining a copy of this software and associated documentation @@ -34,17 +34,23 @@ using System.Collections.Generic; namespace MathNet.Numerics.Providers.Common.Mkl { - internal static class MklProvider + public static class MklProvider { - const int _designTimeRevision = 11; + const int _designTimeRevision = 12; static int _nativeRevision; static Version _mklVersion; static bool _nativeX86; static bool _nativeX64; static bool _nativeIA64; + static bool _loaded; internal static bool IsAvailable(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return true; + } + try { if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) @@ -65,6 +71,11 @@ namespace MathNet.Numerics.Providers.Common.Mkl internal static void Load(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return; + } + int a, b; try { @@ -101,11 +112,22 @@ namespace MathNet.Numerics.Providers.Common.Mkl throw new NotSupportedException("MKL Native Provider too old. Consider upgrading to a newer version."); } - ConfigureThreading(); + // set threading settings, if supported + if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) + { + SafeNativeMethods.set_max_threads(Control.MaxDegreeOfParallelism); + } + + _loaded = true; } - static void ConfigureThreading() + internal static void ConfigureThreading() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + // set threading settings, if supported if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) { @@ -113,8 +135,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl } } - public static void ConfigurePrecision(MklConsistency consistency, MklPrecision precision, MklAccuracy accuracy) + internal static void ConfigurePrecision(MklConsistency consistency, MklPrecision precision, MklAccuracy accuracy) { + if (!_loaded) + { + throw new InvalidOperationException(); + } + // set numerical consistency, precision and accuracy modes, if supported if (SafeNativeMethods.query_capability((int)ProviderConfig.Precision) > 0) { @@ -126,8 +153,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Frees the memory allocated to the MKL memory pool. /// - internal static void FreeBuffers() + public static void FreeBuffers() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -139,8 +171,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Frees the memory allocated to the MKL memory pool on the current thread. /// - internal static void ThreadFreeBuffers() + public static void ThreadFreeBuffers() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -152,8 +189,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Disable the MKL memory pool. May impact performance. /// - internal static void DisableMemoryPool() + public static void DisableMemoryPool() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -167,8 +209,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// On output, returns the number of memory buffers allocated. /// Returns the number of bytes allocated to all memory buffers. - internal static long MemoryStatistics(out int allocatedBuffers) + public static long MemoryStatistics(out int allocatedBuffers) { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -180,8 +227,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Enable gathering of peak memory statistics of the MKL memory pool. /// - internal static void EnablePeakMemoryStatistics() + public static void EnablePeakMemoryStatistics() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -193,8 +245,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Disable gathering of peak memory statistics of the MKL memory pool. /// - internal static void DisablePeakMemoryStatistics() + public static void DisablePeakMemoryStatistics() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -208,8 +265,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl /// /// Whether the usage counter should be reset. /// The peak number of bytes allocated to all memory buffers. - internal static long PeakMemoryStatistics(bool reset = true) + public static long PeakMemoryStatistics(bool reset = true) { + if (!_loaded) + { + throw new InvalidOperationException(); + } + if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) { throw new NotSupportedException("MKL Native Provider does not support memory management functions. Consider upgrading to a newer version."); @@ -218,8 +280,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl return SafeNativeMethods.peak_mem_usage((int)(reset ? MklMemoryRequestMode.PeakMemoryReset : MklMemoryRequestMode.PeakMemory)); } - internal static string Describe() + public static string Describe() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + var parts = new List(); if (_nativeX86) parts.Add("x86"); if (_nativeX64) parts.Add("x64"); diff --git a/src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs b/src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs index 4775543f..ee0497a0 100644 --- a/src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs +++ b/src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs @@ -41,9 +41,15 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas static bool _nativeX64; static bool _nativeIA64; static bool _nativeARM; + static bool _loaded; internal static bool IsAvailable(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return true; + } + try { if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) @@ -64,6 +70,11 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas internal static void Load(int minRevision, string hintPath) { + if (_loaded && _nativeRevision >= minRevision) + { + return; + } + int a, b; try { @@ -96,11 +107,22 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas throw new NotSupportedException("OpenBLAS Native Provider too old. Consider upgrading to a newer version."); } - ConfigureThreading(); + // set threading settings, if supported + if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) + { + SafeNativeMethods.set_max_threads(Control.MaxDegreeOfParallelism); + } + + _loaded = true; } - static void ConfigureThreading() + internal static void ConfigureThreading() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + // set threading settings, if supported if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) { @@ -108,8 +130,13 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas } } - internal static string Describe() + public static string Describe() { + if (!_loaded) + { + throw new InvalidOperationException(); + } + var parts = new List(); if (_nativeX86) parts.Add("x86"); if (_nativeX64) parts.Add("x64");