Browse Source

Native Providers: provide access to features like memory management, skip reloading if already loaded

spatial
Christoph Ruegg 9 years ago
parent
commit
b47d2191d1
  1. 22
      src/Numerics/Providers/Common/Cuda/CudaProvider.cs
  2. 95
      src/Numerics/Providers/Common/Mkl/MklProvider.cs
  3. 33
      src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs

22
src/Numerics/Providers/Common/Cuda/CudaProvider.cs

@ -34,15 +34,21 @@ using System.Collections.Generic;
namespace MathNet.Numerics.Providers.Common.Cuda namespace MathNet.Numerics.Providers.Common.Cuda
{ {
internal static class CudaProvider public static class CudaProvider
{ {
static int _nativeRevision; static int _nativeRevision;
static bool _nativeX86; static bool _nativeX86;
static bool _nativeX64; static bool _nativeX64;
static bool _nativeIA64; static bool _nativeIA64;
static bool _loaded;
internal static bool IsAvailable(int minRevision, string hintPath) internal static bool IsAvailable(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return true;
}
try try
{ {
if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath))
@ -63,6 +69,11 @@ namespace MathNet.Numerics.Providers.Common.Cuda
internal static void Load(int minRevision, string hintPath) internal static void Load(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return;
}
int a, b; int a, b;
try 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."); 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<string>(); var parts = new List<string>();
if (_nativeX86) parts.Add("x86"); if (_nativeX86) parts.Add("x86");
if (_nativeX64) parts.Add("x64"); if (_nativeX64) parts.Add("x64");

95
src/Numerics/Providers/Common/Mkl/MklProvider.cs

@ -3,7 +3,7 @@
// http://numerics.mathdotnet.com // http://numerics.mathdotnet.com
// http://github.com/mathnet/mathnet-numerics // 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 // Permission is hereby granted, free of charge, to any person
// obtaining a copy of this software and associated documentation // obtaining a copy of this software and associated documentation
@ -34,17 +34,23 @@ using System.Collections.Generic;
namespace MathNet.Numerics.Providers.Common.Mkl 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 int _nativeRevision;
static Version _mklVersion; static Version _mklVersion;
static bool _nativeX86; static bool _nativeX86;
static bool _nativeX64; static bool _nativeX64;
static bool _nativeIA64; static bool _nativeIA64;
static bool _loaded;
internal static bool IsAvailable(int minRevision, string hintPath) internal static bool IsAvailable(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return true;
}
try try
{ {
if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath))
@ -65,6 +71,11 @@ namespace MathNet.Numerics.Providers.Common.Mkl
internal static void Load(int minRevision, string hintPath) internal static void Load(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return;
}
int a, b; int a, b;
try 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."); 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 // set threading settings, if supported
if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) 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 // set numerical consistency, precision and accuracy modes, if supported
if (SafeNativeMethods.query_capability((int)ProviderConfig.Precision) > 0) if (SafeNativeMethods.query_capability((int)ProviderConfig.Precision) > 0)
{ {
@ -126,8 +153,13 @@ namespace MathNet.Numerics.Providers.Common.Mkl
/// <summary> /// <summary>
/// Frees the memory allocated to the MKL memory pool. /// Frees the memory allocated to the MKL memory pool.
/// </summary> /// </summary>
internal static void FreeBuffers() public static void FreeBuffers()
{ {
if (!_loaded)
{
throw new InvalidOperationException();
}
if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) 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."); 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
/// <summary> /// <summary>
/// Frees the memory allocated to the MKL memory pool on the current thread. /// Frees the memory allocated to the MKL memory pool on the current thread.
/// </summary> /// </summary>
internal static void ThreadFreeBuffers() public static void ThreadFreeBuffers()
{ {
if (!_loaded)
{
throw new InvalidOperationException();
}
if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) 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."); 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
/// <summary> /// <summary>
/// Disable the MKL memory pool. May impact performance. /// Disable the MKL memory pool. May impact performance.
/// </summary> /// </summary>
internal static void DisableMemoryPool() public static void DisableMemoryPool()
{ {
if (!_loaded)
{
throw new InvalidOperationException();
}
if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) 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."); 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
/// </summary> /// </summary>
/// <param name="allocatedBuffers">On output, returns the number of memory buffers allocated.</param> /// <param name="allocatedBuffers">On output, returns the number of memory buffers allocated.</param>
/// <returns>Returns the number of bytes allocated to all memory buffers.</returns> /// <returns>Returns the number of bytes allocated to all memory buffers.</returns>
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) 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."); 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
/// <summary> /// <summary>
/// Enable gathering of peak memory statistics of the MKL memory pool. /// Enable gathering of peak memory statistics of the MKL memory pool.
/// </summary> /// </summary>
internal static void EnablePeakMemoryStatistics() public static void EnablePeakMemoryStatistics()
{ {
if (!_loaded)
{
throw new InvalidOperationException();
}
if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) 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."); 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
/// <summary> /// <summary>
/// Disable gathering of peak memory statistics of the MKL memory pool. /// Disable gathering of peak memory statistics of the MKL memory pool.
/// </summary> /// </summary>
internal static void DisablePeakMemoryStatistics() public static void DisablePeakMemoryStatistics()
{ {
if (!_loaded)
{
throw new InvalidOperationException();
}
if (SafeNativeMethods.query_capability((int)ProviderConfig.Memory) < 1) 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."); 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
/// </summary> /// </summary>
/// <param name="reset">Whether the usage counter should be reset.</param> /// <param name="reset">Whether the usage counter should be reset.</param>
/// <returns>The peak number of bytes allocated to all memory buffers.</returns> /// <returns>The peak number of bytes allocated to all memory buffers.</returns>
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) 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."); 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)); 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<string>(); var parts = new List<string>();
if (_nativeX86) parts.Add("x86"); if (_nativeX86) parts.Add("x86");
if (_nativeX64) parts.Add("x64"); if (_nativeX64) parts.Add("x64");

33
src/Numerics/Providers/Common/OpenBlas/OpenBlasProvider.cs

@ -41,9 +41,15 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas
static bool _nativeX64; static bool _nativeX64;
static bool _nativeIA64; static bool _nativeIA64;
static bool _nativeARM; static bool _nativeARM;
static bool _loaded;
internal static bool IsAvailable(int minRevision, string hintPath) internal static bool IsAvailable(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return true;
}
try try
{ {
if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath)) if (!NativeProviderLoader.TryLoad(SafeNativeMethods.DllName, hintPath))
@ -64,6 +70,11 @@ namespace MathNet.Numerics.Providers.Common.OpenBlas
internal static void Load(int minRevision, string hintPath) internal static void Load(int minRevision, string hintPath)
{ {
if (_loaded && _nativeRevision >= minRevision)
{
return;
}
int a, b; int a, b;
try 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."); 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 // set threading settings, if supported
if (SafeNativeMethods.query_capability((int)ProviderConfig.Threading) > 0) 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<string>(); var parts = new List<string>();
if (_nativeX86) parts.Add("x86"); if (_nativeX86) parts.Add("x86");
if (_nativeX64) parts.Add("x64"); if (_nativeX64) parts.Add("x64");

Loading…
Cancel
Save