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 @@
-
+