diff --git a/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs index 6d6d575e..7deaf8d1 100644 --- a/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Complex/DenseMatrix.cs @@ -767,12 +767,7 @@ namespace MathNet.Numerics.LinearAlgebra.Complex { var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeThisAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.Transpose, @@ -786,7 +781,32 @@ 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(RowCount, other.ColumnCount); + if (d < other.ColumnCount) + { + result.ClearSubMatrix(0, ColumnCount, RowCount, other.ColumnCount - RowCount); + } + int index = 0; + for (int i = 0; i < ColumnCount; i++) + { + for (int j = 0; j < d; j++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + index += (RowCount - d); + } + return; + } + + base.DoTransposeThisAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs index b9537954..1337b9aa 100644 --- a/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Complex32/DenseMatrix.cs @@ -762,12 +762,7 @@ namespace MathNet.Numerics.LinearAlgebra.Complex32 { var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeThisAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.Transpose, @@ -781,7 +776,32 @@ 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(RowCount, other.ColumnCount); + if (d < other.ColumnCount) + { + result.ClearSubMatrix(0, ColumnCount, RowCount, other.ColumnCount - RowCount); + } + int index = 0; + for (int i = 0; i < ColumnCount; i++) + { + for (int j = 0; j < d; j++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + index += (RowCount - d); + } + return; + } + + base.DoTransposeThisAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs index 74831a77..a45cd15f 100644 --- a/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Double/DenseMatrix.cs @@ -742,12 +742,7 @@ namespace MathNet.Numerics.LinearAlgebra.Double { var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeThisAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.Transpose, @@ -761,7 +756,32 @@ 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(RowCount, other.ColumnCount); + if (d < other.ColumnCount) + { + result.ClearSubMatrix(0, ColumnCount, RowCount, other.ColumnCount - RowCount); + } + int index = 0; + for (int i = 0; i < ColumnCount; i++) + { + for (int j = 0; j < d; j++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + index += (RowCount - d); + } + return; + } + + base.DoTransposeThisAndMultiply(other, result); } /// diff --git a/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs b/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs index 0a17c439..61dc0c34 100644 --- a/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs +++ b/src/Numerics/LinearAlgebra/Single/DenseMatrix.cs @@ -742,12 +742,7 @@ namespace MathNet.Numerics.LinearAlgebra.Single { var denseOther = other as DenseMatrix; var denseResult = result as DenseMatrix; - - if (denseOther == null || denseResult == null) - { - base.DoTransposeThisAndMultiply(other, result); - } - else + if (denseOther != null && denseResult != null) { Control.LinearAlgebraProvider.MatrixMultiplyWithUpdate( Providers.LinearAlgebra.Transpose.Transpose, @@ -761,7 +756,32 @@ 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(RowCount, other.ColumnCount); + if (d < other.ColumnCount) + { + result.ClearSubMatrix(0, ColumnCount, RowCount, other.ColumnCount - RowCount); + } + int index = 0; + for (int i = 0; i < ColumnCount; i++) + { + for (int j = 0; j < d; j++) + { + result.At(i, j, _values[index]*diagonal[j]); + index++; + } + index += (RowCount - d); + } + return; + } + + base.DoTransposeThisAndMultiply(other, result); } /// diff --git a/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs index d39d57cc..55de88b8 100644 --- a/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Complex/DiagonalMatrixTests.cs @@ -390,7 +390,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex public void DenseDiagonalMatrixMultiply() { 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); @@ -408,7 +407,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex 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); @@ -421,5 +419,22 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex 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))); } + + [Test] + public void DenseDiagonalMatrixTransposeThisAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2d)).Equals(wide.Transpose().Multiply(2d))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 5, 2d)).Equals(wide.Transpose().Multiply(2d).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 2, 2d)).Equals(wide.Transpose().Multiply(2d).SubMatrix(0, 8, 0, 2))); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2d)).Equals(tall.Transpose().Multiply(2d))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 10, 2d)).Equals(tall.Transpose().Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 2, 2d)).Equals(tall.Transpose().Multiply(2d).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs index c6db5d0d..c4423a74 100644 --- a/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Complex32/DiagonalMatrixTests.cs @@ -386,7 +386,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex32 public void DenseDiagonalMatrixMultiply() { 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); @@ -404,7 +403,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex32 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); @@ -417,5 +415,22 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Complex32 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))); } + + [Test] + public void DenseDiagonalMatrixTransposeThisAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2f)).Equals(wide.Transpose().Multiply(2f))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 5, 2f)).Equals(wide.Transpose().Multiply(2f).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 2, 2f)).Equals(wide.Transpose().Multiply(2f).SubMatrix(0, 8, 0, 2))); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2f)).Equals(tall.Transpose().Multiply(2f))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 10, 2f)).Equals(tall.Transpose().Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 2, 2f)).Equals(tall.Transpose().Multiply(2f).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs index d90c7c1e..bb27d8b5 100644 --- a/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Double/DiagonalMatrixTests.cs @@ -417,7 +417,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double public void DenseDiagonalMatrixMultiply() { 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); @@ -435,7 +434,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double 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); @@ -448,5 +446,22 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Double 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))); } + + [Test] + public void DenseDiagonalMatrixTransposeThisAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2d)).Equals(wide.Transpose().Multiply(2d))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 5, 2d)).Equals(wide.Transpose().Multiply(2d).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 2, 2d)).Equals(wide.Transpose().Multiply(2d).SubMatrix(0, 8, 0, 2))); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2d)).Equals(tall.Transpose().Multiply(2d))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 10, 2d)).Equals(tall.Transpose().Multiply(2d).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 2, 2d)).Equals(tall.Transpose().Multiply(2d).SubMatrix(0, 3, 0, 2))); + } } } diff --git a/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs b/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs index f6c17198..6c253abd 100644 --- a/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs +++ b/src/UnitTests/LinearAlgebraTests/Single/DiagonalMatrixTests.cs @@ -384,7 +384,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Single public void DenseDiagonalMatrixMultiply() { 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); @@ -402,7 +401,6 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Single 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); @@ -415,5 +413,22 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests.Single 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))); } + + [Test] + public void DenseDiagonalMatrixTransposeThisAndMultiply() + { + var dist = new ContinuousUniform(-1.0, 1.0, new MersenneTwister()); + Assert.IsInstanceOf(Matrix.Build.DiagonalIdentity(3, 3)); + + var wide = Matrix.Build.Random(3, 8, dist); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(3).Multiply(2f)).Equals(wide.Transpose().Multiply(2f))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 5, 2f)).Equals(wide.Transpose().Multiply(2f).Append(Matrix.Build.Dense(8, 2)))); + Assert.IsTrue(wide.TransposeThisAndMultiply(Matrix.Build.Diagonal(3, 2, 2f)).Equals(wide.Transpose().Multiply(2f).SubMatrix(0, 8, 0, 2))); + + var tall = Matrix.Build.Random(8, 3, dist); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.DiagonalIdentity(8).Multiply(2f)).Equals(tall.Transpose().Multiply(2f))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 10, 2f)).Equals(tall.Transpose().Multiply(2f).Append(Matrix.Build.Dense(3, 2)))); + Assert.IsTrue(tall.TransposeThisAndMultiply(Matrix.Build.Diagonal(8, 2, 2f)).Equals(tall.Transpose().Multiply(2f).SubMatrix(0, 3, 0, 2))); + } } }