Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,72 @@ private interface IAggregationOperator<T> : IBinaryOperator<T>
static abstract T Invoke(Vector512<T> x);

static virtual T IdentityValue => throw new NotSupportedException();

/// <summary>Gets whether combining values in a different order produces the same result.</summary>
static virtual bool CanReassociate => false;
}

/// <summary>Aggregates four vectors into <paramref name="vresult"/>, as a tree when the operator allows it.</summary>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static Vector128<T> AggregateFour<T, TAggregationOperator>(
Vector128<T> vresult, Vector128<T> vector1, Vector128<T> vector2, Vector128<T> vector3, Vector128<T> vector4)
where TAggregationOperator : struct, IAggregationOperator<T>
{
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);
}

/// <inheritdoc cref="AggregateFour{T, TAggregationOperator}(Vector128{T}, Vector128{T}, Vector128{T}, Vector128{T}, Vector128{T})"/>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static Vector256<T> AggregateFour<T, TAggregationOperator>(
Vector256<T> vresult, Vector256<T> vector1, Vector256<T> vector2, Vector256<T> vector3, Vector256<T> vector4)
where TAggregationOperator : struct, IAggregationOperator<T>
{
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);
}

/// <inheritdoc cref="AggregateFour{T, TAggregationOperator}(Vector128{T}, Vector128{T}, Vector128{T}, Vector128{T}, Vector128{T})"/>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
private static Vector512<T> AggregateFour<T, TAggregationOperator>(
Vector512<T> vresult, Vector512<T> vector1, Vector512<T> vector2, Vector512<T> vector3, Vector512<T> vector4)
where TAggregationOperator : struct, IAggregationOperator<T>
{
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);
}

/// <summary>Adapts a stateless <see cref="IUnaryOperator{TInput, TOutput}"/> to be used as a stateful <see cref="IStatefulUnaryOperator{T}"/>.</summary>
Expand Down Expand Up @@ -218,10 +284,7 @@ static T Vectorized128(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 2)));
vector4 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -230,10 +293,7 @@ static T Vectorized128(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 6)));
vector4 = transform.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down Expand Up @@ -400,10 +460,7 @@ static T Vectorized256(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 2)));
vector4 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -412,10 +469,7 @@ static T Vectorized256(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 6)));
vector4 = transform.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down Expand Up @@ -582,10 +636,7 @@ static T Vectorized512(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 2)));
vector4 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -594,10 +645,7 @@ static T Vectorized512(ref T xRef, nuint remainder, TTransformOperator transform
vector3 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 6)));
vector4 = transform.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down Expand Up @@ -1351,10 +1399,7 @@ static T Vectorized128(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 3)),
Vector128.Load(yPtr + (uint)(Vector128<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -1367,10 +1412,7 @@ static T Vectorized128(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector128.Load(xPtr + (uint)(Vector128<T>.Count * 7)),
Vector128.Load(yPtr + (uint)(Vector128<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down Expand Up @@ -1558,10 +1600,7 @@ static T Vectorized256(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 3)),
Vector256.Load(yPtr + (uint)(Vector256<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -1574,10 +1613,7 @@ static T Vectorized256(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector256.Load(xPtr + (uint)(Vector256<T>.Count * 7)),
Vector256.Load(yPtr + (uint)(Vector256<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down Expand Up @@ -1765,10 +1801,7 @@ static T Vectorized512(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 3)),
Vector512.Load(yPtr + (uint)(Vector512<T>.Count * 3)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We load, process, and store the next four vectors

Expand All @@ -1781,10 +1814,7 @@ static T Vectorized512(ref T xRef, ref T yRef, nuint remainder)
vector4 = TBinaryOperator.Invoke(Vector512.Load(xPtr + (uint)(Vector512<T>.Count * 7)),
Vector512.Load(yPtr + (uint)(Vector512<T>.Count * 7)));

vresult = TAggregationOperator.Invoke(vresult, vector1);
vresult = TAggregationOperator.Invoke(vresult, vector2);
vresult = TAggregationOperator.Invoke(vresult, vector3);
vresult = TAggregationOperator.Invoke(vresult, vector4);
vresult = AggregateFour<T, TAggregationOperator>(vresult, vector1, vector2, vector3, vector4);

// We adjust the source and destination references, then update
// the count of remaining elements to process.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,8 @@ public static void Add<T>(ReadOnlySpan<T> x, T y, Span<T> 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<T> Invoke(Vector128<T> x, Vector128<T> y) => x + y;
public static Vector256<T> Invoke(Vector256<T> x, Vector256<T> y) => x + y;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,8 @@ public static void Multiply<T>(ReadOnlySpan<T> x, T y, Span<T> 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<T> Invoke(Vector128<T> x, Vector128<T> y) => x * y;
public static Vector256<T> Invoke(Vector256<T> x, Vector256<T> y) => x * y;
Expand Down
3 changes: 3 additions & 0 deletions src/libraries/System.Numerics.Tensors/tests/Helpers.cs
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,9 @@ public static class Helpers
public static IEnumerable<int> TensorLengthsIncluding0 => Enumerable.Range(0, 257);

public static IEnumerable<int> TensorLengths => Enumerable.Range(1, 256);

public static IEnumerable<int> 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<nint[]> 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]];
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1842,6 +1842,42 @@ public void IndexOfMinMagnitude_HandlesMinValue()
public unsafe abstract class GenericIntegerTensorPrimitivesTests<T> : GenericNumberTensorPrimitivesTests<T>
where T : unmanaged, IBinaryInteger<T>, IMinMaxValue<T>
{
#region Sum
[Fact]
public void Sum_MatchesScalarSumAcrossUnrolledBlocks()
{
Assert.All(Helpers.TensorLengthsSpanningUnrolledBlocks, tensorLength =>
{
using BoundedMemory<T> 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<T> 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()
Expand Down
Loading