diff --git a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/Common/TensorPrimitives.IAggregationOperator.cs b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/Common/TensorPrimitives.IAggregationOperator.cs index f3115744143737..45a8745ad65504 100644 --- a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/Common/TensorPrimitives.IAggregationOperator.cs +++ b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/Common/TensorPrimitives.IAggregationOperator.cs @@ -20,6 +20,72 @@ private interface IAggregationOperator : IBinaryOperator static abstract T Invoke(Vector512 x); static virtual T IdentityValue => throw new NotSupportedException(); + + /// Gets whether combining values in a different order produces the same result. + static virtual bool CanReassociate => false; + } + + /// Aggregates four vectors into , as a tree when the operator allows it. + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static Vector128 AggregateFour( + Vector128 vresult, Vector128 vector1, Vector128 vector2, Vector128 vector3, Vector128 vector4) + where TAggregationOperator : struct, IAggregationOperator + { + if (TAggregationOperator.CanReassociate) + { + return TAggregationOperator.Invoke( + vresult, + TAggregationOperator.Invoke( + TAggregationOperator.Invoke(vector1, vector2), + TAggregationOperator.Invoke(vector3, vector4))); + } + + vresult = TAggregationOperator.Invoke(vresult, vector1); + vresult = TAggregationOperator.Invoke(vresult, vector2); + vresult = TAggregationOperator.Invoke(vresult, vector3); + return TAggregationOperator.Invoke(vresult, vector4); + } + + /// + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static Vector256 AggregateFour( + Vector256 vresult, Vector256 vector1, Vector256 vector2, Vector256 vector3, Vector256 vector4) + where TAggregationOperator : struct, IAggregationOperator + { + if (TAggregationOperator.CanReassociate) + { + return TAggregationOperator.Invoke( + vresult, + TAggregationOperator.Invoke( + TAggregationOperator.Invoke(vector1, vector2), + TAggregationOperator.Invoke(vector3, vector4))); + } + + vresult = TAggregationOperator.Invoke(vresult, vector1); + vresult = TAggregationOperator.Invoke(vresult, vector2); + vresult = TAggregationOperator.Invoke(vresult, vector3); + return TAggregationOperator.Invoke(vresult, vector4); + } + + /// + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static Vector512 AggregateFour( + Vector512 vresult, Vector512 vector1, Vector512 vector2, Vector512 vector3, Vector512 vector4) + where TAggregationOperator : struct, IAggregationOperator + { + if (TAggregationOperator.CanReassociate) + { + return TAggregationOperator.Invoke( + vresult, + TAggregationOperator.Invoke( + TAggregationOperator.Invoke(vector1, vector2), + TAggregationOperator.Invoke(vector3, vector4))); + } + + vresult = TAggregationOperator.Invoke(vresult, vector1); + vresult = TAggregationOperator.Invoke(vresult, vector2); + vresult = TAggregationOperator.Invoke(vresult, vector3); + return TAggregationOperator.Invoke(vresult, vector4); } /// Adapts a stateless to be used as a stateful . @@ -218,10 +284,7 @@ static T Vectorized128(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 2))); vector4 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -230,10 +293,7 @@ static T Vectorized128(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 6))); vector4 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. @@ -400,10 +460,7 @@ static T Vectorized256(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 2))); vector4 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -412,10 +469,7 @@ static T Vectorized256(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 6))); vector4 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. @@ -582,10 +636,7 @@ static T Vectorized512(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 2))); vector4 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -594,10 +645,7 @@ static T Vectorized512(ref T xRef, nuint remainder, TTransformOperator transform vector3 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 6))); vector4 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. @@ -1351,10 +1399,7 @@ static T Vectorized128(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 3)), Vector128.Load(yPtr + (uint)(Vector128.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -1367,10 +1412,7 @@ static T Vectorized128(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector128.Load(xPtr + (uint)(Vector128.Count * 7)), Vector128.Load(yPtr + (uint)(Vector128.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. @@ -1558,10 +1600,7 @@ static T Vectorized256(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 3)), Vector256.Load(yPtr + (uint)(Vector256.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -1574,10 +1613,7 @@ static T Vectorized256(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector256.Load(xPtr + (uint)(Vector256.Count * 7)), Vector256.Load(yPtr + (uint)(Vector256.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. @@ -1765,10 +1801,7 @@ static T Vectorized512(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 3)), Vector512.Load(yPtr + (uint)(Vector512.Count * 3))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We load, process, and store the next four vectors @@ -1781,10 +1814,7 @@ static T Vectorized512(ref T xRef, ref T yRef, nuint remainder) vector4 = TBinaryOperator.Invoke(Vector512.Load(xPtr + (uint)(Vector512.Count * 7)), Vector512.Load(yPtr + (uint)(Vector512.Count * 7))); - vresult = TAggregationOperator.Invoke(vresult, vector1); - vresult = TAggregationOperator.Invoke(vresult, vector2); - vresult = TAggregationOperator.Invoke(vresult, vector3); - vresult = TAggregationOperator.Invoke(vresult, vector4); + vresult = AggregateFour(vresult, vector1, vector2, vector3, vector4); // We adjust the source and destination references, then update // the count of remaining elements to process. diff --git a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Add.cs b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Add.cs index 74dfe0030d7efb..8c3c7c0fcd3cc4 100644 --- a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Add.cs +++ b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Add.cs @@ -64,6 +64,8 @@ public static void Add(ReadOnlySpan x, T y, Span destination) { public static bool Vectorizable => true; + public static bool CanReassociate => typeof(T) != typeof(float) && typeof(T) != typeof(double); + public static T Invoke(T x, T y) => x + y; public static Vector128 Invoke(Vector128 x, Vector128 y) => x + y; public static Vector256 Invoke(Vector256 x, Vector256 y) => x + y; diff --git a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Multiply.cs b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Multiply.cs index 937bf8ee164d53..eb55bfdc6bcb8f 100644 --- a/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Multiply.cs +++ b/src/libraries/System.Numerics.Tensors/src/System/Numerics/Tensors/netcore/TensorPrimitives.Multiply.cs @@ -65,6 +65,8 @@ public static void Multiply(ReadOnlySpan x, T y, Span destination) { public static bool Vectorizable => true; + public static bool CanReassociate => typeof(T) != typeof(float) && typeof(T) != typeof(double); + public static T Invoke(T x, T y) => x * y; public static Vector128 Invoke(Vector128 x, Vector128 y) => x * y; public static Vector256 Invoke(Vector256 x, Vector256 y) => x * y; diff --git a/src/libraries/System.Numerics.Tensors/tests/Helpers.cs b/src/libraries/System.Numerics.Tensors/tests/Helpers.cs index d3ebdecefa92d9..5fa90fc4624111 100644 --- a/src/libraries/System.Numerics.Tensors/tests/Helpers.cs +++ b/src/libraries/System.Numerics.Tensors/tests/Helpers.cs @@ -17,6 +17,9 @@ public static class Helpers public static IEnumerable TensorLengthsIncluding0 => Enumerable.Range(0, 257); public static IEnumerable TensorLengths => Enumerable.Range(1, 256); + + public static IEnumerable TensorLengthsSpanningUnrolledBlocks => + [1, 3, 8, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 128, 129, 255, 256, 257, 511, 512, 513, 1023, 1024, 1025, 4099]; public static IEnumerable TensorShapes => [[1], [2], [10], [1, 1], [1, 2], [2, 2], [5, 5], [2, 2, 2], [5, 5, 5], [3, 3, 3, 3], [4, 4, 4, 4, 4], [1, 2, 3, 4, 5, 6, 7, 1, 2]]; public static nint[][] TensorSliceShapes => [[1], [1], [5], [1, 1], [1, 1], [1, 2], [3, 3], [2, 2, 1], [5, 3, 5], [3, 2, 1, 3], [4, 3, 2, 1, 2], [1, 2, 2, 2, 2, 1, 1, 1, 1]]; public static nint[][] TensorSliceShapesForBroadcast => [[1], [1], [1], [1, 1], [1, 1], [1, 2], [1, 1], [2, 2, 1], [1, 5, 5], [3, 1, 1, 3], [4, 1, 4, 1, 4], [1, 2, 1, 4, 1, 1, 7, 1, 1]]; diff --git a/src/libraries/System.Numerics.Tensors/tests/TensorPrimitives.Generic.cs b/src/libraries/System.Numerics.Tensors/tests/TensorPrimitives.Generic.cs index 295955d9497ea6..12a00a0cb2fb17 100644 --- a/src/libraries/System.Numerics.Tensors/tests/TensorPrimitives.Generic.cs +++ b/src/libraries/System.Numerics.Tensors/tests/TensorPrimitives.Generic.cs @@ -1842,6 +1842,42 @@ public void IndexOfMinMagnitude_HandlesMinValue() public unsafe abstract class GenericIntegerTensorPrimitivesTests : GenericNumberTensorPrimitivesTests where T : unmanaged, IBinaryInteger, IMinMaxValue { + #region Sum + [Fact] + public void Sum_MatchesScalarSumAcrossUnrolledBlocks() + { + Assert.All(Helpers.TensorLengthsSpanningUnrolledBlocks, tensorLength => + { + using BoundedMemory x = CreateAndFillTensor(tensorLength); + + T expected = T.Zero; + foreach (T value in x.Span) + { + expected += value; + } + + Assert.Equal(expected, Sum(x.Span)); + }); + } + + [Fact] + public void SumOfSquares_MatchesScalarSumAcrossUnrolledBlocks() + { + Assert.All(Helpers.TensorLengthsSpanningUnrolledBlocks, tensorLength => + { + using BoundedMemory x = CreateAndFillTensor(tensorLength); + + T expected = T.Zero; + foreach (T value in x.Span) + { + expected += value * value; + } + + Assert.Equal(expected, SumOfSquares(x.Span)); + }); + } + #endregion + #region Divide [Fact] public void Divide_TwoTensors_ByZero_Throws()