Browse Source

Add rectangular AV1 motion-error metrics

Extend the shared residual operators with rectangular SAD and signed/squared moments, including alternate-row sampling and SIMD tails. Keep raw sample precision until the search controller normalizes the totals.

Verified in the preserved worktree with Release .NET 11 and serialized Visual Studio VSTest: 329 focused tests passed, including independent scalar comparisons across hardware widths. This is a motion-search dependency checkpoint; integrated encoder parity remains unproven.
pull/2633/head
James Jackson-South 4 weeks ago
parent
commit
e9babd55a3
  1. 300
      src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs
  2. 135
      tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs

300
src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs

@ -17,6 +17,306 @@ internal static partial class Av1ResidualBuilder
/// </summary> /// </summary>
private const int SearchBlockDimension = 8; private const int SearchBlockDimension = 8;
/// <summary>
/// Measures absolute prediction error over a rectangular block, optionally sampling alternate rows.
/// </summary>
/// <param name="source">Source samples beginning at the block origin.</param>
/// <param name="sourceStride">The source row stride, in samples.</param>
/// <param name="prediction">Prediction samples beginning at the block origin.</param>
/// <param name="predictionStride">The prediction row stride, in samples.</param>
/// <param name="width">The block width, in samples.</param>
/// <param name="height">The block height, in samples.</param>
/// <param name="rowStep">One for every row, or two for alternating rows with doubled error.</param>
/// <returns>The unnormalized absolute difference over the block.</returns>
public static int SumAbsoluteDifferences(
ReadOnlySpan<byte> source,
int sourceStride,
ReadOnlySpan<byte> prediction,
int predictionStride,
int width,
int height,
int rowStep)
=> SumAbsoluteDifferences<byte, ByteOperator>(source, sourceStride, prediction, predictionStride, width, height, rowStep);
/// <summary>
/// Measures signed residual sum and squared error over a rectangular prediction block.
/// </summary>
/// <param name="source">Source samples beginning at the block origin.</param>
/// <param name="sourceStride">The source row stride, in samples.</param>
/// <param name="prediction">Prediction samples beginning at the block origin.</param>
/// <param name="predictionStride">The prediction row stride, in samples.</param>
/// <param name="width">The block width, in samples.</param>
/// <param name="height">The block height, in samples.</param>
/// <param name="sum">The unnormalized signed residual sum.</param>
/// <param name="sumOfSquares">The unnormalized squared residual sum.</param>
public static void GetMoments(
ReadOnlySpan<byte> source,
int sourceStride,
ReadOnlySpan<byte> prediction,
int predictionStride,
int width,
int height,
out int sum,
out long sumOfSquares)
=> GetMoments<byte, ByteOperator>(source, sourceStride, prediction, predictionStride, width, height, out sum, out sumOfSquares);
/// <summary>
/// Measures absolute prediction error over a rectangular block, optionally sampling alternate rows.
/// </summary>
/// <param name="source">Source samples beginning at the block origin.</param>
/// <param name="sourceStride">The source row stride, in samples.</param>
/// <param name="prediction">Prediction samples beginning at the block origin.</param>
/// <param name="predictionStride">The prediction row stride, in samples.</param>
/// <param name="width">The block width, in samples.</param>
/// <param name="height">The block height, in samples.</param>
/// <param name="rowStep">One for every row, or two for alternating rows with doubled error.</param>
/// <returns>The unnormalized absolute difference over the block.</returns>
public static int SumAbsoluteDifferences(
ReadOnlySpan<ushort> source,
int sourceStride,
ReadOnlySpan<ushort> prediction,
int predictionStride,
int width,
int height,
int rowStep)
=> SumAbsoluteDifferences<ushort, UInt16Operator>(source, sourceStride, prediction, predictionStride, width, height, rowStep);
/// <summary>
/// Measures signed residual sum and squared error over a rectangular prediction block.
/// </summary>
/// <param name="source">Source samples beginning at the block origin.</param>
/// <param name="sourceStride">The source row stride, in samples.</param>
/// <param name="prediction">Prediction samples beginning at the block origin.</param>
/// <param name="predictionStride">The prediction row stride, in samples.</param>
/// <param name="width">The block width, in samples.</param>
/// <param name="height">The block height, in samples.</param>
/// <param name="sum">The unnormalized signed residual sum.</param>
/// <param name="sumOfSquares">The unnormalized squared residual sum.</param>
public static void GetMoments(
ReadOnlySpan<ushort> source,
int sourceStride,
ReadOnlySpan<ushort> prediction,
int predictionStride,
int width,
int height,
out int sum,
out long sumOfSquares)
=> GetMoments<ushort, UInt16Operator>(source, sourceStride, prediction, predictionStride, width, height, out sum, out sumOfSquares);
/// <summary>
/// Traverses rectangular SAD candidates using the selected sample operator and descending vector widths.
/// </summary>
private static int SumAbsoluteDifferences<TSample, TOperator>(
ReadOnlySpan<TSample> source,
int sourceStride,
ReadOnlySpan<TSample> prediction,
int predictionStride,
int width,
int height,
int rowStep)
where TSample : unmanaged
where TOperator : struct, IResidualOperator<TSample>
{
int sum = 0;
// Wider blocks use complete native loads. The eight-sample tail retains the existing compact load,
// which reads eight bytes or eight words without crossing a short row's boundary.
for (int y = 0; y < height; y += rowStep)
{
ReadOnlySpan<TSample> sourceRow = source.Slice(y * sourceStride, width);
ReadOnlySpan<TSample> predictionRow = prediction.Slice(y * predictionStride, width);
ref TSample sourceBase = ref MemoryMarshal.GetReference(sourceRow);
ref TSample predictionBase = ref MemoryMarshal.GetReference(predictionRow);
int x = 0;
if (Vector512.IsHardwareAccelerated)
{
for (; x <= width - Vector512<TSample>.Count; x += Vector512<TSample>.Count)
{
Vector512<TSample> sourceVector = Vector512.LoadUnsafe(ref sourceBase, (nuint)x);
Vector512<TSample> predictionVector = Vector512.LoadUnsafe(ref predictionBase, (nuint)x);
Vector512<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector512<short> upper);
// Absolute residuals fit short, but their horizontal sum may not. Widen before reducing.
Vector512<short> absolute = Vector512.Abs(lower);
sum += Vector512.Sum(Vector512.WidenLower(absolute)) + Vector512.Sum(Vector512.WidenUpper(absolute));
if (Vector512<TSample>.Count != Vector512<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
absolute = Vector512.Abs(upper);
sum += Vector512.Sum(Vector512.WidenLower(absolute)) + Vector512.Sum(Vector512.WidenUpper(absolute));
}
}
}
if (Vector256.IsHardwareAccelerated)
{
for (; x <= width - Vector256<TSample>.Count; x += Vector256<TSample>.Count)
{
Vector256<TSample> sourceVector = Vector256.LoadUnsafe(ref sourceBase, (nuint)x);
Vector256<TSample> predictionVector = Vector256.LoadUnsafe(ref predictionBase, (nuint)x);
Vector256<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector256<short> upper);
// Absolute residuals fit short, but their horizontal sum may not. Widen before reducing.
Vector256<short> absolute = Vector256.Abs(lower);
sum += Vector256.Sum(Vector256.WidenLower(absolute)) + Vector256.Sum(Vector256.WidenUpper(absolute));
if (Vector256<TSample>.Count != Vector256<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
absolute = Vector256.Abs(upper);
sum += Vector256.Sum(Vector256.WidenLower(absolute)) + Vector256.Sum(Vector256.WidenUpper(absolute));
}
}
}
if (Vector128.IsHardwareAccelerated)
{
for (; x <= width - Vector128<TSample>.Count; x += Vector128<TSample>.Count)
{
Vector128<TSample> sourceVector = Vector128.LoadUnsafe(ref sourceBase, (nuint)x);
Vector128<TSample> predictionVector = Vector128.LoadUnsafe(ref predictionBase, (nuint)x);
Vector128<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector128<short> upper);
// Absolute residuals fit short, but their horizontal sum may not. Widen before reducing.
Vector128<short> absolute = Vector128.Abs(lower);
sum += Vector128.Sum(Vector128.WidenLower(absolute)) + Vector128.Sum(Vector128.WidenUpper(absolute));
if (Vector128<TSample>.Count != Vector128<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
absolute = Vector128.Abs(upper);
sum += Vector128.Sum(Vector128.WidenLower(absolute)) + Vector128.Sum(Vector128.WidenUpper(absolute));
}
}
}
if (Vector128.IsHardwareAccelerated && x <= width - SearchBlockDimension)
{
Vector128<TSample> sourceVector = LoadSearchRow(sourceRow[x..]);
Vector128<TSample> predictionVector = LoadSearchRow(predictionRow[x..]);
sum += TOperator.SumAbsoluteDifferences(sourceVector, predictionVector);
x += SearchBlockDimension;
}
for (; x < width; x++)
{
sum += TOperator.SumAbsoluteDifferences(sourceRow[x], predictionRow[x]);
}
}
// Alternate-row search represents the complete even-height block by doubling the sampled row total.
// Precision normalization follows this scaling so fractional error units are truncated only once.
return sum * rowStep;
}
/// <summary>
/// Accumulates rectangular residual moments without storing an intermediate residual plane.
/// </summary>
private static void GetMoments<TSample, TOperator>(
ReadOnlySpan<TSample> source,
int sourceStride,
ReadOnlySpan<TSample> prediction,
int predictionStride,
int width,
int height,
out int sum,
out long sumOfSquares)
where TSample : unmanaged
where TOperator : struct, IResidualOperator<TSample>
{
sum = 0;
sumOfSquares = 0;
// Wider blocks use complete native loads. The eight-sample tail retains the existing compact load,
// which reads eight bytes or eight words without crossing a short row's boundary.
for (int y = 0; y < height; y++)
{
ReadOnlySpan<TSample> sourceRow = source.Slice(y * sourceStride, width);
ReadOnlySpan<TSample> predictionRow = prediction.Slice(y * predictionStride, width);
ref TSample sourceBase = ref MemoryMarshal.GetReference(sourceRow);
ref TSample predictionBase = ref MemoryMarshal.GetReference(predictionRow);
int x = 0;
if (Vector512.IsHardwareAccelerated)
{
for (; x <= width - Vector512<TSample>.Count; x += Vector512<TSample>.Count)
{
Vector512<TSample> sourceVector = Vector512.LoadUnsafe(ref sourceBase, (nuint)x);
Vector512<TSample> predictionVector = Vector512.LoadUnsafe(ref predictionBase, (nuint)x);
Vector512<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector512<short> upper);
// Widen signed lanes before summing: a vector of twelve-bit residuals can exceed short.
// Each vector's squared sum fits int; the block total needs long for large twelve-bit blocks.
sum += Vector512.Sum(Vector512.WidenLower(lower)) + Vector512.Sum(Vector512.WidenUpper(lower));
sumOfSquares += SumSquares(lower);
if (Vector512<TSample>.Count != Vector512<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
sum += Vector512.Sum(Vector512.WidenLower(upper)) + Vector512.Sum(Vector512.WidenUpper(upper));
sumOfSquares += SumSquares(upper);
}
}
}
if (Vector256.IsHardwareAccelerated)
{
for (; x <= width - Vector256<TSample>.Count; x += Vector256<TSample>.Count)
{
Vector256<TSample> sourceVector = Vector256.LoadUnsafe(ref sourceBase, (nuint)x);
Vector256<TSample> predictionVector = Vector256.LoadUnsafe(ref predictionBase, (nuint)x);
Vector256<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector256<short> upper);
// Widen signed lanes before summing: a vector of twelve-bit residuals can exceed short.
// Each vector's squared sum fits int; the block total needs long for large twelve-bit blocks.
sum += Vector256.Sum(Vector256.WidenLower(lower)) + Vector256.Sum(Vector256.WidenUpper(lower));
sumOfSquares += SumSquares(lower);
if (Vector256<TSample>.Count != Vector256<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
sum += Vector256.Sum(Vector256.WidenLower(upper)) + Vector256.Sum(Vector256.WidenUpper(upper));
sumOfSquares += SumSquares(upper);
}
}
}
if (Vector128.IsHardwareAccelerated)
{
for (; x <= width - Vector128<TSample>.Count; x += Vector128<TSample>.Count)
{
Vector128<TSample> sourceVector = Vector128.LoadUnsafe(ref sourceBase, (nuint)x);
Vector128<TSample> predictionVector = Vector128.LoadUnsafe(ref predictionBase, (nuint)x);
Vector128<short> lower = TOperator.Subtract(sourceVector, predictionVector, out Vector128<short> upper);
// Widen signed lanes before summing: a vector of twelve-bit residuals can exceed short.
// Each vector's squared sum fits int; the block total needs long for large twelve-bit blocks.
sum += Vector128.Sum(Vector128.WidenLower(lower)) + Vector128.Sum(Vector128.WidenUpper(lower));
sumOfSquares += SumSquares(lower);
if (Vector128<TSample>.Count != Vector128<short>.Count)
{
// Byte subtraction produces two widened halves; word subtraction has only the lower half.
sum += Vector128.Sum(Vector128.WidenLower(upper)) + Vector128.Sum(Vector128.WidenUpper(upper));
sumOfSquares += SumSquares(upper);
}
}
}
if (Vector128.IsHardwareAccelerated && x <= width - SearchBlockDimension)
{
Vector128<TSample> sourceVector = LoadSearchRow(sourceRow[x..]);
Vector128<TSample> predictionVector = LoadSearchRow(predictionRow[x..]);
sumOfSquares += TOperator.SumSquaredDifferences(sourceVector, predictionVector, out int tailSum);
sum += tailSum;
x += SearchBlockDimension;
}
for (; x < width; x++)
{
int difference = TOperator.Subtract(sourceRow[x], predictionRow[x]);
sum += difference;
sumOfSquares += difference * difference;
}
}
}
/// <summary> /// <summary>
/// Calculates the sum of absolute differences for an 8x8 block. /// Calculates the sum of absolute differences for an 8x8 block.
/// </summary> /// </summary>

