diff --git a/src/Numerics/LinearAlgebra/Builder.cs b/src/Numerics/LinearAlgebra/Builder.cs index 6147a514..05ef6c43 100644 --- a/src/Numerics/LinearAlgebra/Builder.cs +++ b/src/Numerics/LinearAlgebra/Builder.cs @@ -1363,7 +1363,8 @@ namespace MathNet.Numerics.LinearAlgebra /// /// Create a new vector with the same kind of the provided example. /// - public Vector SameAs(Matrix example, int length) + public Vector SameAs(Matrix example, int length) + where TU : struct, IEquatable, IFormattable { return example.Storage.IsDense ? Dense(length) : Sparse(length); } diff --git a/src/Numerics/LinearAlgebra/Matrix.cs b/src/Numerics/LinearAlgebra/Matrix.cs index dfd0bd00..96ffa61b 100644 --- a/src/Numerics/LinearAlgebra/Matrix.cs +++ b/src/Numerics/LinearAlgebra/Matrix.cs @@ -1572,5 +1572,29 @@ namespace MathNet.Numerics.LinearAlgebra Storage.MapIndexedToUnchecked(result.Storage, f, zeros, ExistingData.AssumeZeros); return result; } + + /// + /// For each row, applies a function f to each element of the row, threading an accumulator argument through the computation. + /// Returns a vector with the resulting accumulator states for each row. + /// + public Vector FoldRows(Func f, TU state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + var result = Vector.Build.SameAs(this, RowCount); + Storage.FoldRowsUnchecked(result.Storage, f, (x, c) => x, DenseVectorStorage.OfInit(RowCount, i => state), zeros); + return result; + } + + /// + /// For each column, applies a function f to each element of the column, threading an accumulator argument through the computation. + /// Returns a vector with the resulting accumulator states for each column. + /// + public Vector FoldColumns(Func f, TU state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + var result = Vector.Build.SameAs(this, ColumnCount); + Storage.FoldColumnsUnchecked(result.Storage, f, (x, c) => x, DenseVectorStorage.OfInit(ColumnCount, i => state), zeros); + return result; + } } } diff --git a/src/Numerics/LinearAlgebra/Storage/DenseColumnMajorMatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/DenseColumnMajorMatrixStorage.cs index ac284a4c..c68c8e21 100644 --- a/src/Numerics/LinearAlgebra/Storage/DenseColumnMajorMatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/DenseColumnMajorMatrixStorage.cs @@ -598,7 +598,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } - // FUNCTIONAL COMBINATORS + // FUNCTIONAL COMBINATORS: MAP public override void MapInplace(Func f, Zeros zeros = Zeros.AllowSkip) { @@ -724,5 +724,34 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } } + + // FUNCTIONAL COMBINATORS: FOLD + + internal override void FoldRowsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + for (int i = 0; i < RowCount; i++) + { + TU s = state.At(i); + for (int j = 0; j < ColumnCount; j++) + { + s = f(s, Data[j*RowCount + i]); + } + target.At(i, finalize(s, ColumnCount)); + } + } + + internal override void FoldColumnsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + for (int j = 0; j < ColumnCount; j++) + { + int offset = j*RowCount; + TU s = state.At(j); + for (int i = 0; i < RowCount; i++) + { + s = f(s, Data[offset + i]); + } + target.At(j, finalize(s, RowCount)); + } + } } } diff --git a/src/Numerics/LinearAlgebra/Storage/DiagonalMatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/DiagonalMatrixStorage.cs index 9a4bc881..389a0cf6 100644 --- a/src/Numerics/LinearAlgebra/Storage/DiagonalMatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/DiagonalMatrixStorage.cs @@ -595,7 +595,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } - // FUNCTIONAL COMBINATORS + // FUNCTIONAL COMBINATORS: MAP public override void MapInplace(Func f, Zeros zeros = Zeros.AllowSkip) { @@ -898,5 +898,63 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } } + + // FUNCTIONAL COMBINATORS: FOLD + + internal override void FoldRowsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + if (zeros == Zeros.AllowSkip) + { + for (int k = 0; k < Data.Length; k++) + { + target.At(k, finalize(f(state.At(k), Data[k]), 1)); + } + + for (int k = Data.Length; k < RowCount; k++) + { + target.At(k, finalize(state.At(k), 0)); + } + } + else + { + for (int i = 0; i < RowCount; i++) + { + TU s = state.At(i); + for (int j = 0; j < ColumnCount; j++) + { + s = f(s, i == j ? Data[i] : Zero); + } + target.At(i, finalize(s, ColumnCount)); + } + } + } + + internal override void FoldColumnsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + if (zeros == Zeros.AllowSkip) + { + for (int k = 0; k < Data.Length; k++) + { + target.At(k, finalize(f(state.At(k), Data[k]), 1)); + } + + for (int k = Data.Length; k < ColumnCount; k++) + { + target.At(k, finalize(state.At(k), 0)); + } + } + else + { + for (int j = 0; j < ColumnCount; j++) + { + TU s = state.At(j); + for (int i = 0; i < RowCount; i++) + { + s = f(s, i == j ? Data[i] : Zero); + } + target.At(j, finalize(s, RowCount)); + } + } + } } } diff --git a/src/Numerics/LinearAlgebra/Storage/MatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/MatrixStorage.cs index 1156a348..0cf4a15f 100644 --- a/src/Numerics/LinearAlgebra/Storage/MatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/MatrixStorage.cs @@ -539,7 +539,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } - // FUNCTIONAL COMBINATORS + // FUNCTIONAL COMBINATORS: MAP public virtual void MapInplace(Func f, Zeros zeros = Zeros.AllowSkip) { @@ -667,5 +667,83 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } } + + // FUNCTIONAL COMBINATORS: FOLD + + public void FoldRows(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + if (target == null) + { + throw new ArgumentNullException("target"); + } + if (target.Length != RowCount) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength, "target"); + } + + if (state == null) + { + throw new ArgumentNullException("state"); + } + if (state.Length != RowCount) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength, "state"); + } + + FoldRowsUnchecked(target, f, finalize, state, zeros); + } + + internal virtual void FoldRowsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + for (int i = 0; i < RowCount; i++) + { + TU s = state.At(i); + for (int j = 0; j < ColumnCount; j++) + { + s = f(s, At(i, j)); + } + target.At(i, finalize(s, ColumnCount)); + } + } + + public void FoldColumns(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + if (target == null) + { + throw new ArgumentNullException("target"); + } + if (target.Length != ColumnCount) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength, "target"); + } + + if (state == null) + { + throw new ArgumentNullException("state"); + } + if (state.Length != ColumnCount) + { + throw new ArgumentException(Resources.ArgumentVectorsSameLength, "state"); + } + + FoldColumnsUnchecked(target, f, finalize, state, zeros); + } + + internal virtual void FoldColumnsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + where TU : struct, IEquatable, IFormattable + { + for (int j = 0; j < ColumnCount; j++) + { + TU s = state.At(j); + for (int i = 0; i < RowCount; i++) + { + s = f(s, At(i, j)); + } + target.At(j, finalize(s, RowCount)); + } + } } } diff --git a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs index 00727315..7f328689 100644 --- a/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs +++ b/src/Numerics/LinearAlgebra/Storage/SparseCompressedRowMatrixStorage.cs @@ -31,7 +31,6 @@ using System; using System.Collections.Generic; using System.Linq; -using System.Security.Policy; using MathNet.Numerics.Properties; namespace MathNet.Numerics.LinearAlgebra.Storage @@ -1248,7 +1247,7 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } - // FUNCTIONAL COMBINATORS + // FUNCTIONAL COMBINATORS: MAP public override void MapInplace(Func f, Zeros zeros = Zeros.AllowSkip) { @@ -1425,9 +1424,9 @@ namespace MathNet.Numerics.LinearAlgebra.Storage var endIndex = RowPointers[row + 1]; for (int j = 0; j < ColumnCount; j++) { - if (j == ColumnIndices[index]) + if (index < endIndex && j == ColumnIndices[index]) { - target.At(row, j, f(Values[j])); + target.At(row, j, f(Values[index])); index = Math.Min(index + 1, endIndex); } else @@ -1520,9 +1519,9 @@ namespace MathNet.Numerics.LinearAlgebra.Storage var endIndex = RowPointers[row + 1]; for (int j = 0; j < ColumnCount; j++) { - if (j == ColumnIndices[index]) + if (index < endIndex && j == ColumnIndices[index]) { - target.At(row, j, f(row, j, Values[j])); + target.At(row, j, f(row, j, Values[index])); index = Math.Min(index + 1, endIndex); } else @@ -1755,5 +1754,108 @@ namespace MathNet.Numerics.LinearAlgebra.Storage } } } + + // FUNCTIONAL COMBINATORS: FOLD + + internal override void FoldRowsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + if (zeros == Zeros.AllowSkip) + { + for (int row = 0; row < RowCount; row++) + { + var startIndex = RowPointers[row]; + var endIndex = RowPointers[row + 1]; + TU s = state.At(row); + for (var j = startIndex; j < endIndex; j++) + { + s = f(s, Values[j]); + } + target.At(row, finalize(s, endIndex - startIndex)); + } + } + else + { + for (int row = 0; row < RowCount; row++) + { + var index = RowPointers[row]; + var endIndex = RowPointers[row + 1]; + TU s = state.At(row); + for (int j = 0; j < ColumnCount; j++) + { + if (index < endIndex && j == ColumnIndices[index]) + { + s = f(s, Values[index]); + index = Math.Min(index + 1, endIndex); + } + else + { + s = f(s, Zero); + } + } + target.At(row, finalize(s, ColumnCount)); + } + } + } + + internal override void FoldColumnsUnchecked(VectorStorage target, Func f, Func finalize, VectorStorage state, Zeros zeros = Zeros.AllowSkip) + { + var denseResult = target as DenseVectorStorage; + if (denseResult == null) + { + denseResult = new DenseVectorStorage(ColumnCount); + } + + state.CopyTo(denseResult); + TU[] result = denseResult.Data; + + if (zeros == Zeros.AllowSkip) + { + int[] count = new int[ColumnCount]; + for (int row = 0; row < RowCount; row++) + { + var startIndex = RowPointers[row]; + var endIndex = RowPointers[row + 1]; + for (var j = startIndex; j < endIndex; j++) + { + var column = ColumnIndices[j]; + result[column] = f(result[column], Values[j]); + count[column]++; + } + } + for (int j = 0; j < ColumnCount; j++) + { + result[j] = finalize(result[j], count[j]); + } + } + else + { + for (int row = 0; row < RowCount; row++) + { + var index = RowPointers[row]; + var endIndex = RowPointers[row + 1]; + for (int j = 0; j < ColumnCount; j++) + { + if (index < endIndex && j == ColumnIndices[index]) + { + result[j] = f(result[j], Values[index]); + index = Math.Min(index + 1, endIndex); + } + else + { + result[j] = f(result[j], Zero); + } + } + } + for (int j = 0; j < ColumnCount; j++) + { + result[j] = finalize(result[j], RowCount); + } + } + + if (!ReferenceEquals(denseResult, target)) + { + denseResult.CopyTo(target); + } + } } } diff --git a/src/UnitTests/GenericMath.cs b/src/UnitTests/GenericMath.cs new file mode 100644 index 00000000..746e2302 --- /dev/null +++ b/src/UnitTests/GenericMath.cs @@ -0,0 +1,763 @@ +// TAKEN FROM: +// Miscellaneous Utility Library +// http://www.yoda.arachsys.com/csharp/miscutil/ +// +// "Miscellaneous Utility Library" Software Licence +// +// Version 1.0 +// +// Copyright (c) 2004-2008 Jon Skeet and Marc Gravell. +// All rights reserved. +// +// Redistribution and use in source and binary forms, with or without +// modification, are permitted provided that the following conditions +// are met: +// +// 1. Redistributions of source code must retain the above copyright +// notice, this list of conditions and the following disclaimer. +// +// 2. Redistributions in binary form must reproduce the above copyright +// notice, this list of conditions and the following disclaimer in the +// documentation and/or other materials provided with the distribution. +// +// 3. The end-user documentation included with the redistribution, if +// any, must include the following acknowledgment: +// +// "This product includes software developed by Jon Skeet +// and Marc Gravell. Contact skeet@pobox.com, or see +// http://www.pobox.com/~skeet/)." +// +// Alternately, this acknowledgment may appear in the software itself, +// if and wherever such third-party acknowledgments normally appear. +// +// 4. The name "Miscellaneous Utility Library" must not be used to endorse +// or promote products derived from this software without prior written +// permission. For written permission, please contact skeet@pobox.com. +// +// 5. Products derived from this software may not be called +// "Miscellaneous Utility Library", nor may "Miscellaneous Utility Library" +// appear in their name, without prior written permission of Jon Skeet. +// +// THIS SOFTWARE IS PROVIDED "AS IS" AND ANY EXPRESSED OR IMPLIED +// WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF +// MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. +// IN NO EVENT SHALL JON SKEET BE LIABLE FOR ANY DIRECT, INDIRECT, +// INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, +// BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; +// LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +// CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT +// LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN +// ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +// POSSIBILITY OF SUCH DAMAGE. + +using System; +using System.Linq.Expressions; + +namespace MathNet.Numerics.UnitTests +{ + /// + /// The Operator class provides easy access to the standard operators + /// (addition, etc) for generic types, using type inference to simplify + /// usage. + /// + internal static class Operator + { + + /// + /// Indicates if the supplied value is non-null, + /// for reference-types or Nullable<T> + /// + /// True for non-null values, else false + public static bool HasValue(T value) + { + return Operator.NullOp.HasValue(value); + } + + /// + /// Increments the accumulator only + /// if the value is non-null. If the accumulator + /// is null, then the accumulator is given the new + /// value; otherwise the accumulator and value + /// are added. + /// + /// The current total to be incremented (can be null) + /// The value to be tested and added to the accumulator + /// True if the value is non-null, else false - i.e. + /// "has the accumulator been updated?" + public static bool AddIfNotNull(ref T accumulator, T value) + { + return Operator.NullOp.AddIfNotNull(ref accumulator, value); + } + + /// + /// Evaluates unary negation (-) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Negate(T value) + { + return Operator.Negate(value); + } + + /// + /// Evaluates bitwise not (~) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Not(T value) + { + return Operator.Not(value); + } + + /// + /// Evaluates bitwise or (|) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Or(T value1, T value2) + { + return Operator.Or(value1, value2); + } + + /// + /// Evaluates bitwise and (&) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T And(T value1, T value2) + { + return Operator.And(value1, value2); + } + + /// + /// Evaluates bitwise xor (^) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Xor(T value1, T value2) + { + return Operator.Xor(value1, value2); + } + + /// + /// Performs a conversion between the given types; this will throw + /// an InvalidOperationException if the type T does not provide a suitable cast, or for + /// Nullable<TInner> if TInner does not provide this cast. + /// + public static TTo Convert(TFrom value) + { + return Operator.Convert(value); + } + + /// + /// Evaluates binary addition (+) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Add(T value1, T value2) + { + return Operator.Add(value1, value2); + } + + /// + /// Evaluates binary addition (+) for the given type(s); this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static TArg1 AddAlternative(TArg1 value1, TArg2 value2) + { + return Operator.Add(value1, value2); + } + + /// + /// Evaluates binary subtraction (-) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Subtract(T value1, T value2) + { + return Operator.Subtract(value1, value2); + } + + /// + /// Evaluates binary subtraction(-) for the given type(s); this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static TArg1 SubtractAlternative(TArg1 value1, TArg2 value2) + { + return Operator.Subtract(value1, value2); + } + + /// + /// Evaluates binary multiplication (*) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Multiply(T value1, T value2) + { + return Operator.Multiply(value1, value2); + } + + /// + /// Evaluates binary multiplication (*) for the given type(s); this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static TArg1 MultiplyAlternative(TArg1 value1, TArg2 value2) + { + return Operator.Multiply(value1, value2); + } + + /// + /// Evaluates binary division (/) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static T Divide(T value1, T value2) + { + return Operator.Divide(value1, value2); + } + + /// + /// Evaluates binary division (/) for the given type(s); this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static TArg1 DivideAlternative(TArg1 value1, TArg2 value2) + { + return Operator.Divide(value1, value2); + } + + /// + /// Evaluates binary equality (==) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool Equal(T value1, T value2) + { + return Operator.Equal(value1, value2); + } + + /// + /// Evaluates binary inequality (!=) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool NotEqual(T value1, T value2) + { + return Operator.NotEqual(value1, value2); + } + + /// + /// Evaluates binary greater-than (>) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool GreaterThan(T value1, T value2) + { + return Operator.GreaterThan(value1, value2); + } + + /// + /// Evaluates binary less-than (<) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool LessThan(T value1, T value2) + { + return Operator.LessThan(value1, value2); + } + + /// + /// Evaluates binary greater-than-on-eqauls (>=) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool GreaterThanOrEqual(T value1, T value2) + { + return Operator.GreaterThanOrEqual(value1, value2); + } + + /// + /// Evaluates binary less-than-or-equal (<=) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static bool LessThanOrEqual(T value1, T value2) + { + return Operator.LessThanOrEqual(value1, value2); + } + + /// + /// Evaluates integer division (/) for the given type; this will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + /// This operation is particularly useful for computing averages and + /// similar aggregates. + /// + public static T DivideInt32(T value, int divisor) + { + return Operator.Divide(value, divisor); + } + } + + /// + /// Provides standard operators (such as addition) that operate over operands of + /// different types. For operators, the return type is assumed to match the first + /// operand. + /// + /// + /// + internal static class Operator + { + static readonly Func convert; + + /// + /// Returns a delegate to convert a value between two types; this delegate will throw + /// an InvalidOperationException if the type T does not provide a suitable cast, or for + /// Nullable<TInner> if TInner does not provide this cast. + /// + public static Func Convert + { + get { return convert; } + } + + static Operator() + { + convert = ExpressionUtil.CreateExpression(body => Expression.Convert(body, typeof (TResult))); + add = ExpressionUtil.CreateExpression(Expression.Add, true); + subtract = ExpressionUtil.CreateExpression(Expression.Subtract, true); + multiply = ExpressionUtil.CreateExpression(Expression.Multiply, true); + divide = ExpressionUtil.CreateExpression(Expression.Divide, true); + } + + static readonly Func add, subtract, multiply, divide; + + /// + /// Returns a delegate to evaluate binary addition (+) for the given types; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Add + { + get { return add; } + } + + /// + /// Returns a delegate to evaluate binary subtraction (-) for the given types; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Subtract + { + get { return subtract; } + } + + /// + /// Returns a delegate to evaluate binary multiplication (*) for the given types; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Multiply + { + get { return multiply; } + } + + /// + /// Returns a delegate to evaluate binary division (/) for the given types; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Divide + { + get { return divide; } + } + } + + /// + /// Provides standard operators (such as addition) over a single type + /// + /// + /// + internal static class Operator + { + static readonly INullOp nullOp; + + internal static INullOp NullOp + { + get { return nullOp; } + } + + static readonly T zero; + + /// + /// Returns the zero value for value-types (even full Nullable<TInner>) - or null for reference types + /// + public static T Zero + { + get { return zero; } + } + + static readonly Func negate, not; + static readonly Func or, and, xor; + + /// + /// Returns a delegate to evaluate unary negation (-) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Negate + { + get { return negate; } + } + + /// + /// Returns a delegate to evaluate bitwise not (~) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Not + { + get { return not; } + } + + /// + /// Returns a delegate to evaluate bitwise or (|) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Or + { + get { return or; } + } + + /// + /// Returns a delegate to evaluate bitwise and (&) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func And + { + get { return and; } + } + + /// + /// Returns a delegate to evaluate bitwise xor (^) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Xor + { + get { return xor; } + } + + static readonly Func add, subtract, multiply, divide; + + /// + /// Returns a delegate to evaluate binary addition (+) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Add + { + get { return add; } + } + + /// + /// Returns a delegate to evaluate binary subtraction (-) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Subtract + { + get { return subtract; } + } + + /// + /// Returns a delegate to evaluate binary multiplication (*) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Multiply + { + get { return multiply; } + } + + /// + /// Returns a delegate to evaluate binary division (/) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Divide + { + get { return divide; } + } + + + static readonly Func equal, notEqual, greaterThan, lessThan, greaterThanOrEqual, lessThanOrEqual; + + /// + /// Returns a delegate to evaluate binary equality (==) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func Equal + { + get { return equal; } + } + + /// + /// Returns a delegate to evaluate binary inequality (!=) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func NotEqual + { + get { return notEqual; } + } + + /// + /// Returns a delegate to evaluate binary greater-then (>) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func GreaterThan + { + get { return greaterThan; } + } + + /// + /// Returns a delegate to evaluate binary less-than (<) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func LessThan + { + get { return lessThan; } + } + + /// + /// Returns a delegate to evaluate binary greater-than-or-equal (>=) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func GreaterThanOrEqual + { + get { return greaterThanOrEqual; } + } + + /// + /// Returns a delegate to evaluate binary less-than-or-equal (<=) for the given type; this delegate will throw + /// an InvalidOperationException if the type T does not provide this operator, or for + /// Nullable<TInner> if TInner does not provide this operator. + /// + public static Func LessThanOrEqual + { + get { return lessThanOrEqual; } + } + + static Operator() + { + add = ExpressionUtil.CreateExpression(Expression.Add); + subtract = ExpressionUtil.CreateExpression(Expression.Subtract); + divide = ExpressionUtil.CreateExpression(Expression.Divide); + multiply = ExpressionUtil.CreateExpression(Expression.Multiply); + + greaterThan = ExpressionUtil.CreateExpression(Expression.GreaterThan); + greaterThanOrEqual = ExpressionUtil.CreateExpression(Expression.GreaterThanOrEqual); + lessThan = ExpressionUtil.CreateExpression(Expression.LessThan); + lessThanOrEqual = ExpressionUtil.CreateExpression(Expression.LessThanOrEqual); + equal = ExpressionUtil.CreateExpression(Expression.Equal); + notEqual = ExpressionUtil.CreateExpression(Expression.NotEqual); + + negate = ExpressionUtil.CreateExpression(Expression.Negate); + and = ExpressionUtil.CreateExpression(Expression.And); + or = ExpressionUtil.CreateExpression(Expression.Or); + not = ExpressionUtil.CreateExpression(Expression.Not); + xor = ExpressionUtil.CreateExpression(Expression.ExclusiveOr); + + Type typeT = typeof (T); + if (typeT.IsValueType && typeT.IsGenericType && (typeT.GetGenericTypeDefinition() == typeof (Nullable<>))) + { + // get the *inner* zero (not a null Nullable, but default(TValue)) + Type nullType = typeT.GetGenericArguments()[0]; + zero = (T)Activator.CreateInstance(nullType); + nullOp = (INullOp)Activator.CreateInstance( + typeof (StructNullOp<>).MakeGenericType(nullType)); + } + else + { + zero = default(T); + if (typeT.IsValueType) + { + nullOp = (INullOp)Activator.CreateInstance( + typeof (StructNullOp<>).MakeGenericType(typeT)); + } + else + { + nullOp = (INullOp)Activator.CreateInstance( + typeof (ClassNullOp<>).MakeGenericType(typeT)); + } + } + } + } + + /// + /// General purpose Expression utilities + /// + internal static class ExpressionUtil + { + /// + /// Create a function delegate representing a unary operation + /// + /// The parameter type + /// The return type + /// Body factory + /// Compiled function delegate + public static Func CreateExpression( + Func body) + { + ParameterExpression inp = Expression.Parameter(typeof (TArg1), "inp"); + try + { + return Expression.Lambda>(body(inp), inp).Compile(); + } + catch (Exception ex) + { + string msg = ex.Message; // avoid capture of ex itself + return delegate { throw new InvalidOperationException(msg); }; + } + } + + /// + /// Create a function delegate representing a binary operation + /// + /// The first parameter type + /// The second parameter type + /// The return type + /// Body factory + /// Compiled function delegate + public static Func CreateExpression( + Func body) + { + return CreateExpression(body, false); + } + + /// + /// Create a function delegate representing a binary operation + /// + /// + /// If no matching operation is possible, attempt to convert + /// TArg1 and TArg2 to TResult for a match? For example, there is no + /// "decimal operator /(decimal, int)", but by converting TArg2 (int) to + /// TResult (decimal) a match is found. + /// + /// The first parameter type + /// The second parameter type + /// The return type + /// Body factory + /// Compiled function delegate + public static Func CreateExpression( + Func body, bool castArgsToResultOnFailure) + { + ParameterExpression lhs = Expression.Parameter(typeof (TArg1), "lhs"); + ParameterExpression rhs = Expression.Parameter(typeof (TArg2), "rhs"); + try + { + try + { + return Expression.Lambda>(body(lhs, rhs), lhs, rhs).Compile(); + } + catch (InvalidOperationException) + { + if (castArgsToResultOnFailure && !( // if we show retry + typeof (TArg1) == typeof (TResult) && // and the args aren't + typeof (TArg2) == typeof (TResult))) + { + // already "TValue, TValue, TValue"... + // convert both lhs and rhs to TResult (as appropriate) + Expression castLhs = typeof (TArg1) == typeof (TResult) ? + (Expression)lhs : + (Expression)Expression.Convert(lhs, typeof (TResult)); + Expression castRhs = typeof (TArg2) == typeof (TResult) ? + (Expression)rhs : + (Expression)Expression.Convert(rhs, typeof (TResult)); + + return Expression.Lambda>( + body(castLhs, castRhs), lhs, rhs).Compile(); + } + else throw; + } + } + catch (Exception ex) + { + string msg = ex.Message; // avoid capture of ex itself + return delegate { throw new InvalidOperationException(msg); }; + } + } + } + + internal interface INullOp + { + bool HasValue(T value); + bool AddIfNotNull(ref T accumulator, T value); + } + + internal sealed class StructNullOp + : INullOp, INullOp + where T : struct + { + public bool HasValue(T value) + { + return true; + } + + public bool AddIfNotNull(ref T accumulator, T value) + { + accumulator = Operator.Add(accumulator, value); + return true; + } + + public bool HasValue(T? value) + { + return value.HasValue; + } + + public bool AddIfNotNull(ref T? accumulator, T? value) + { + if (value.HasValue) + { + accumulator = accumulator.HasValue ? + Operator.Add( + accumulator.GetValueOrDefault(), + value.GetValueOrDefault()) + : value; + return true; + } + return false; + } + } + + internal sealed class ClassNullOp + : INullOp + where T : class + { + public bool HasValue(T value) + { + return value != null; + } + + public bool AddIfNotNull(ref T accumulator, T value) + { + if (value != null) + { + accumulator = accumulator == null ? + value : Operator.Add(accumulator, value); + return true; + } + return false; + } + } +} diff --git a/src/UnitTests/MatrixHelpers.cs b/src/UnitTests/LinearAlgebraTests/MatrixHelpers.cs similarity index 100% rename from src/UnitTests/MatrixHelpers.cs rename to src/UnitTests/LinearAlgebraTests/MatrixHelpers.cs diff --git a/src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Map.cs b/src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Functional.cs similarity index 87% rename from src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Map.cs rename to src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Functional.cs index 99f70a40..73e9a732 100644 --- a/src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Map.cs +++ b/src/UnitTests/LinearAlgebraTests/MatrixStructureTheory.Functional.cs @@ -1,4 +1,4 @@ -// +// // Math.NET Numerics, part of the Math.NET Project // http://numerics.mathdotnet.com // http://github.com/mathnet/mathnet-numerics @@ -256,5 +256,47 @@ namespace MathNet.Numerics.UnitTests.LinearAlgebraTests sparse.SetSubMatrix(1, 0, matrix.RowCount - 1, 1, 0, matrix.ColumnCount, Matrix.Build.Dense(matrix.RowCount - 1, matrix.ColumnCount, one)); Assert.That(sparse.Enumerate().All(one.Equals), Is.True); } + + [Theory] + public void CanFoldRows(Matrix matrix) + { + // not forced + var rowSum = matrix.FoldRows((s, x) => Operator.Add(s, x), Operator.Zero, Zeros.AllowSkip); + for (int i = 0; i < rowSum.Count; i++) + { + Assert.That(rowSum.At(i), Is.EqualTo(matrix.Row(i).Enumerate().Aggregate((a, b) => Operator.Add(a, b))), "not forced"); + } + + // forced + rowSum = matrix.FoldRows((s, x) => Operator.Add(s, x), Operator.Zero, Zeros.Include); + for (int i = 0; i < rowSum.Count; i++) + { + Assert.That(rowSum.At(i), Is.EqualTo(matrix.Row(i).Enumerate().Aggregate((a, b) => Operator.Add(a, b))), "forced"); + } + + Assert.That(matrix.FoldRows((s, x) => s + 1.0, 0.0, Zeros.Include), + Is.EqualTo(Vector.Build.Dense(matrix.RowCount, matrix.ColumnCount)), "forced - full coverage"); + } + + [Theory] + public void CanFoldColumns(Matrix matrix) + { + // not forced + var colSum = matrix.FoldColumns((s, x) => Operator.Add(s, x), Operator.Zero, Zeros.AllowSkip); + for (int i = 0; i < colSum.Count; i++) + { + Assert.That(colSum.At(i), Is.EqualTo(matrix.Column(i).Enumerate().Aggregate((a, b) => Operator.Add(a, b))), "not forced"); + } + + // forced + colSum = matrix.FoldColumns((s, x) => Operator.Add(s, x), Operator.Zero, Zeros.Include); + for (int i = 0; i < colSum.Count; i++) + { + Assert.That(colSum.At(i), Is.EqualTo(matrix.Column(i).Enumerate().Aggregate((a, b) => Operator.Add(a, b))), "forced"); + } + + Assert.That(matrix.FoldColumns((s, x) => s + 1.0, 0.0, Zeros.Include), + Is.EqualTo(Vector.Build.Dense(matrix.ColumnCount, matrix.RowCount)), "forced - full coverage"); + } } } diff --git a/src/UnitTests/UnitTests.csproj b/src/UnitTests/UnitTests.csproj index fc3e9f17..c131a104 100644 --- a/src/UnitTests/UnitTests.csproj +++ b/src/UnitTests/UnitTests.csproj @@ -303,9 +303,10 @@ + - + @@ -358,7 +359,7 @@ - +