diff --git a/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs b/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs index 3a3db74803..ba4d553a19 100644 --- a/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs +++ b/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1ResidualBuilder.cs @@ -17,6 +17,306 @@ internal static partial class Av1ResidualBuilder /// private const int SearchBlockDimension = 8; + /// + /// Measures absolute prediction error over a rectangular block, optionally sampling alternate rows. + /// + /// Source samples beginning at the block origin. + /// The source row stride, in samples. + /// Prediction samples beginning at the block origin. + /// The prediction row stride, in samples. + /// The block width, in samples. + /// The block height, in samples. + /// One for every row, or two for alternating rows with doubled error. + /// The unnormalized absolute difference over the block. + public static int SumAbsoluteDifferences( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + int rowStep) + => SumAbsoluteDifferences(source, sourceStride, prediction, predictionStride, width, height, rowStep); + + /// + /// Measures signed residual sum and squared error over a rectangular prediction block. + /// + /// Source samples beginning at the block origin. + /// The source row stride, in samples. + /// Prediction samples beginning at the block origin. + /// The prediction row stride, in samples. + /// The block width, in samples. + /// The block height, in samples. + /// The unnormalized signed residual sum. + /// The unnormalized squared residual sum. + public static void GetMoments( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + out int sum, + out long sumOfSquares) + => GetMoments(source, sourceStride, prediction, predictionStride, width, height, out sum, out sumOfSquares); + + /// + /// Measures absolute prediction error over a rectangular block, optionally sampling alternate rows. + /// + /// Source samples beginning at the block origin. + /// The source row stride, in samples. + /// Prediction samples beginning at the block origin. + /// The prediction row stride, in samples. + /// The block width, in samples. + /// The block height, in samples. + /// One for every row, or two for alternating rows with doubled error. + /// The unnormalized absolute difference over the block. + public static int SumAbsoluteDifferences( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + int rowStep) + => SumAbsoluteDifferences(source, sourceStride, prediction, predictionStride, width, height, rowStep); + + /// + /// Measures signed residual sum and squared error over a rectangular prediction block. + /// + /// Source samples beginning at the block origin. + /// The source row stride, in samples. + /// Prediction samples beginning at the block origin. + /// The prediction row stride, in samples. + /// The block width, in samples. + /// The block height, in samples. + /// The unnormalized signed residual sum. + /// The unnormalized squared residual sum. + public static void GetMoments( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + out int sum, + out long sumOfSquares) + => GetMoments(source, sourceStride, prediction, predictionStride, width, height, out sum, out sumOfSquares); + + /// + /// Traverses rectangular SAD candidates using the selected sample operator and descending vector widths. + /// + private static int SumAbsoluteDifferences( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + int rowStep) + where TSample : unmanaged + where TOperator : struct, IResidualOperator + { + 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 sourceRow = source.Slice(y * sourceStride, width); + ReadOnlySpan 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.Count; x += Vector512.Count) + { + Vector512 sourceVector = Vector512.LoadUnsafe(ref sourceBase, (nuint)x); + Vector512 predictionVector = Vector512.LoadUnsafe(ref predictionBase, (nuint)x); + Vector512 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector512 upper); + + // Absolute residuals fit short, but their horizontal sum may not. Widen before reducing. + Vector512 absolute = Vector512.Abs(lower); + sum += Vector512.Sum(Vector512.WidenLower(absolute)) + Vector512.Sum(Vector512.WidenUpper(absolute)); + if (Vector512.Count != Vector512.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.Count; x += Vector256.Count) + { + Vector256 sourceVector = Vector256.LoadUnsafe(ref sourceBase, (nuint)x); + Vector256 predictionVector = Vector256.LoadUnsafe(ref predictionBase, (nuint)x); + Vector256 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector256 upper); + + // Absolute residuals fit short, but their horizontal sum may not. Widen before reducing. + Vector256 absolute = Vector256.Abs(lower); + sum += Vector256.Sum(Vector256.WidenLower(absolute)) + Vector256.Sum(Vector256.WidenUpper(absolute)); + if (Vector256.Count != Vector256.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.Count; x += Vector128.Count) + { + Vector128 sourceVector = Vector128.LoadUnsafe(ref sourceBase, (nuint)x); + Vector128 predictionVector = Vector128.LoadUnsafe(ref predictionBase, (nuint)x); + Vector128 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector128 upper); + + // Absolute residuals fit short, but their horizontal sum may not. Widen before reducing. + Vector128 absolute = Vector128.Abs(lower); + sum += Vector128.Sum(Vector128.WidenLower(absolute)) + Vector128.Sum(Vector128.WidenUpper(absolute)); + if (Vector128.Count != Vector128.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 sourceVector = LoadSearchRow(sourceRow[x..]); + Vector128 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; + } + + /// + /// Accumulates rectangular residual moments without storing an intermediate residual plane. + /// + private static void GetMoments( + ReadOnlySpan source, + int sourceStride, + ReadOnlySpan prediction, + int predictionStride, + int width, + int height, + out int sum, + out long sumOfSquares) + where TSample : unmanaged + where TOperator : struct, IResidualOperator + { + 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 sourceRow = source.Slice(y * sourceStride, width); + ReadOnlySpan 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.Count; x += Vector512.Count) + { + Vector512 sourceVector = Vector512.LoadUnsafe(ref sourceBase, (nuint)x); + Vector512 predictionVector = Vector512.LoadUnsafe(ref predictionBase, (nuint)x); + Vector512 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector512 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.Count != Vector512.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.Count; x += Vector256.Count) + { + Vector256 sourceVector = Vector256.LoadUnsafe(ref sourceBase, (nuint)x); + Vector256 predictionVector = Vector256.LoadUnsafe(ref predictionBase, (nuint)x); + Vector256 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector256 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.Count != Vector256.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.Count; x += Vector128.Count) + { + Vector128 sourceVector = Vector128.LoadUnsafe(ref sourceBase, (nuint)x); + Vector128 predictionVector = Vector128.LoadUnsafe(ref predictionBase, (nuint)x); + Vector128 lower = TOperator.Subtract(sourceVector, predictionVector, out Vector128 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.Count != Vector128.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 sourceVector = LoadSearchRow(sourceRow[x..]); + Vector128 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; + } + } + } + /// /// Calculates the sum of absolute differences for an 8x8 block. /// diff --git a/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs b/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs index 64ac67a3e6..5a4fd219ff 100644 --- a/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs +++ b/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1ResidualBuilderTests.cs @@ -140,6 +140,13 @@ public class Av1ResidualBuilderTests private const HwIntrinsics ResidualConfigurations = HwIntrinsics.AllowAll | HwIntrinsics.DisableAVX512F | HwIntrinsics.DisableAVX | HwIntrinsics.DisableHWIntrinsic; + /// + /// Checks rectangular search costs, alternate-row sampling, and wide moments across hardware paths. + /// + [Fact] + public void RectangularSearchMetricsMatchScalarAcrossHardwareWidths() + => FeatureTestRunner.RunWithHwIntrinsicsFeature(ValidateRectangularSearchMetrics, ResidualConfigurations); + /// /// Verifies 8-bit, 10-bit, and 12-bit residuals across misaligned planes, independent strides, and SIMD tails. /// @@ -294,6 +301,134 @@ public class Av1ResidualBuilderTests Assert.Equal(0, allocated); } + /// + /// Exercises independent strides and exact final-row lengths, including every coded block dimension and vector tails. + /// + 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); + } + } + } + + /// + /// Compares each metric with scalar arithmetic for mixed residuals and both signs at maximum magnitude. + /// + 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 sourcePlane = source.AsSpan(1); + Span 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)); + } + } + + /// + /// Compares each metric with scalar arithmetic for mixed residuals and both signs at maximum magnitude. + /// + 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 sourcePlane = source.AsSpan(1); + Span 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() { short[] residual = new short[127];