135
tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs

@ -140,6 +140,13 @@ public class Av1ResidualBuilderTests
private const HwIntrinsics ResidualConfigurations = private const HwIntrinsics ResidualConfigurations =
HwIntrinsics.AllowAll | HwIntrinsics.DisableAVX512F | HwIntrinsics.DisableAVX | HwIntrinsics.DisableHWIntrinsic; HwIntrinsics.AllowAll | HwIntrinsics.DisableAVX512F | HwIntrinsics.DisableAVX | HwIntrinsics.DisableHWIntrinsic;
/// <summary>
/// Checks rectangular search costs, alternate-row sampling, and wide moments across hardware paths.
/// </summary>
[Fact]
public void RectangularSearchMetricsMatchScalarAcrossHardwareWidths()
=> FeatureTestRunner.RunWithHwIntrinsicsFeature(ValidateRectangularSearchMetrics, ResidualConfigurations);
/// <summary> /// <summary>
/// Verifies 8-bit, 10-bit, and 12-bit residuals across misaligned planes, independent strides, and SIMD tails. /// Verifies 8-bit, 10-bit, and 12-bit residuals across misaligned planes, independent strides, and SIMD tails.
/// </summary> /// </summary>
@ -294,6 +301,134 @@ public class Av1ResidualBuilderTests
Assert.Equal(0, allocated); Assert.Equal(0, allocated);
} }
/// <summary>
/// Exercises independent strides and exact final-row lengths, including every coded block dimension and vector tails.
/// </summary>
private static void ValidateRectangularSearchMetrics()
{
foreach (int width in new[] { 4, 8, 12, 16, 24, 31, 32, 63, 64, 127, 128 })
{
foreach (int height in new[] { 4, 8, 16, 32, 64, 128 })
{
ValidateByteRectangularSearchMetrics(width, height);
ValidateUInt16RectangularSearchMetrics(width, height, 1023);
ValidateUInt16RectangularSearchMetrics(width, height, 4095);
}
}
}
/// <summary>
/// Compares each metric with scalar arithmetic for mixed residuals and both signs at maximum magnitude.
/// </summary>
private static void ValidateByteRectangularSearchMetrics(int width, int height)
{
int sourceStride = width + 3;
int predictionStride = width + 7;
byte[] source = new byte[1 + ((height - 1) * sourceStride) + width];
byte[] prediction = new byte[3 + ((height - 1) * predictionStride) + width];
Span<byte> sourcePlane = source.AsSpan(1);
Span<byte> predictionPlane = prediction.AsSpan(3);
for (int pattern = 0; pattern < 3; pattern++)
{
FillBytePlanes(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height);
int expectedSum = 0;
long expectedSquares = 0;
int expectedSad = 0;
int expectedAlternateSad = 0;
for (int y = 0; y < height; y++)
{
for (int x = 0; x < width; x++)
{
if (pattern != 0)
{
// Constant extrema expose overflowing signed reductions and squared block totals.
sourcePlane[(y * sourceStride) + x] = (byte)(pattern == 1 ? byte.MaxValue : 0);
predictionPlane[(y * predictionStride) + x] = (byte)(pattern == 2 ? byte.MaxValue : 0);
}
int difference = sourcePlane[(y * sourceStride) + x] - predictionPlane[(y * predictionStride) + x];
expectedSum += difference;
expectedSquares += (long)difference * difference;
expectedSad += Math.Abs(difference);
if ((y & 1) == 0)
{
expectedAlternateSad += 2 * Math.Abs(difference);
}
}
}
Av1ResidualBuilder.GetMoments(
sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, out int sum, out long squares);
Assert.Equal(expectedSum, sum);
Assert.Equal(expectedSquares, squares);
Assert.Equal(
expectedSad,
Av1ResidualBuilder.SumAbsoluteDifferences(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, 1));
Assert.Equal(
expectedAlternateSad,
Av1ResidualBuilder.SumAbsoluteDifferences(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, 2));
}
}
/// <summary>
/// Compares each metric with scalar arithmetic for mixed residuals and both signs at maximum magnitude.
/// </summary>
private static void ValidateUInt16RectangularSearchMetrics(int width, int height, int maximumSample)
{
int sourceStride = width + 3;
int predictionStride = width + 7;
ushort[] source = new ushort[1 + ((height - 1) * sourceStride) + width];
ushort[] prediction = new ushort[3 + ((height - 1) * predictionStride) + width];
Span<ushort> sourcePlane = source.AsSpan(1);
Span<ushort> predictionPlane = prediction.AsSpan(3);
for (int pattern = 0; pattern < 3; pattern++)
{
FillUInt16Planes(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, maximumSample);
int expectedSum = 0;
long expectedSquares = 0;
int expectedSad = 0;
int expectedAlternateSad = 0;
for (int y = 0; y < height; y++)
{
for (int x = 0; x < width; x++)
{
if (pattern != 0)
{
// Constant extrema expose overflowing signed reductions and squared block totals.
sourcePlane[(y * sourceStride) + x] = (ushort)(pattern == 1 ? maximumSample : 0);
predictionPlane[(y * predictionStride) + x] = (ushort)(pattern == 2 ? maximumSample : 0);
}
int difference = sourcePlane[(y * sourceStride) + x] - predictionPlane[(y * predictionStride) + x];
expectedSum += difference;
expectedSquares += (long)difference * difference;
expectedSad += Math.Abs(difference);
if ((y & 1) == 0)
{
expectedAlternateSad += 2 * Math.Abs(difference);
}
}
}
Av1ResidualBuilder.GetMoments(
sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, out int sum, out long squares);
Assert.Equal(expectedSum, sum);
Assert.Equal(expectedSquares, squares);
Assert.Equal(
expectedSad,
Av1ResidualBuilder.SumAbsoluteDifferences(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, 1));
Assert.Equal(
expectedAlternateSad,
Av1ResidualBuilder.SumAbsoluteDifferences(sourcePlane, sourceStride, predictionPlane, predictionStride, width, height, 2));
}
}
private static void ValidateSumSquares() private static void ValidateSumSquares()
{ {
short[] residual = new short[127]; short[] residual = new short[127];

Loading…
Cancel
Save