From 19873cb3046c82698359a0a36142811c41615d75 Mon Sep 17 00:00:00 2001 From: Christoph Ruegg Date: Sun, 1 Dec 2013 16:16:06 +0100 Subject: [PATCH] LA: dense-diagonal TransposeAndMultiply --- .../LinearAlgebra/Complex/DenseMatrix.cs | 33 +++++++++++++++---- .../LinearAlgebra/Complex32/DenseMatrix.cs | 33 +++++++++++++++---- .../LinearAlgebra/Double/DenseMatrix.cs | 30 ++++++++++++++--- .../LinearAlgebra/Single/DenseMatrix.cs | 33 +++++++++++++++---- .../Complex/DiagonalMatrixTests.cs | 20 ++++++++++- .../Complex32/DiagonalMatrixTests.cs | 21 +++++++++++- .../Double/DiagonalMatrixTests.cs | 20 ++++++++++- .../Single/DiagonalMatrixTests.cs | 20 ++++++++++- 8 files changed, 180 insertions(+), 30 deletions(-) diff --git a/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs index 43adc606..6d6d575e 100644 --- a/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs @@ -684,14 +684,9 @@ namespace MathNet.Numerics.LinearAlgebra.Complex /// The result of the multiplication. protected override void DoTransposeAndMultiply(Matrix other, Matrix result) { - var denseOther = other as DenseMatrix; + var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.DontTranspose, @@ -705,7 +700,31 @@ namespace MathNet.Numerics.LinearAlgebra.Complex denseOther._columnCount, 0.0, denseResult._values); + return; } + + var diagonalOther = other.Storage as DiagonalMatrixStorage; + if (diagonalOther != null) + { + var diagonal = diagonalOther.Data; + var d = Math.Min(ColumnCount, other.RowCount); + if (d < other.RowCount) + { + result.ClearSubMatrix(0, RowCount, ColumnCount, other.RowCount - ColumnCount); + } + int index = 0; + for (int j = 0; j < d; j++) + { + for (int i = 0; i < RowCount; i++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + } + return; + } + + base.DoTransposeAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs index 37721c1a..b9537954 100644 --- a/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs @@ -679,14 +679,9 @@ namespace MathNet.Numerics.LinearAlgebra.Complex32 /// The result of the multiplication. protected override void DoTransposeAndMultiply(Matrix other, Matrix result) { - var denseOther = other as DenseMatrix; + var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.DontTranspose, @@ -700,7 +695,31 @@ namespace MathNet.Numerics.LinearAlgebra.Complex32 denseOther._columnCount, 0.0f, denseResult._values); + return; } + + var diagonalOther = other.Storage as DiagonalMatrixStorage; + if (diagonalOther != null) + { + var diagonal = diagonalOther.Data; + var d = Math.Min(ColumnCount, other.RowCount); + if (d < other.RowCount) + { + result.ClearSubMatrix(0, RowCount, ColumnCount, other.RowCount - ColumnCount); + } + int index = 0; + for (int j = 0; j < d; j++) + { + for (int i = 0; i < RowCount; i++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + } + return; + } + + base.DoTransposeAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs index f9370001..74831a77 100644 --- a/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs @@ -661,11 +661,7 @@ namespace MathNet.Numerics.LinearAlgebra.Double { var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - if (denseOther == null || denseResult == null) - { - base.DoTransposeAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.DontTranspose, @@ -679,7 +675,31 @@ namespace MathNet.Numerics.LinearAlgebra.Double denseOther._columnCount, 0.0, denseResult._values); + return; } + + var diagonalOther = other.Storage as DiagonalMatrixStorage; + if (diagonalOther != null) + { + var diagonal = diagonalOther.Data; + var d = Math.Min(ColumnCount, other.RowCount); + if (d < other.RowCount) + { + result.ClearSubMatrix(0, RowCount, ColumnCount, other.RowCount - ColumnCount); + } + int index = 0; + for (int j = 0; j < d; j++) + { + for (int i = 0; i < RowCount; i++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + } + return; + } + + base.DoTransposeAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs index b112612d..0a17c439 100644 --- a/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs @@ -659,14 +659,9 @@ namespace MathNet.Numerics.LinearAlgebra.Single /// The result of the multiplication. protected override void DoTransposeAndMultiply(Matrix other, Matrix result) { - var denseOther = other as DenseMatrix; + var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.DontTranspose, @@ -680,7 +675,31 @@ namespace MathNet.Numerics.LinearAlgebra.Single denseOther._columnCount, 0.0f, denseResult._values); + return; } + + var diagonalOther = other.Storage as DiagonalMatrixStorage; + if (diagonalOther != null) + { + var diagonal = diagonalOther.Data; + var d = Math.Min(ColumnCount, other.RowCount); + if (d < other.RowCount) + { + result.ClearSubMatrix(0, RowCount, ColumnCount, other.RowCount - ColumnCount); + } + int index = 0; + for (int j = 0; j < d; j++) + { + for (int i = 0; i < RowCount; i++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + } + return; + } + + base.DoTransposeAndMultiply(other, result); } /// diff --git a/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs index c54dc08c..d39d57cc 100644 --- a/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs @@ -387,7 +387,7 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex } [Test] - public void DenseDiagonalMatrixMultiplication() + public void DenseDiagonalMatrixMultiply() { var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); @@ -403,5 +403,23 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 10, 2d)).Equals(wide.Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 2, 2d)).Equals(wide.Multiply(2d).SubMatrix(0, 3, 0, 2))); } + + [Test] + public void DenseDiagonalMatrixTransposeAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2d)).Equals(tall.Multiply(2d))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(5, 3, 2d)).Equals(tall.Multiply(2d).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(2, 3, 2d)).Equals(tall.Multiply(2d).SubMatrix(0, 8, 0, 2))); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2d)).Equals(wide.Multiply(2d))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(10, 8, 2d)).Equals(wide.Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(2, 8, 2d)).Equals(wide.Multiply(2d).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs index 2623229d..c6db5d0d 100644 --- a/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs @@ -381,8 +381,9 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex32 var matrix = TestMatrices["Square3x3"]; Assert.IsTrue(matrix.IsSymmetric); } + [Test] - public void DenseDiagonalMatrixMultiplication() + public void DenseDiagonalMatrixMultiply() { var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); @@ -398,5 +399,23 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex32 Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 10, 2f)).Equals(wide.Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 2, 2f)).Equals(wide.Multiply(2f).SubMatrix(0, 3, 0, 2))); } + + [Test] + public void DenseDiagonalMatrixTransposeAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2f)).Equals(tall.Multiply(2f))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(5, 3, 2f)).Equals(tall.Multiply(2f).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(2, 3, 2f)).Equals(tall.Multiply(2f).SubMatrix(0, 8, 0, 2))); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2f)).Equals(wide.Multiply(2f))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(10, 8, 2f)).Equals(wide.Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(2, 8, 2f)).Equals(wide.Multiply(2f).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs index cc3e5240..d90c7c1e 100644 --- a/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs @@ -414,7 +414,7 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double } [Test] - public void DenseDiagonalMatrixMultiplication() + public void DenseDiagonalMatrixMultiply() { var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); @@ -430,5 +430,23 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 10, 2d)).Equals(wide.Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 2, 2d)).Equals(wide.Multiply(2d).SubMatrix(0, 3, 0, 2))); } + + [Test] + public void DenseDiagonalMatrixTransposeAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2d)).Equals(tall.Multiply(2d))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(5, 3, 2d)).Equals(tall.Multiply(2d).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(2, 3, 2d)).Equals(tall.Multiply(2d).SubMatrix(0, 8, 0, 2))); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2d)).Equals(wide.Multiply(2d))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(10, 8, 2d)).Equals(wide.Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(2, 8, 2d)).Equals(wide.Multiply(2d).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs index 8773fce1..f6c17198 100644 --- a/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs @@ -381,7 +381,7 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Single } [Test] - public void DenseDiagonalMatrixMultiplication() + public void DenseDiagonalMatrixMultiply() { var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); @@ -397,5 +397,23 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Single Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 10, 2f)).Equals(wide.Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); Assert.IsTrue((wide*Matrix.Build.Diagonal(8, 2, 2f)).Equals(wide.Multiply(2f).SubMatrix(0, 3, 0, 2))); } + + [Test] + public void DenseDiagonalMatrixTransposeAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2f)).Equals(tall.Multiply(2f))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(5, 3, 2f)).Equals(tall.Multiply(2f).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(tall.TransposeAndMultiply(Matrix.Build.Diagonal(2, 3, 2f)).Equals(tall.Multiply(2f).SubMatrix(0, 8, 0, 2))); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2f)).Equals(wide.Multiply(2f))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(10, 8, 2f)).Equals(wide.Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(wide.TransposeAndMultiply(Matrix.Build.Diagonal(2, 8, 2f)).Equals(wide.Multiply(2f).SubMatrix(0, 3, 0, 2))); + } } }