diff --git a/src/NativeProviders/MKL/capabilities.cpp b/src/NativeProviders/MKL/capabilities.cpp index 025bd9fa..b6675190 100644 --- a/src/NativeProviders/MKL/capabilities.cpp +++ b/src/NativeProviders/MKL/capabilities.cpp @@ -84,6 +84,10 @@ extern "C" { case 384: return 1; // basic FFT (major - breaking) case 385: return 0; // basic FFT (minor - non-breaking) + // SPARSE SOLVER + case 512: return 1; + case 513: return 0; + default: return 0; // unknown or not supported } diff --git a/src/NativeProviders/MKL/dss.c b/src/NativeProviders/MKL/dss.c new file mode 100644 index 00000000..21838774 --- /dev/null +++ b/src/NativeProviders/MKL/dss.c @@ -0,0 +1,152 @@ +#include "wrapper_common.h" +#include "dss.h" + +#if __cplusplus +extern "C" { +#endif + + // Notes: zero-based indexing is used for rowIdx[] and colPtr[]. + + DLLEXPORT dss_int s_dss_solve(const dss_int matrixStructure, const dss_int matrixType, const dss_int systemType, + const dss_int nRows, const dss_int nCols, const dss_int nnz, const dss_int rowIdx[], const dss_int colPtr[], const float values[], + const dss_int nRhs, const float rhsValues[], float solValues[]) + { + _MKL_DSS_HANDLE_t handle; + dss_int error; + + dss_int opt = MKL_DSS_MSG_LVL_WARNING + MKL_DSS_TERM_LVL_ERROR + MKL_DSS_ZERO_BASED_INDEXING + MKL_DSS_AUTO_ORDER; + opt += systemType; + + // Initialize the solver + error = dss_create(handle, opt); + if (error != MKL_DSS_SUCCESS) return error; + + // Define the non-zero structure of the matrix + error = dss_define_structure(handle, matrixStructure, rowIdx, nRows, nCols, colPtr, nnz); + if (error != MKL_DSS_SUCCESS) return error; + + // Reorder the matrix + error = dss_reorder(handle, opt, 0); + if (error != MKL_DSS_SUCCESS) return error; + + // Factor the matrix + error = dss_factor_real(handle, matrixType, values); + if (error != MKL_DSS_SUCCESS) return error; + + // Get the solution vector + error = dss_solve_real(handle, opt, rhsValues, nRhs, solValues); + if (error != MKL_DSS_SUCCESS) return error; + + // Deallocate solver storage + error = dss_delete(handle, opt); + return error; + } + + DLLEXPORT dss_int d_dss_solve(const dss_int matrixStructure, const dss_int matrixType, const dss_int systemType, + const dss_int nRows, const dss_int nCols, const dss_int nnz, const dss_int rowIdx[], const dss_int colPtr[], const double values[], + const dss_int nRhs, const double rhsValues[], double solValues[]) + { + _MKL_DSS_HANDLE_t handle; + dss_int error; + + dss_int opt = MKL_DSS_MSG_LVL_WARNING + MKL_DSS_TERM_LVL_ERROR + MKL_DSS_ZERO_BASED_INDEXING + MKL_DSS_AUTO_ORDER; + opt += systemType; + + // Initialize the solver + error = dss_create(handle, opt); + if (error != MKL_DSS_SUCCESS) return error; + + // Define the non-zero structure of the matrix + error = dss_define_structure(handle, matrixStructure, rowIdx, nRows, nCols, colPtr, nnz); + if (error != MKL_DSS_SUCCESS) return error; + + // Reorder the matrix + error = dss_reorder(handle, opt, 0); + if (error != MKL_DSS_SUCCESS) return error; + + // Factor the matrix + error = dss_factor_real(handle, matrixType, values); + if (error != MKL_DSS_SUCCESS) return error; + + // Get the solution vector + error = dss_solve_real(handle, opt, rhsValues, nRhs, solValues); + if (error != MKL_DSS_SUCCESS) return error; + + // Deallocate solver storage + error = dss_delete(handle, opt); + return error; + } + + DLLEXPORT dss_int c_dss_solve(const dss_int matrixStructure, const dss_int matrixType, const dss_int systemType, + const dss_int nRows, const dss_int nCols, const dss_int nnz, const dss_int rowIdx[], const dss_int colPtr[], const dss_complex_float values[], + const dss_int nRhs, const dss_complex_float rhsValues[], dss_complex_float solValues[]) + { + _MKL_DSS_HANDLE_t handle; + dss_int error; + + dss_int opt = MKL_DSS_MSG_LVL_WARNING + MKL_DSS_TERM_LVL_ERROR + MKL_DSS_ZERO_BASED_INDEXING + MKL_DSS_AUTO_ORDER; + opt += systemType; + + // Initialize the solver + error = dss_create(handle, opt); + if (error != MKL_DSS_SUCCESS) return error; + + // Define the non-zero structure of the matrix + error = dss_define_structure(handle, matrixStructure, rowIdx, nRows, nCols, colPtr, nnz); + if (error != MKL_DSS_SUCCESS) return error; + + // Reorder the matrix + error = dss_reorder(handle, opt, 0); + if (error != MKL_DSS_SUCCESS) return error; + + // Factor the matrix + error = dss_factor_complex(handle, matrixType, values); + if (error != MKL_DSS_SUCCESS) return error; + + // Get the solution vector + error = dss_solve_real(handle, opt, rhsValues, nRhs, solValues); + if (error != MKL_DSS_SUCCESS) return error; + + // Deallocate solver storage + error = dss_delete(handle, opt); + return error; + } + + DLLEXPORT dss_int z_dss_solve(const dss_int matrixStructure, const dss_int matrixType, const dss_int systemType, + const dss_int nRows, const dss_int nCols, const dss_int nnz, const dss_int rowIdx[], const dss_int colPtr[], const dss_complex_double values[], + const dss_int nRhs, const dss_complex_double rhsValues[], dss_complex_double solValues[]) + { + _MKL_DSS_HANDLE_t handle; + dss_int error; + + dss_int opt = MKL_DSS_MSG_LVL_WARNING + MKL_DSS_TERM_LVL_ERROR + MKL_DSS_ZERO_BASED_INDEXING + MKL_DSS_AUTO_ORDER; + opt += systemType; + + // Initialize the solver + error = dss_create(handle, opt); + if (error != MKL_DSS_SUCCESS) return error; + + // Define the non-zero structure of the matrix + error = dss_define_structure(handle, matrixStructure, rowIdx, nRows, nCols, colPtr, nnz); + if (error != MKL_DSS_SUCCESS) return error; + + // Reorder the matrix + error = dss_reorder(handle, opt, 0); + if (error != MKL_DSS_SUCCESS) return error; + + // Factor the matrix + error = dss_factor_complex(handle, matrixType, values); + if (error != MKL_DSS_SUCCESS) return error; + + // Get the solution vector + error = dss_solve_real(handle, opt, rhsValues, nRhs, solValues); + if (error != MKL_DSS_SUCCESS) return error; + + // Deallocate solver storage + error = dss_delete(handle, opt); + return error; + } + +#if __cplusplus +} +#endif diff --git a/src/NativeProviders/MKL/dss.h b/src/NativeProviders/MKL/dss.h new file mode 100644 index 00000000..e7fd398a --- /dev/null +++ b/src/NativeProviders/MKL/dss.h @@ -0,0 +1,9 @@ +#pragma once + +#include "mkl_dss.h" +#include "mkl_types.h" + +#define dss_int MKL_INT +#define dss_complex_float MKL_Complex8 +#define dss_complex_double MKL_Complex16 + diff --git a/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj b/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj index 12fbbca2..08a201bf 100644 --- a/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj +++ b/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj @@ -194,6 +194,7 @@ + @@ -204,10 +205,11 @@ + - \ No newline at end of file + diff --git a/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj.filters b/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj.filters index 7e91fc28..ac9baa12 100644 --- a/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj.filters +++ b/src/NativeProviders/Windows/MKL/MKLWrapper.vcxproj.filters @@ -36,6 +36,9 @@ Source Files + + Source Files + @@ -55,6 +58,9 @@ Header Files + + Header Files + diff --git a/src/Numerics.Tests/LinearAlgebraTests/MatrixStructureTheory.cs b/src/Numerics.Tests/LinearAlgebraTests/MatrixStructureTheory.cs index 2cc08b2f..6a32c910 100644 --- a/src/Numerics.Tests/LinearAlgebraTests/MatrixStructureTheory.cs +++ b/src/Numerics.Tests/LinearAlgebraTests/MatrixStructureTheory.cs @@ -30,6 +30,7 @@ using System; using System.Linq; using MathNet.Numerics.LinearAlgebra; +using MathNet.Numerics.LinearAlgebra.Storage; using NUnit.Framework; namespace MathNet.Numerics.UnitTests.LinearAlgebraTests @@ -580,6 +581,131 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests Assert.That(matrix[i, j], Is.EqualTo(rows[i][j])); } + [Test] + public void CanCreateSparseFromCoordinateFormat() + { + var rows = new[] + { + Vector.Build.Random(4, 0), + Vector.Build.Random(4, 1), + Vector.Build.Random(4, 3) + }; + + var rowCount = rows.Length; + var columnCount = 4; + var valueCount = rowCount * columnCount; + + var cooRowIndices = new int[valueCount]; + var cooColumnIndices = new int[valueCount]; + var cooValues = new T[valueCount]; + + int loc = 0; + for (int i = 0; i < rowCount; i++) + { + for (int j = 0; j < columnCount; j++) + { + cooRowIndices[loc] = i; + cooColumnIndices[loc] = j; + cooValues[loc] = rows[i].At(j); + loc++; + } + } + + var matrix = Matrix.Build.SparseFromCoordinateFormat(rowCount, columnCount, valueCount, cooRowIndices, cooColumnIndices, cooValues); + Assert.That(matrix.GetType().Name, Is.EqualTo("SparseMatrix")); + Assert.That(matrix.RowCount, Is.EqualTo(3)); + Assert.That(matrix.ColumnCount, Is.EqualTo(4)); + for (int j = 0; j < 4; j++) + for (int i = 0; i < 3; i++) + Assert.That(matrix[i, j], Is.EqualTo(rows[i][j])); + } + + [Test] + public void CanCreateSparseFromCompressedSparseRowFormat() + { + var rows = new[] + { + Vector.Build.Random(4, 0), + Vector.Build.Random(4, 1), + Vector.Build.Random(4, 3) + }; + + var rowCount = rows.Length; + var columnCount = 4; + var valueCount = rowCount * columnCount; + + var csrRowPointers = new int[rowCount + 1]; + var csrColumnIndices = new int[valueCount]; + var csrValues = new T[valueCount]; + + int loc = 0; + for (int i = 0; i < rowCount; i++) + { + for (int j = 0; j < columnCount; j++) + { + csrRowPointers[i + 1]++; + csrColumnIndices[loc] = j; + csrValues[loc] = rows[i].At(j); + loc++; + } + } + for (int i = 1; i < rowCount + 1; i++) + { + csrRowPointers[i] += csrRowPointers[i - 1]; + } + + var matrix = Matrix.Build.SparseFromCompressedSparseRowFormat(rowCount, columnCount, valueCount, csrRowPointers, csrColumnIndices, csrValues); + Assert.That(matrix.GetType().Name, Is.EqualTo("SparseMatrix")); + Assert.That(matrix.RowCount, Is.EqualTo(3)); + Assert.That(matrix.ColumnCount, Is.EqualTo(4)); + for (int j = 0; j < 4; j++) + for (int i = 0; i < 3; i++) + Assert.That(matrix[i, j], Is.EqualTo(rows[i][j])); + } + + [Test] + public void CanCreateSparseFromCompressedSparseColumnFormat() + { + var rows = new[] + { + Vector.Build.Random(4, 0), + Vector.Build.Random(4, 1), + Vector.Build.Random(4, 3) + }; + + var rowCount = rows.Length; + var columnCount = 4; + var valueCount = rowCount * columnCount; + + var cscRowIndices = new int[valueCount]; + var cscColumnPointers = new int[columnCount + 1]; + var cscValues = new T[valueCount]; + + int loc = 0; + for (int j = 0; j < columnCount; j++) + { + for (int i = 0; i < rowCount; i++) + { + cscColumnPointers[j + 1]++; + cscRowIndices[loc] = i; + cscValues[loc] = rows[i].At(j); + loc++; + } + } + for (int i = 1; i < columnCount + 1; i++) + { + cscColumnPointers[i] += cscColumnPointers[i - 1]; + } + + var matrix = Matrix.Build.SparseFromCompressedSparseColumnFormat(rowCount, columnCount, valueCount, cscRowIndices, cscColumnPointers, cscValues); + Assert.That(matrix.GetType().Name, Is.EqualTo("SparseMatrix")); + Assert.That(matrix.RowCount, Is.EqualTo(3)); + Assert.That(matrix.ColumnCount, Is.EqualTo(4)); + for (int j = 0; j < 4; j++) + for (int i = 0; i < 3; i++) + Assert.That(matrix[i, j], Is.EqualTo(rows[i][j])); + } + [Test] public void CanEnumerateWithIndex() { diff --git a/src/Numerics.Tests/Providers/SparseSolver/double/SparseSolverProviderTests.cs b/src/Numerics.Tests/Providers/SparseSolver/double/SparseSolverProviderTests.cs new file mode 100644 index 00000000..0bc86a34 --- /dev/null +++ b/src/Numerics.Tests/Providers/SparseSolver/double/SparseSolverProviderTests.cs @@ -0,0 +1,585 @@ +using MathNet.Numerics.LinearAlgebra; +using MathNet.Numerics.LinearAlgebra.Double; +using MathNet.Numerics.LinearAlgebra.Storage; +using MathNet.Numerics.Providers.SparseSolver; +using NUnit.Framework; +using System; +using System.Collections.Generic; + +namespace MathNet.Numerics.UnitTests.Providers.SparseSolver.Double +{ + +#if NATIVE +#if MKL + + /// + /// Base class for sparse solver provider tests. + /// + [TestFixture, Category("SparseSolverProvider")] + public class SparseSolverProviderTests + { + readonly double[] _b4 = { 1.0, 2.0, 3.0, 4.0}; + readonly double[] _b5 = { 1.0, 2.0, 3.0, 4.0, 5.0 }; + + /// + /// Test matrix to use. + /// + readonly IDictionary _matrices = new Dictionary + { + {"SymmetricPositiveDefinite5x5", (SparseMatrix)Matrix.Build.SparseOfColumnArrays(new [] {9.0, 1.5, 6.0, 0.75, 3.0}, new [] {1.5, 0.5, 0.0, 0.0, 0.0}, new [] {6.0, 0.0, 12.0, 0.0, 0.0 }, new [] { 0.75, 0.0, 0.0, 0.625, 0.0}, new [] {3.0, 0.0, 0.0, 0.0, 16.0})}, + {"Triangle5x5", (SparseMatrix)Matrix.Build.SparseOfColumnArrays(new [] {1.0, 0.0, 0.0, 0.0, 0.0}, new [] {5.0, 2.0, 0.0, 0.0, 0.0}, new [] {0.0, 8.0, 3.0, 0.0, 0.0 }, new [] { 0.0, 0.0, 9.0, 4.0, 0.0}, new [] {0.0, 0.0, 0.0, 10.0, 5.0})}, + {"Square4x4", (SparseMatrix)Matrix.Build.SparseOfColumnArrays(new [] {1.0, 1.0, 1.0, 2.0 },new [] {2.0, 0.0, 0.0, 2.0 },new [] {0.0, 0.0, 2.0, 1.0 },new [] {4.0, 1.0, 1.0, 0.0 })}, + }; + + /// + /// Can solve Ax=b using direct sparse solver. + /// + [Test] + public void CanSolveSymmetricPositiveDefiniteMatrix() + { + var A = _matrices["SymmetricPositiveDefinite5x5"].UpperTriangle(); + + var csr = A.Storage as SparseCompressedRowMatrixStorage; + csr.PopulateExplicitZerosOnDiagonal(); + + var rowCount = csr.RowCount; + var columnCount = csr.ColumnCount; + var valueCount = csr.ValueCount; + var values = csr.Values; + var rowPointers = csr.RowPointers; + var columnIndices = csr.ColumnIndices; + + var xactual = new double[rowCount]; + + var error = SparseSolverControl.Provider.Solve(DssMatrixStructure.Symmetric, DssMatrixType.PositiveDefinite, DssSystemType.DontTranspose, + rowCount, columnCount, valueCount, rowPointers, columnIndices, values, + 1, _b5, xactual); + + Assert.That(error, Is.EqualTo(DssStatus.MKL_DSS_SUCCESS)); + + var xtrue = new double[] { -979.0 / 3.0, 983.0, 1961.0 / 12.0, 398.0, 123.0 / 2.0 }; + + for (int i = 0; i < xtrue.Length; i++) + AssertHelpers.AlmostEqualRelative(xtrue[i], xactual[i], 12); + } + + /// + /// Can solve Ax=b using direct sparse solver. + /// + [Test] + public void CanSolveUpperTriangularMatrix() + { + var A = _matrices["Triangle5x5"].UpperTriangle(); + + var csr = A.Storage as SparseCompressedRowMatrixStorage; + csr.PopulateExplicitZerosOnDiagonal(); + + var rowCount = csr.RowCount; + var columnCount = csr.ColumnCount; + var valueCount = csr.ValueCount; + var values = csr.Values; + var rowPointers = csr.RowPointers; + var columnIndices = csr.ColumnIndices; + + var xactual = new double[rowCount]; + + var error = SparseSolverControl.Provider.Solve(DssMatrixStructure.Nonsymmetric, DssMatrixType.Indefinite, DssSystemType.DontTranspose, + rowCount, columnCount, valueCount, rowPointers, columnIndices, values, + 1, _b5, xactual); + + Assert.That(error, Is.EqualTo(DssStatus.MKL_DSS_SUCCESS)); + + var xtrue = new double[] { 106.0, -21.0, 5.5, -1.5, 1.0 }; + + for (int i = 0; i < xtrue.Length; i++) + AssertHelpers.AlmostEqualRelative(xtrue[i], xactual[i], 13); + } + + /// + /// Can solve Ax=b using direct sparse solver. + /// + [Test] + public void CanSolveSquareMatrix() + { + var A = _matrices["Square4x4"]; + + var csr = A.Storage as SparseCompressedRowMatrixStorage; + csr.PopulateExplicitZerosOnDiagonal(); + + var rowCount = csr.RowCount; + var columnCount = csr.ColumnCount; + var valueCount = csr.ValueCount; + var values = csr.Values; + var rowPointers = csr.RowPointers; + var columnIndices = csr.ColumnIndices; + + var xactual = new double[rowCount]; + + var error = SparseSolverControl.Provider.Solve(DssMatrixStructure.Nonsymmetric, DssMatrixType.Indefinite, DssSystemType.DontTranspose, + rowCount, columnCount, valueCount, rowPointers, columnIndices, values, + 1, _b4, xactual); + + Assert.That(error, Is.EqualTo(DssStatus.MKL_DSS_SUCCESS)); + + var xtrue = new double[] { 2.1, -0.35, 0.5, -0.1 }; + + for (int i = 0; i < xtrue.Length; i++) + AssertHelpers.AlmostEqualRelative(xtrue[i], xactual[i], 10); + } + + /// + /// Can inverse A by using AX = I. + /// + [Test] + public void CanInverseSquareMatrix() + { + var A = _matrices["SymmetricPositiveDefinite5x5"]; + var Atr = A.UpperTriangle(); + + var csr = Atr.Storage as SparseCompressedRowMatrixStorage; + csr.PopulateExplicitZerosOnDiagonal(); + + var rowCount = csr.RowCount; + var columnCount = csr.ColumnCount; + var valueCount = csr.ValueCount; + var values = csr.Values; + var rowPointers = csr.RowPointers; + var columnIndices = csr.ColumnIndices; + + var Identity = Matrix.Build.DenseIdentity(columnCount); + var b = Identity.ToColumnMajorArray(); + var Xactual = new double[rowCount * columnCount]; + + var error = SparseSolverControl.Provider.Solve(DssMatrixStructure.Symmetric, DssMatrixType.PositiveDefinite, DssSystemType.DontTranspose, + rowCount, columnCount, valueCount, rowPointers, columnIndices, values, + Identity.ColumnCount, b, Xactual); + + Assert.That(error, Is.EqualTo(DssStatus.MKL_DSS_SUCCESS)); + + var Ainverse_actual = Matrix.Build.SparseOfColumnMajor(rowCount, columnCount, Xactual); + + var Ainverse_expected = Matrix.Build.DenseOfColumnArrays( + new[] { 80.0/3.0, -80.0, -40.0/3.0, -32.0, -5.0 }, + new[] { -80.0, 242.0, 40.0, 96.0, 15.0 }, + new[] { -40.0/3.0, 40.0, 6.75, 16.0, 2.5 }, + new[] { -32.0, 96.0, 16.0, 40.0, 6.0 }, + new[] { -5.0, 15.0, 2.5, 6.0, 1.0}); + + for (int i = 0; i < Ainverse_actual.RowCount; i++) + for (int j = 0; j < Ainverse_actual.ColumnCount; j++) + AssertHelpers.AlmostEqualRelative(Ainverse_actual[i, j], Ainverse_expected[i, j], 10); + } + + /// + /// Can solve 1D boundary problem. + /// + [TestCase(4, 8001, 10, 1000)] + [TestCase(40, 8001, 0.1, 100)] + [TestCase(400, 8001, 0.001, 10)] + [TestCase(4000, 8001, 0.00001, 1)] + public void CanSolvePoissonEquation(int elementCount, int interpolationCount, double errV, double errE) + { + var domain = new Domain(elementCount, 0.08); + domain.DefineBoundaries(1.0, 0.0); + domain.DefineProblem(); + domain.SolveProblem(); + + var xGrid = Generate.LinearSpaced(interpolationCount, 0.0, domain.Length); + + var interpolated = domain.InterpolateAt(xGrid); + var Vactual = interpolated.Item1; // interpolated electric potential + var Eactual = interpolated.Item2; // interpolated electric field + + var exact = domain.GetExactSolution(xGrid); + var Vexpected = exact.Item1; // expected electric potential + var Eexpected = exact.Item2; // expected electric field + + var Vdiff = Vector.Build.Dense(Vactual.Length, (i) => Vexpected[i] - Vactual[i]); + var Vnorm2 = Vdiff.L2Norm(); + Assert.LessOrEqual(Vnorm2, errV); + + var Ediff = Vector.Build.Dense(Eactual.Length, (i) => Eexpected[i] - Eactual[i]); + var Enorm2 = Ediff.L2Norm(); + Assert.LessOrEqual(Enorm2, errE); + } + + #region Finite element method to solve Poisson's equation + + class Node + { + public int ID; + public double X; + public double PrimaryValue; // electric potential at the node + public double SecondaryValue; // electric field at the node + + public Node(double x) + { + X = x; + } + } + + class Element // Linear element + { + public int ID; + public Node[] Nodes; + public double Alpha; + public double Gamma; + public Matrix Kmatrix; + public Vector Rhs; + + public Element(Node node1, Node node2, double alpha, double gamma) + { + Nodes = new[] { node1, node2 }; + Alpha = alpha; + Gamma = gamma; + } + + public void ComputeMatrices() + { + // The master equation: + // ∇(α∇V) + γ = 0 + // + // Let's define a transformation from u to x + // x = a * u + b + // + // for Node1, u = -1 gives -a + b = x1 + // for Node2, u = +1 gives a + b = x2 + // + // So, + // x = (x2 - x1)/2*u + (x2 + x1)/2 + // u = 2*(x - x1)/(x2 - x1) - 1 + // + // This gives + // dx = l/2*du where l = x2 - x1 = x21 is length of the line + // + // Interpolation function, Ni(u) = ai + bi*u for i = 1, 2 + // N1(-1) = a1 - b1 = 1 + // N1(+1) = a1 + b1 = 0 -> a1 = 1/2, b1 = -1/2 + // N2(-1) = a2 - b2 = 0 + // N2(+1) = a2 + b2 = 1 -> a2 = 1/2, b2 = 1/2 + // so, + // N1 = (1 - u) / 2 + // N2 = (1 + u) / 2 + // + // Using the chain rule of differentiation, + // ∂Ni(x)/∂u = ∂Ni/∂x ∂x/∂u + // + // Here, + // J = ∂x/∂u = x21/2 = l/2 + // + // |J| = l/2 where l is the length of the line + // + // Therefore, we can get + // ∂N1/∂x = 2/l ∂N1/∂u = - 1/l + // ∂N2/∂x = 2/l ∂N2/∂u = 1/l + //----------------------------------------------------------------------------- + // By using V(u) = Σ Vi*Ni(u) and ω = Ni for i = 1, 2 + // the weak form of the master equation is given as + // + // ∫ ω [d/dx(α dV/dx) + γ] dx = 0 + // + // K v = f + p + // + // where + // Kij = ∫ (dNi/dx) α (dNj/dx) dx + // fi = ∫ Ni γ dx + // pi = -Ni(x2)D(x2) + Ni(x1)D(x1) where D(x) = -α dV(x)/dx + //----------------------------------------------------------------------------- + // Kij = (l/2) ∫ α (∂Ni/∂x) (∂Nj/∂x) du, where α = const. + // = α*(∂Ni/∂x)*(∂Nj/∂x)*l + // K11 = α/l + // K12 = M21 = - α/l + // K22 = α/l + //----------------------------------------------------------------------------- + // fi = γ*(l/2) ∫ Ni du, where γ = const. + // = γ*(l/2) + // f1 = γ*l/2 + // f2 = γ*l/2 + //----------------------------------------------------------------------------- + // p1 = -N1(x2)D(x2) + N1(x1)D(x1) = D(x1) = D1 + // p2 = -N2(x2)D(x2) + N2(x1)D(x1) = -D(x2) = -D2 + // For a sufficiently large number of elements in the domain, we can ignore p. + + var x21 = Nodes[1].X - Nodes[0].X; // element length + + Kmatrix = Matrix.Build.Dense(2, 2); + Kmatrix[0, 0] = Alpha / x21; + Kmatrix[0, 1] = -Alpha / x21; + Kmatrix[1, 0] = -Alpha / x21; + Kmatrix[1, 1] = Alpha / x21; + + Rhs = Vector.Build.Dense(2); + Rhs[0] = Gamma * x21 / 2.0; + Rhs[1] = Gamma * x21 / 2.0; + } + + public bool Contains(Node point) + { + // A point P can be described with a line A-B + // P = A + s*(B - A) + // s = PA/BA + // if (s >= 0 && s <= 1), then P is inside the line A-B. + + var s = (point.X - Nodes[0].X) / (Nodes[1].X - Nodes[0].X); + return s >= 0.0 && (1.0 - s) >= 0.0; + } + + public void InterpolateAt(Node point) + { + // V(u) can be described with interpolation functions + // V(u) = V1*N1(u) + V2*N2(u) + // where + // N1 = (1 - u) / 2 + // N2 = (1 + u) / 2 + // The transformation from x to u, + // u = 2*(x - x1)/(x2 - x1) - 1 + // gives + // V(x) = V1*(x2 - x)/(x2 - x1) + V2*(x - x1)/(x2 - x1) + + var V1 = Nodes[0].PrimaryValue; + var V2 = Nodes[1].PrimaryValue; + + var x21 = Nodes[1].X - Nodes[0].X; + var xp1 = point.X - Nodes[0].X; + var u = 2.0 * xp1 / x21 - 1.0; + + var Vx = V1 * (1 - u) * 0.5 + V2 * (1 + u) * 0.5; + + // Electric field, + // E = -∇V + // where + // ∇V = [ (∂/∂x) ∑ViNi ] + // = [ (∂/∂x)(V1*N1 + V2*N2) ] + // = [ V1*(∂N1/∂x) + V2*(∂N2/∂x) ] + // + // The derivatives of Ni are + // ∂N1/∂x = 2/l ∂N1/∂u = - 1/l where l = x21 + // ∂N2/∂x = 2/l ∂N2/∂u = 1/l + // + // Therefore, + // E = (V1 - V2) / x21 + + var Ex = (V1 - V2) / x21; + + point.PrimaryValue = Vx; + point.SecondaryValue = Ex; + } + } + + class Domain + { + // Length of the domain in [m] + public double Length; + + // Dielectric constant of the domain + public double Permittivity; + + // Charge density in [C/m^3] + public double ChargeDensity; + + // Boundary conditions + public double VoltageAtLeft; + public double VoltageAtRight; + public List> Boundaries; + + public Node[] Nodes; + public Element[] Elements; + + public Matrix Kmatrix; + public Vector Rhs; + + public Domain(int elementCount = 4, double length = 0.08, double relativePermittivity = 1, double chargeDensity = 1E-8) + { + Length = length; + + Permittivity = Constants.ElectricPermittivity * relativePermittivity; + ChargeDensity = chargeDensity; + + // Create nodes and elements + Nodes = new Node[elementCount + 1]; + Elements = new Element[elementCount]; + + var dx = length / (double)elementCount; + + // nodes: 0 1 2 3 ... n+1 + // elements: 0 1 2 ... n + // +---+---+---+ ... ---+ + // x axis: 0 length + + for (int i = 0; i < Nodes.Length; i++) + { + Nodes[i] = new Node(dx * i) { ID = i }; + } + for (int i = 0; i < Elements.Length; i++) + { + Elements[i] = new Element(Nodes[i], Nodes[i + 1], Permittivity, -ChargeDensity) { ID = i }; + } + + // Initialization of the global K matrix and right-hand side vector + // We know the K matrix is a symmetric matrix, so we will only handle the upper triangular parts. + Kmatrix = Matrix.Build.Sparse(Nodes.Length, Nodes.Length); + Rhs = Vector.Build.Dense(Nodes.Length); + } + + public void DefineBoundaries(double Vleft, double Vright) + { + VoltageAtLeft = Vleft; + VoltageAtRight = Vright; + + // Dirichlet boundary conditions + Boundaries = new List>(); + foreach (var node in Nodes) + { + if (node.X == 0d) + { + Boundaries.Add(new Tuple(node, Vleft)); + } + else if (node.X == Length) + { + Boundaries.Add(new Tuple(node, Vright)); + } + } + } + + public void DefineProblem() + { + // Form the element matrices and assemble to the global matrix + foreach (var element in Elements) + { + element.ComputeMatrices(); + + // Assemble element matrix into the global K matrix + for (int i = 0; i < element.Nodes.Length; i++) + { + var row = element.Nodes[i].ID; + for (int j = 0; j < element.Nodes.Length; j++) + { + var col = element.Nodes[j].ID; + if (row <= col) // only upper triangular parts are handled + { + Kmatrix[row, col] += element.Kmatrix[i, j]; + } + } + Rhs[row] += element.Rhs[i]; + } + } + + // Imposition of Dirichlet boundary conditions + // + // If xn is given, i.e. x2 = b0, then the linear equations, + // [ K11 K12 K13 K14 ][ x1 ] = [ b1 ] + // [ K21 K22 K23 K24 ][ x2 ] [ b2 ] + // [ K31 K32 K33 K34 ][ x3 ] [ b3 ] + // [ K41 K42 K43 K44 ][ x4 ] [ b4 ] + // can be changed to + // [ K11 0 K13 K14 ][ x1 ] = [ b1 - K12*b0 ] + // [ 0 1 0 0 ][ x2 ] [ b0 ] + // [ K31 0 K33 K34 ][ x3 ] [ b3 - K32*b0 ] + // [ K41 0 K43 K44 ][ x4 ] [ b4 - K42*b0 ] + + foreach (var boundary in Boundaries) + { + var node = boundary.Item1; + var i = node.ID; + var val = boundary.Item2; + + for (int j = 0; j < Nodes.Length; j++) + { + if (Nodes[j].ID != i) + { + Rhs[j] -= (j <= i) + ? Kmatrix[j, i] * val + : Kmatrix[i, j] * val; // Kmatrix has only upper triangular parts + } + } + + var storage = Kmatrix.Storage as SparseCompressedRowMatrixStorage; + storage.MapIndexedInplace( + (row, col, x) => (row == i && col == i) + ? 1d + : (row == i) || (col == i) + ? 0d + : x, + Zeros.AllowSkip); + + Rhs[i] = val; + } + } + + public void SolveProblem() + { + // Note that Kmatrix is actually an upper triangular, but considered as a symmetric. + var storage = Kmatrix.Storage as SparseCompressedRowMatrixStorage; + var rowCount = storage.RowCount; + var columnCount = storage.ColumnCount; + var valueCount = storage.ValueCount; + var values = storage.Values; + var rowPointers = storage.RowPointers; + var columnIndices = storage.ColumnIndices; + + var rhs = Rhs.ToArray(); + var solution = new double[rowCount]; + + SparseSolverControl.Provider.Solve(DssMatrixStructure.Symmetric, DssMatrixType.Indefinite, DssSystemType.DontTranspose, + rowCount, columnCount, valueCount, rowPointers, columnIndices, values, + 1, rhs, solution); + + for (int i = 0; i < solution.Length; i++) + { + Nodes[i].PrimaryValue = solution[i]; + } + } + + public Tuple InterpolateAt(double[] xGrid) + { + var Vactual = new double[xGrid.Length]; + var Eactual = new double[xGrid.Length]; + + for (int i = 0; i < xGrid.Length; i++) + { + var point = new Node(xGrid[i]); + for (int j = 0; j < Elements.Length; j++) + { + if (Elements[j].Contains(point)) + { + Elements[j].InterpolateAt(point); + Vactual[i] = point.PrimaryValue; + Eactual[i] = point.SecondaryValue; + break; + } + } + } + + return new Tuple(Vactual, Eactual); + } + + public Tuple GetExactSolution(double[] xGrid) + { + // Poisson's equation: ∇(ε∇V) = ρ + // Solution: + // V(x) = ρ/ε/2*x^2 - (ρ/ε/2*d + (Va - Vb)/d)*x + Va, where d = length + // E(x) = -∇V = ρ/ε*x - (ρ/ε/2*d + (Va - Vb)/d) + + double[] Vexact = new double[xGrid.Length]; // electric potentials + double[] Eexact = new double[xGrid.Length]; // electric fields + var factor = ChargeDensity / Permittivity * 0.5; // ρ/ε/2 + + for (int i = 0; i < xGrid.Length; i++) + { + var x = xGrid[i]; + Vexact[i] = factor * x * x - factor * Length * x - (VoltageAtLeft - VoltageAtRight) / Length * x + VoltageAtLeft; + Eexact[i] = -2d * factor * x + factor * Length + (VoltageAtLeft - VoltageAtRight) / Length; + } + + return new Tuple(Vexact, Eexact); + } + } + + #endregion + } + +#endif +#endif + +} + diff --git a/src/Numerics/Control.cs b/src/Numerics/Control.cs index 3028af03..9c8f1b35 100644 --- a/src/Numerics/Control.cs +++ b/src/Numerics/Control.cs @@ -33,6 +33,7 @@ using System.Reflection; using System.Runtime.InteropServices; using System.Text; using System.Threading.Tasks; +using MathNet.Numerics.Providers.SparseSolver; using MathNet.Numerics.Providers.FourierTransform; using MathNet.Numerics.Providers.LinearAlgebra; @@ -70,12 +71,14 @@ namespace MathNet.Numerics { LinearAlgebraControl.UseManaged(); FourierTransformControl.UseManaged(); + SparseSolverControl.UseManaged(); } public static void UseManagedReference() { LinearAlgebraControl.UseManagedReference(); FourierTransformControl.UseManaged(); + SparseSolverControl.UseManaged(); } /// @@ -86,6 +89,7 @@ namespace MathNet.Numerics { LinearAlgebraControl.UseDefault(); FourierTransformControl.UseDefault(); + SparseSolverControl.UseDefault(); } /// @@ -95,6 +99,7 @@ namespace MathNet.Numerics { LinearAlgebraControl.UseBest(); FourierTransformControl.UseBest(); + SparseSolverControl.UseBest(); } #if NATIVE @@ -107,6 +112,7 @@ namespace MathNet.Numerics { LinearAlgebraControl.UseNativeMKL(); FourierTransformControl.UseNativeMKL(); + SparseSolverControl.UseNativeMKL(); } /// @@ -121,6 +127,7 @@ namespace MathNet.Numerics { LinearAlgebraControl.UseNativeMKL(consistency, precision, accuracy); FourierTransformControl.UseNativeMKL(); + SparseSolverControl.UseNativeMKL(); } /// @@ -134,7 +141,8 @@ namespace MathNet.Numerics { bool linearAlgebra = LinearAlgebraControl.TryUseNativeMKL(); bool fourierTransform = FourierTransformControl.TryUseNativeMKL(); - return linearAlgebra || fourierTransform; + bool directSparseSolver = SparseSolverControl.TryUseNativeMKL(); + return linearAlgebra || fourierTransform || directSparseSolver; } /// @@ -200,6 +208,7 @@ namespace MathNet.Numerics { LinearAlgebraControl.FreeResources(); FourierTransformControl.FreeResources(); + SparseSolverControl.FreeResources(); } public static void UseSingleThread() @@ -209,6 +218,7 @@ namespace MathNet.Numerics LinearAlgebraControl.Provider.InitializeVerify(); FourierTransformControl.Provider.InitializeVerify(); + SparseSolverControl.Provider.InitializeVerify(); } public static void UseMultiThreading() @@ -218,6 +228,7 @@ namespace MathNet.Numerics LinearAlgebraControl.Provider.InitializeVerify(); FourierTransformControl.Provider.InitializeVerify(); + SparseSolverControl.Provider.InitializeVerify(); } /// @@ -247,6 +258,7 @@ namespace MathNet.Numerics _nativeProviderHintPath = value; LinearAlgebraControl.HintPath = value; FourierTransformControl.HintPath = value; + SparseSolverControl.HintPath = value; } } @@ -265,6 +277,7 @@ namespace MathNet.Numerics // Reinitialize providers: LinearAlgebraControl.Provider.InitializeVerify(); FourierTransformControl.Provider.InitializeVerify(); + SparseSolverControl.Provider.InitializeVerify(); } } @@ -323,6 +336,7 @@ namespace MathNet.Numerics #endif sb.AppendLine($"Linear Algebra Provider: {LinearAlgebraControl.Provider}"); sb.AppendLine($"Fourier Transform Provider: {FourierTransformControl.Provider}"); + sb.AppendLine($"Sparse Solver Provider: {SparseSolverControl.Provider}"); sb.AppendLine($"Max Degree of Parallelism: {MaxDegreeOfParallelism}"); sb.AppendLine($"Parallelize Elements: {ParallelizeElements}"); sb.AppendLine($"Parallelize Order: {ParallelizeOrder}"); diff --git a/src/Numerics/LinearAlgebra/Builder.cs b/src/Numerics/LinearAlgebra/Builder.cs index b5a58900..e0ea1a46 100644 --- a/src/Numerics/LinearAlgebra/Builder.cs +++ b/src/Numerics/LinearAlgebra/Builder.cs @@ -33,6 +33,7 @@ using System.Linq; using MathNet.Numerics.Distributions; using MathNet.Numerics.LinearAlgebra.Solvers; using MathNet.Numerics.LinearAlgebra.Storage; +using MathNet.Numerics.Properties; using MathNet.Numerics.Random; namespace MathNet.Numerics.LinearAlgebra.Double @@ -1128,6 +1129,174 @@ namespace MathNet.Numerics.LinearAlgebra return m; } + + // Representation of Sparse Matrix + // + // Matrix A = [ 0 b 0 h 0 0 ] + // [ a c e i 0 0 ] + // [ 0 0 f j l n ] + // [ 0 d g k m 0 ] + // + // Rows = 4, Columns = 6, NonZeroCount = 14 + // + // (1) COO, Coordinate Format: + // cooRowIndices = { 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3 } + // cooColumnIndices = { 1, 3, 0, 1, 2, 3, 2, 3, 4, 5, 1, 2, 3, 4 } + // cooValues = { b, h, a, c, e, i, f, j, l, n, d, g, k, m } + // + // (2) CSR, Compressed Sparse Row representation: + // csrRowPointers = { 0, 2, 6, 10, 14 } + // csrColumnIndices = { 1, 3, 0, 1, 2, 3, 2, 3, 4, 5, 1, 2, 3, 4 } + // csrValues = { b, h, a, c, e, i, f, j, l, n, d, g, k, m } + // + // (3) CSC, Compressed Sparse Column representation: + // csrColumnPointers = { 0, 1, 4, 7, 11, 13, 14 } + // csrRowIndices = { 1, 0, 1, 3, 1, 2, 3, 0, 1, 2, 3, 2, 3, 2 } + // csrValues = { a, b, c, d, e, f, g, h, i, j, k, l, m, n } + + + /// + /// Create a new sparse matrix from a coordinate format. + /// This new matrix will be independent from the given arrays. + /// A new memory block will be allocated for storing the matrix. + /// + public Matrix SparseFromCoordinateFormat(int rows, int columns, int nonZeroCount, int[] cooRowIndices, int[] cooColumnIndices, T[] cooValues) + { + if (cooValues == null) + throw new NullReferenceException(nameof(cooValues)); + if (cooRowIndices == null) + throw new NullReferenceException(nameof(cooRowIndices)); + if (cooColumnIndices == null) + throw new NullReferenceException(nameof(cooColumnIndices)); + + if (cooRowIndices.Length < nonZeroCount || cooColumnIndices.Length < nonZeroCount || cooValues.Length < nonZeroCount) + { + var message = string.Format(Resources.ArgumentArrayWrongLength, nonZeroCount); + throw new Exception(message); + } + + // convert from COO to CSR + + var csrValues = new T[nonZeroCount]; + var csrColumnIndices = new int[nonZeroCount]; + var csrRowPointers = new int[rows + 1]; + + for (int i = 0; i < nonZeroCount; i++) + { + csrRowPointers[cooRowIndices[i] + 1]++; + } + for (int i = 1; i < rows + 1; i++) + { + csrRowPointers[i] += csrRowPointers[i - 1]; + } + var curr = new int[rows]; + for (int i = 0; i < nonZeroCount; i++) + { + int row = cooRowIndices[i]; + var loc = csrRowPointers[row] + curr[row]; + curr[row]++; + + csrColumnIndices[loc] = cooColumnIndices[i]; + csrValues[loc] = cooValues[i]; + } + + var storage = new SparseCompressedRowMatrixStorage(rows, columns, csrRowPointers, csrColumnIndices, csrValues); + return Sparse(storage); + } + + + /// + /// Create a new sparse matrix from a compressed sparse row format. + /// This new matrix will be independent from the given arrays. + /// A new memory block will be allocated for storing the matrix. + /// + public Matrix SparseFromCompressedSparseRowFormat(int rows, int columns, int nonZeroCount, int[] csrRowPointers, int[] csrColumnIndices, T[] csrValues) + { + if (csrValues == null) + throw new NullReferenceException(nameof(csrValues)); + if (csrColumnIndices == null) + throw new NullReferenceException(nameof(csrColumnIndices)); + if (csrRowPointers == null) + throw new NullReferenceException(nameof(csrRowPointers)); + if (csrRowPointers.Length < rows) + { + var message = string.Format(Resources.ArgumentArrayWrongLength, rows + 1); + throw new Exception(message); + } + if (nonZeroCount != csrRowPointers[rows]) + { + var message = string.Format("{0} should be same to {1}", nameof(nonZeroCount), csrRowPointers[rows]); + throw new Exception(message); + } + + var values = new T[nonZeroCount]; + Array.Copy(csrValues, values, nonZeroCount); + var columnIndices = new int[nonZeroCount]; + Array.Copy(csrColumnIndices, columnIndices, nonZeroCount); + var rowPointers = new int[rows + 1]; + Array.Copy(csrRowPointers, rowPointers, rows + 1); + + var storage = new SparseCompressedRowMatrixStorage(rows, columns, rowPointers, columnIndices, values); + return Sparse(storage); + } + + /// + /// Create a new sparse matrix from a compressed sparse column format. + /// This new matrix will be independent from the given arrays. + /// A new memory block will be allocated for storing the matrix. + /// + public Matrix SparseFromCompressedSparseColumnFormat(int rows, int columns, int nonZeroCount, int[] cscRowIndices, int[] cscColumnPointers, T[] cscValues) + { + if (cscValues == null) + throw new NullReferenceException(nameof(cscValues)); + if (cscRowIndices == null) + throw new NullReferenceException(nameof(cscRowIndices)); + if (cscColumnPointers == null) + throw new NullReferenceException(nameof(cscColumnPointers)); + if (cscColumnPointers.Length < columns) + { + var message = string.Format(Resources.ArgumentArrayWrongLength, columns + 1); + throw new Exception(message); + } + if (nonZeroCount != cscColumnPointers[columns]) + { + var message = string.Format("{0} should be same to {1}", nameof(nonZeroCount), cscColumnPointers[columns]); + throw new Exception(message); + } + + // convert from CSC to CSR + + var csrValues = new T[nonZeroCount]; + var csrRowPointers = new int[rows + 1]; + var csrColumnIndices = new int[nonZeroCount]; + + for (int i = 0; i < columns; i++) + { + for (int j = cscColumnPointers[i]; j < cscColumnPointers[i + 1]; j++) + { + csrRowPointers[cscRowIndices[j] + 1]++; + } + } + for (int i = 1; i < rows + 1; i++) + { + csrRowPointers[i] += csrRowPointers[i - 1]; + } + var curr = new int[rows]; + for (int i = 0; i < columns; i++) + { + for (int j = cscColumnPointers[i]; j < cscColumnPointers[i + 1]; j++) + { + var loc = csrRowPointers[cscRowIndices[j]] + curr[cscRowIndices[j]]; + curr[cscRowIndices[j]]++; + csrColumnIndices[loc] = i; + csrValues[loc] = cscValues[j]; + } + } + + var storage = new SparseCompressedRowMatrixStorage(rows, columns, csrRowPointers, csrColumnIndices, csrValues); + return Sparse(storage); + } + /// /// Create a new diagonal matrix straight from an initialized matrix storage instance. /// The storage is used directly without copying. diff --git a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs index 5067dfd2..824a6c89 100644 --- a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs @@ -78,6 +78,14 @@ namespace MathNet.Numerics.LinearAlgebra.Storage Values = new T[0]; } + internal SparseCompressedRowMatrixStorage(int rows, int columns, int[] rowPointers, int[] columnIndices, T[] values) + : base(rows, columns) + { + RowPointers = rowPointers; + ColumnIndices = columnIndices; + Values = values; + } + /// /// True if the matrix storage format is dense. /// @@ -281,6 +289,75 @@ namespace MathNet.Numerics.LinearAlgebra.Storage MapInplace(x => x, Zeros.AllowSkip); } + /// + /// Fill zeros explicitly on the diagonal entries as required by the Intel MKL direct sparse solver. + /// + public void PopulateExplicitZerosOnDiagonal() + { + var delta = 0; // number of missing diagonal entries + + for (int row = 0; row < RowCount; row++) + { + var found = false; + for (int j = RowPointers[row]; j < RowPointers[row + 1]; j++) + { + if (ColumnIndices[j] == row) + { + found = true; + break; + } + } + if (!found) delta++; + } + + if (delta > 0) + { + var size = Values.Length + delta; + if (size > int.MaxValue) + { + throw new NotSupportedException(Resources.TooManyElements); + } + + var newRowPointers = new int[RowCount + 1]; + var newColumnIndices = new int[size]; + var newValues = new T[size]; + + delta = 0; + for (int row = 0; row < RowCount; row++) + { + var found = false; + for (int j = RowPointers[row]; j < RowPointers[row + 1]; j++) + { + newColumnIndices[j + delta] = ColumnIndices[j]; + newValues[j + delta] = Values[j]; + if (ColumnIndices[j] == row) + { + found = true; + } + } + if (!found) + { + var start = RowPointers[row] + delta; + var end = RowPointers[row + 1] + delta; + var count = end - start + 1; + + newColumnIndices[end] = row; + newValues[end] = Zero; + + // Ordering may be not necessary + Sorting.Sort(newColumnIndices, newValues, start, count); + + delta++; + } + newRowPointers[row + 1] = RowPointers[row + 1] + delta; + } + + Array.Copy(newRowPointers, RowPointers, RowCount + 1); + ColumnIndices = newColumnIndices; + Values = newValues; + } + } + /// /// Returns a hash code for this instance. /// diff --git a/src/Numerics/Providers/Common/Mkl/MklProviderCapabilities.cs b/src/Numerics/Providers/Common/Mkl/MklProviderCapabilities.cs index 711c15b4..0e90dd2e 100644 --- a/src/Numerics/Providers/Common/Mkl/MklProviderCapabilities.cs +++ b/src/Numerics/Providers/Common/Mkl/MklProviderCapabilities.cs @@ -56,7 +56,9 @@ namespace MathNet.Numerics.Providers.Common.Mkl VectorFunctionsMajor = 130, VectorFunctionsMinor = 131, FourierTransformMajor = 384, - FourierTransformMinor = 385 + FourierTransformMinor = 385, + SparseSolverMajor = 512, + SparseSolverMinor = 513, } } diff --git a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs index 0bd306e6..1abaffea 100644 --- a/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs +++ b/src/Numerics/Providers/Common/Mkl/SafeNativeMethods.cs @@ -430,6 +430,30 @@ namespace MathNet.Numerics.Providers.Common.Mkl #endregion FFT + #region Direct Sparse Solver + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int s_dss_solve(int matrixStructure, int matrixType, int systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, float[] values, + int nRhs, [In, Out] float[] rhs, [In, Out] float[] solution); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int d_dss_solve(int matrixStructure, int matrixType, int systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, double[] values, + int nRhs, [In, Out] double[] rhs, [In, Out] double[] solution); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int c_dss_solve(int matrixStructure, int matrixType, int systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, Complex32[] values, + int nRhs, [In, Out] Complex32[] rhs, [In, Out] Complex32[] solution); + + [DllImport(_DllName, ExactSpelling = true, SetLastError = false, CallingConvention = CallingConvention.Cdecl)] + internal static extern int z_dss_solve(int matrixStructure, int matrixType, int systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, Complex[] values, + int nRhs, [In, Out] Complex[] rhs, [In, Out] Complex[] solution); + + #endregion Direct Sparse Solver + // ReSharper restore InconsistentNaming } } diff --git a/src/Numerics/Providers/SparseSolver/ISparseSolverProvider.cs b/src/Numerics/Providers/SparseSolver/ISparseSolverProvider.cs new file mode 100644 index 00000000..5e7ccc94 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/ISparseSolverProvider.cs @@ -0,0 +1,115 @@ +using Complex = System.Numerics.Complex; + +namespace MathNet.Numerics.Providers.SparseSolver +{ + /// + /// Structure option. + /// + public enum DssMatrixStructure : int + { + Symmetric = 536870976, + SymmetricStructure = 536871040, + Nonsymmetric = 536871104, + SymmetricComplex = 536871168, + SymmetricStructureComplex = 536871232, + NonsymmetricComplex = 536871296, + } + + /// + /// Factorization option. + /// + public enum DssMatrixType : int + { + PositiveDefinite = 134217792, + Indefinite = 134217856, + HermitianPositiveDefinite = 134217920, + HermitianIndefinite = 134217984 + } + + /// + /// Solver step's substitution. + /// + public enum DssSystemType : int + { + /// + /// Solve a system, Ax = b. + /// + DontTranspose = 0, + /// + /// Solve a transposed system, A'x = b + /// + Transpose = 262144, + /// + /// Solve a conjugate transposed system, A†x = b + /// + ConjugateTranspose = 524288, + } + + /// + /// Status values + /// + public enum DssStatus : int + { + /// + /// The operation was successful. + /// + MKL_DSS_SUCCESS = 0, + MKL_DSS_ZERO_PIVOT = -1, + MKL_DSS_OUT_OF_MEMORY = -2, + MKL_DSS_FAILURE = -3, + MKL_DSS_ROW_ERR = -4, + MKL_DSS_COL_ERR = -5, + MKL_DSS_TOO_FEW_VALUES = -6, + MKL_DSS_TOO_MANY_VALUES = -7, + MKL_DSS_NOT_SQUARE = -8, + MKL_DSS_STATE_ERR = -9, + MKL_DSS_INVALID_OPTION = -10, + MKL_DSS_OPTION_CONFLICT = -11, + MKL_DSS_MSG_LVL_ERR = -12, + MKL_DSS_TERM_LVL_ERR = -13, + MKL_DSS_STRUCTURE_ERR = -14, + MKL_DSS_REORDER_ERR = -15, + MKL_DSS_VALUES_ERR = -16, + MKL_DSS_STATISTICS_INVALID_MATRIX = -17, + MKL_DSS_STATISTICS_INVALID_STATE = -18, + MKL_DSS_STATISTICS_INVALID_STRING = -19, + MKL_DSS_REORDER1_ERR = -20, + MKL_DSS_PREORDER_ERR = -21, + MKL_DSS_DIAG_ERR = -22, + MKL_DSS_I32BIT_ERR = -23, + MKL_DSS_OOC_MEM_ERR = -24, + MKL_DSS_OOC_OC_ERR = -25, + MKL_DSS_OOC_RW_ERR = -26, + } + + public interface ISparseSolverProvider : + ISparseSolverProvider, + ISparseSolverProvider, + ISparseSolverProvider, + ISparseSolverProvider + { + /// + /// Try to find out whether the provider is available, at least in principle. + /// Verification may still fail if available, but it will certainly fail if unavailable. + /// + bool IsAvailable(); + + /// + /// Initialize and verify that the provided is indeed available. If not, fall back to alternatives like the managed provider + /// + void InitializeVerify(); + + /// + /// Frees memory buffers, caches and handles allocated in or to the provider. + /// Does not unload the provider itself, it is still usable afterwards. + /// + void FreeResources(); + } + + public interface ISparseSolverProvider + where T : struct + { + DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, int rows, int cols, int nnz, int[] rowIdx, int[] colPtr, T[] values, int nRhs, T[] rhs, T[] solution); + } +} + diff --git a/src/Numerics/Providers/SparseSolver/Managed/ManagedSparseSolverProvider.cs b/src/Numerics/Providers/SparseSolver/Managed/ManagedSparseSolverProvider.cs new file mode 100644 index 00000000..3b1e6258 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Managed/ManagedSparseSolverProvider.cs @@ -0,0 +1,68 @@ +using System; +using System.Numerics; + +namespace MathNet.Numerics.Providers.SparseSolver.Managed +{ + /// + /// The managed sparse solver provider + /// + internal partial class ManagedSparseSolverProvider : ISparseSolverProvider + { + /// + /// Try to find out whether the provider is available, at least in principle. + /// Verification may still fail if available, but it will certainly fail if unavailable. + /// + public virtual bool IsAvailable() + { + return true; + } + + /// + /// Initialize and verify that the provided is indeed available. If not, fall back to alternatives like the managed provider + /// + public virtual void InitializeVerify() + { + } + + /// + /// Frees memory buffers, caches and handles allocated in or to the provider. + /// Does not unload the provider itself, it is still usable afterwards. + /// + public virtual void FreeResources() + { + } + + public override string ToString() + { + return "Managed"; + } + + public virtual DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] ColumnIndices, float[] values, + int nRhs, float[] rhs, float[] solution) + { + throw new NotImplementedException(); + } + + public virtual DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] ColumnIndices, double[] values, + int nRhs, double[] rhs, double[] solution) + { + throw new NotImplementedException(); + } + + public virtual DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] ColumnIndices, Complex32[] values, + int nRhs, Complex32[] rhs, Complex32[] solution) + { + throw new NotImplementedException(); + } + + public virtual DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] ColumnIndices, Complex[] values, + int nRhs, Complex[] rhs, Complex[] solution) + { + throw new NotImplementedException(); + } + } +} diff --git a/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex.cs b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex.cs new file mode 100644 index 00000000..d64554a0 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex.cs @@ -0,0 +1,85 @@ +#if NATIVE + +using MathNet.Numerics.Properties; +using MathNet.Numerics.Providers.Common.Mkl; +using System; +using System.Security; +using Complex = System.Numerics.Complex; + +namespace MathNet.Numerics.Providers.SparseSolver.Mkl +{ + /// + /// Intel's Math Kernel Library (MKL) direct sparse solver provider. + /// + internal partial class MklSparseSolverProvider + { + /// + /// Solves sparse linear systems of equations, AX = B. + /// + /// The symmetricity of the matrix. For a symmetric matrix, only upper ot lower triangular matrix is used. + /// The definiteness of the matrix. + /// The type of the systems. + /// The number of rows of matrix. + /// The number of columns of matrix. + /// The number of non zero elements of matrix. + /// The array containing the row indices of the existing rows. + /// The array containing the column indices of the non-zero values + /// The array that contains the non-zero elements of matrix. No diagonal element can be ommitted. + /// The number of columns of the right hand side matrix. + /// The right hand side matrix + /// The left hand side matrix + /// The status of the solver. + [SecuritySafeCritical] + public override DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, Complex[] values, + int nRhs, Complex[] rhs, Complex[] solution) + { + if (rowCount != columnCount) + { + throw new ArgumentNullException(Resources.ArgumentMatrixSymmetric); + } + + if (rowPointers == null) + { + throw new ArgumentNullException(nameof(rowPointers)); + } + + if (columnIndices == null) + { + throw new ArgumentNullException(nameof(columnIndices)); + } + + if (values == null) + { + throw new ArgumentNullException(nameof(values)); + } + + if (rhs == null) + { + throw new ArgumentNullException(nameof(rhs)); + } + + if (solution == null) + { + throw new ArgumentNullException(nameof(solution)); + } + + if (rowCount * nRhs != rhs.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(rhs)); + } + + if (columnCount * nRhs != solution.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(solution)); + } + + var error = SafeNativeMethods.z_dss_solve((int)matrixStructure, (int)matrixType, (int)systemType, + rowCount, columnCount, nonZerosCount, rowPointers, columnIndices, values, + nRhs, rhs, solution); + return (DssStatus)error; + } + } +} + +#endif diff --git a/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex32.cs b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex32.cs new file mode 100644 index 00000000..4c7af106 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Complex32.cs @@ -0,0 +1,84 @@ +#if NATIVE + +using MathNet.Numerics.Properties; +using MathNet.Numerics.Providers.Common.Mkl; +using System; +using System.Security; + +namespace MathNet.Numerics.Providers.SparseSolver.Mkl +{ + /// + /// Intel's Math Kernel Library (MKL) direct sparse solver provider. + /// + internal partial class MklSparseSolverProvider + { + /// + /// Solves sparse linear systems of equations, AX = B. + /// + /// The symmetricity of the matrix. For a symmetric matrix, only upper ot lower triangular matrix is used. + /// The definiteness of the matrix. + /// The type of the systems. + /// The number of rows of matrix. + /// The number of columns of matrix. + /// The number of non zero elements of matrix. + /// The array containing the row indices of the existing rows. + /// The array containing the column indices of the non-zero values + /// The array that contains the non-zero elements of matrix. No diagonal element can be ommitted. + /// The number of columns of the right hand side matrix. + /// The right hand side matrix + /// The left hand side matrix + /// The status of the solver. + [SecuritySafeCritical] + public override DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, Complex32[] values, + int nRhs, Complex32[] rhs, Complex32[] solution) + { + if (rowCount != columnCount) + { + throw new ArgumentException(Resources.ArgumentMatrixSymmetric); + } + + if (rowPointers == null) + { + throw new ArgumentNullException(nameof(rowPointers)); + } + + if (columnIndices == null) + { + throw new ArgumentNullException(nameof(columnIndices)); + } + + if (values == null) + { + throw new ArgumentNullException(nameof(values)); + } + + if (rhs == null) + { + throw new ArgumentNullException(nameof(rhs)); + } + + if (solution == null) + { + throw new ArgumentNullException(nameof(solution)); + } + + if (rowCount * nRhs != rhs.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(rhs)); + } + + if (columnCount * nRhs != solution.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(solution)); + } + + var error = SafeNativeMethods.c_dss_solve((int)matrixStructure, (int)matrixType, (int)systemType, + rowCount, columnCount, nonZerosCount, rowPointers, columnIndices, values, + nRhs, rhs, solution); + return (DssStatus)error; + } + } +} + +#endif diff --git a/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Double.cs b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Double.cs new file mode 100644 index 00000000..f1563f4e --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Double.cs @@ -0,0 +1,84 @@ +#if NATIVE + +using MathNet.Numerics.Properties; +using MathNet.Numerics.Providers.Common.Mkl; +using System; +using System.Security; + +namespace MathNet.Numerics.Providers.SparseSolver.Mkl +{ + /// + /// Intel's Math Kernel Library (MKL) direct sparse solver provider. + /// + internal partial class MklSparseSolverProvider + { + /// + /// Solves sparse linear systems of equations, AX = B. + /// + /// The symmetricity of the matrix. For a symmetric matrix, only upper ot lower triangular matrix is used. + /// The definiteness of the matrix. + /// The type of the systems. + /// The number of rows of matrix. + /// The number of columns of matrix. + /// The number of non zero elements of matrix. + /// The array containing the row indices of the existing rows. + /// The array containing the column indices of the non-zero values + /// The array that contains the non-zero elements of matrix. No diagonal element can be ommitted. + /// The number of columns of the right hand side matrix. + /// The right hand side matrix + /// The left hand side matrix + /// The status of the solver. + [SecuritySafeCritical] + public override DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, double[] values, + int nRhs, double[] rhs, double[] solution) + { + if (rowCount != columnCount) + { + throw new ArgumentException(Resources.ArgumentMatrixSymmetric); + } + + if (rowPointers == null) + { + throw new ArgumentNullException(nameof(rowPointers)); + } + + if (columnIndices == null) + { + throw new ArgumentNullException(nameof(columnIndices)); + } + + if (values == null) + { + throw new ArgumentNullException(nameof(values)); + } + + if (rhs == null) + { + throw new ArgumentNullException(nameof(rhs)); + } + + if (solution == null) + { + throw new ArgumentNullException(nameof(solution)); + } + + if (rowCount * nRhs != rhs.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(rhs)); + } + + if (columnCount * nRhs != solution.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(solution)); + } + + var error = SafeNativeMethods.d_dss_solve((int)matrixStructure, (int)matrixType, (int)systemType, + rowCount, columnCount, nonZerosCount, rowPointers, columnIndices, values, + nRhs, rhs, solution); + return (DssStatus)error; + } + } +} + +#endif diff --git a/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Single.cs b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Single.cs new file mode 100644 index 00000000..475e0da7 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.Single.cs @@ -0,0 +1,83 @@ +#if NATIVE + +using MathNet.Numerics.Properties; +using MathNet.Numerics.Providers.Common.Mkl; +using System; +using System.Security; + +namespace MathNet.Numerics.Providers.SparseSolver.Mkl +{ + /// + /// Intel's Math Kernel Library (MKL) direct sparse solver provider. + /// + internal partial class MklSparseSolverProvider + { + /// + /// Solves sparse linear systems of equations, AX = B. + /// + /// The symmetricity of the matrix. For a symmetric matrix, only upper ot lower triangular matrix is used. + /// The definiteness of the matrix. + /// The type of the systems. + /// The number of rows of matrix. + /// The number of columns of matrix. + /// The number of non zero elements of matrix. + /// The array containing the row indices of the existing rows. + /// The array containing the column indices of the non-zero values + /// The array that contains the non-zero elements of matrix. No diagonal element can be ommitted. + /// The number of columns of the right hand side matrix. + /// The right hand side matrix + /// The left hand side matrix + /// The status of the solver. + [SecuritySafeCritical] + public override DssStatus Solve(DssMatrixStructure matrixStructure, DssMatrixType matrixType, DssSystemType systemType, + int rowCount, int columnCount, int nonZerosCount, int[] rowPointers, int[] columnIndices, float[] values, + int nRhs, float[] rhs, float[] solution) + { + if (rowCount != columnCount) + { + throw new ArgumentException(Resources.ArgumentMatrixSymmetric); + } + + if (rowPointers == null) + { + throw new ArgumentNullException(nameof(rowPointers)); + } + + if (columnIndices == null) + { + throw new ArgumentNullException(nameof(columnIndices)); + } + + if (values == null) + { + throw new ArgumentNullException(nameof(values)); + } + + if (rhs == null) + { + throw new ArgumentNullException(nameof(rhs)); + } + + if (solution == null) + { + throw new ArgumentNullException(nameof(solution)); + } + + if (rowCount * nRhs != rhs.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(rhs)); + } + + if (columnCount * nRhs != solution.Length) + { + throw new ArgumentException(Resources.ArgumentArraysSameLength, nameof(solution)); + } + + var error = SafeNativeMethods.s_dss_solve((int)matrixStructure, (int)matrixType, (int)systemType, + rowCount, columnCount, nonZerosCount, rowPointers, columnIndices, values, + nRhs, rhs, solution); + return (DssStatus)error; + } + } +} +#endif diff --git a/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.cs b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.cs new file mode 100644 index 00000000..e544cde7 --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/Mkl/MklSparseSolverProvider.cs @@ -0,0 +1,77 @@ +#if NATIVE + +using MathNet.Numerics.Providers.Common.Mkl; +using System; + +namespace MathNet.Numerics.Providers.SparseSolver.Mkl +{ + /// + /// Intel's Math Kernel Library (MKL) sparse solver provider. + /// + internal partial class MklSparseSolverProvider : Managed.ManagedSparseSolverProvider, IDisposable + { + const int MinimumCompatibleRevision = 12; + + readonly string _hintPath; + + int sparseSolverMajor; + int sparseSolverMinor; + + /// Hint path where to look for the native binaries + internal MklSparseSolverProvider(string hintPath) + { + _hintPath = hintPath; + } + + /// + /// Try to find out whether the provider is available, at least in principle. + /// Verification may still fail if available, but it will certainly fail if unavailable. + /// + public override bool IsAvailable() + { + return MklProvider.IsAvailable(hintPath: _hintPath); + } + + /// + /// Initialize and verify that the provided is indeed available. + /// If calling this method fails, consider to fall back to alternatives like the managed provider. + /// + public override void InitializeVerify() + { + int revision = MklProvider.Load(_hintPath); + if (revision < MinimumCompatibleRevision) + { + throw new NotSupportedException($"MKL Native Provider revision r{revision} is too old. Consider upgrading to a newer version. Revision r{MinimumCompatibleRevision} and newer are supported."); + } + + sparseSolverMajor = SafeNativeMethods.query_capability((int)ProviderCapability.SparseSolverMajor); + sparseSolverMinor = SafeNativeMethods.query_capability((int)ProviderCapability.SparseSolverMinor); + if (!(sparseSolverMajor == 1 && sparseSolverMinor >= 0)) + { + throw new NotSupportedException(string.Format("MKL Native Provider not compatible. Expecting sparse solver v1 but provider implements v{0}.", sparseSolverMajor)); + } + } + + /// + /// Frees memory buffers, caches and handles allocated in or to the provider. + /// Does not unload the provider itself, it is still usable afterwards. + /// + public override void FreeResources() + { + MklProvider.FreeResources(); + } + + public override string ToString() + { + return MklProvider.Describe(); + } + + public void Dispose() + { + FreeResources(); + } + } +} + +#endif + diff --git a/src/Numerics/Providers/SparseSolver/SparseSolverControl.cs b/src/Numerics/Providers/SparseSolver/SparseSolverControl.cs new file mode 100644 index 00000000..4f3cea0d --- /dev/null +++ b/src/Numerics/Providers/SparseSolver/SparseSolverControl.cs @@ -0,0 +1,165 @@ +using System; + +namespace MathNet.Numerics.Providers.SparseSolver +{ + public static class SparseSolverControl + { + const string EnvVarSSProvider = "MathNetNumericsSSProvider"; + const string EnvVarSSProviderPath = "MathNetNumericsSSProviderPath"; + + static ISparseSolverProvider _sparseSolverProvider; + static readonly object StaticLock = new object(); + + /// + /// Gets or sets the sparse solver provider. Consider to use UseNativeMKL or UseManaged instead. + /// + /// The linear algebra provider. + public static ISparseSolverProvider Provider + { + get + { + if (_sparseSolverProvider == null) + { + lock (StaticLock) + { + if (_sparseSolverProvider == null) + { + UseDefault(); + } + } + } + + return _sparseSolverProvider; + } + set + { + value.InitializeVerify(); + + // only actually set if verification did not throw + _sparseSolverProvider = value; + } + } + + /// + /// Optional path to try to load native provider binaries from. + /// If not set, Numerics will fall back to the environment variable + /// `MathNetNumericsSSProviderPath` or the default probing paths. + /// + public static string HintPath { get; set; } + + public static ISparseSolverProvider CreateManaged() + { + return new Managed.ManagedSparseSolverProvider(); + } + + public static void UseManaged() + { + Provider = CreateManaged(); + } + +#if NATIVE + public static ISparseSolverProvider CreateNativeMKL() + { + return new Mkl.MklSparseSolverProvider(GetCombinedHintPath()); + } + + public static void UseNativeMKL() + { + Provider = CreateNativeMKL(); + } + + public static bool TryUseNativeMKL() + { + return TryUse(CreateNativeMKL()); + } + + /// + /// Try to use a native provider, if available. + /// + public static bool TryUseNative() + { + return TryUseNativeMKL(); + } +#endif + + static bool TryUse(ISparseSolverProvider provider) + { + try + { + if (!provider.IsAvailable()) + { + return false; + } + + Provider = provider; + return true; + } + catch + { + // intentionally swallow exceptions here - use the explicit variants if you're interested in why + return false; + } + } + + /// + /// Use the best provider available. + /// + public static void UseBest() + { +#if NATIVE + if (!TryUseNative()) + { + UseManaged(); + } +#else + UseManaged(); +#endif + } + + /// + /// Use a specific provider if configured, e.g. using the + /// "MathNetNumericsDSSProvider" environment variable, + /// or fall back to the best provider. + /// + public static void UseDefault() + { +#if NATIVE + var value = Environment.GetEnvironmentVariable(EnvVarSSProvider); + switch (value != null ? value.ToUpperInvariant() : string.Empty) + { + + case "MKL": + UseNativeMKL(); + break; + + default: + UseBest(); + break; + } +#else + UseBest(); +#endif + } + + public static void FreeResources() + { + Provider.FreeResources(); + } + + static string GetCombinedHintPath() + { + if (!String.IsNullOrEmpty(HintPath)) + { + return HintPath; + } + + var value = Environment.GetEnvironmentVariable(EnvVarSSProviderPath); + if (!String.IsNullOrEmpty(value)) + { + return value; + } + + return null; + } + } +}