// Copyright (c) Six Labors. // Licensed under the Six Labors Split License. using System.Runtime.Intrinsics; using SixLabors.ImageSharp.Formats.Heif.Av1.Transform; using SixLabors.ImageSharp.Formats.Heif.Av1.Transform.Forward; using SixLabors.ImageSharp.Tests.TestUtilities; namespace SixLabors.ImageSharp.Tests.Formats.Heif.Av1; [Trait("Format", "Avif")] public class Av1ForwardTransformTests { /// /// The hardware configurations covering every transform SIMD tier and the scalar fallback. /// private const HwIntrinsics TransformConfigurations = HwIntrinsics.AllowAll | HwIntrinsics.DisableAVX512F | HwIntrinsics.DisableAVX | HwIntrinsics.DisableHWIntrinsic; /// /// Gets every normative transform size, type, and bit-depth combination exercised by the forward and inverse suites. /// public static TheoryData ValidTransformCases { get; } = CreateValidTransformCases(); /// /// Verifies DCT operator parity across the supported hardware feature levels. /// [Fact] public void DctOperatorsProduceIdenticalScalarAndSimdResults() => FeatureTestRunner.RunWithHwIntrinsicsFeature(AssertDctOperatorParity, TransformConfigurations); /// /// Verifies ADST operator parity across the supported hardware feature levels. /// [Fact] public void AdstOperatorsProduceIdenticalScalarAndSimdResults() => FeatureTestRunner.RunWithHwIntrinsicsFeature(AssertAdstOperatorParity, TransformConfigurations); /// /// Verifies identity operator parity across the supported hardware feature levels. /// [Fact] public void IdentityOperatorsProduceIdenticalScalarAndSimdResults() => FeatureTestRunner.RunWithHwIntrinsicsFeature(AssertIdentityOperatorParity, TransformConfigurations); /// /// Verifies the complete sixteen-lane two-dimensional traversal matrix across hardware feature levels. /// [Fact] public void Vector512KernelsMatchScalarForEveryApplicableConfiguration() => FeatureTestRunner.RunWithHwIntrinsicsFeature(AssertVector512TransformParity, TransformConfigurations); /// /// Verifies the forward DCT operators against their scalar implementations. /// private static void AssertDctOperatorParity() { AssertOperatorParity(4); AssertOperatorParity(8); AssertOperatorParity(16); AssertOperatorParity(32); AssertOperatorParity(64); } /// /// Verifies the forward ADST operators against their scalar implementations. /// private static void AssertAdstOperatorParity() { AssertOperatorParity(4); AssertOperatorParity(8); AssertOperatorParity(16); } /// /// Verifies the forward identity operators against their scalar implementations. /// private static void AssertIdentityOperatorParity() { AssertOperatorParity(4); AssertOperatorParity(8); AssertOperatorParity(16); AssertOperatorParity(32); } /// /// Verifies that every applicable SIMD traversal produces the same coefficients as the scalar traversal. /// /// The integral value. /// The integral value. /// The coded sample bit depth. [Theory] [MemberData(nameof(ValidTransformCases))] public void TwoDimensionalSimdKernelsMatchScalarForEveryValidConfiguration( int transformTypeValue, int transformSizeValue, int bitDepth) { Av1TransformType transformType = (Av1TransformType)transformTypeValue; Av1TransformSize transformSize = (Av1TransformSize)transformSizeValue; Av1Transform2dFlipConfiguration config = Av1Transform2dFlipConfiguration.CreateForward(transformType, transformSize, bitDepth); DispatchColumn(transformType, transformSize, bitDepth, ref config); } [Fact] public void TransformDispatchDoesNotAllocatePerBlock() { const int width = 8; short[] input = new short[width * width]; int[] output = new int[input.Length]; int[] workspace = new int[Av1TransformWorkspace.MaximumLength]; Av1ForwardTransformer.Transform2d(input, output, width, Av1TransformType.DctDct, Av1TransformSize.Size8x8, 8, workspace); long before = GC.GetAllocatedBytesForCurrentThread(); for (int iteration = 0; iteration < 32; iteration++) { Av1ForwardTransformer.Transform2d(input, output, width, Av1TransformType.DctDct, Av1TransformSize.Size8x8, 8, workspace); } long allocated = GC.GetAllocatedBytesForCurrentThread() - before; Assert.Equal(0, allocated); } /// /// Compares one forward transform operator across scalar and all SIMD lane widths. /// /// The forward transform operator. /// The transform length. private static void AssertOperatorParity(int length) where TOperator : struct, IAv1Transform1dOperator { const int cosBit = 12; Av1TransformStageRange stageRange = default; for (int index = 0; index < Av1Transform2dFlipConfiguration.MaxStageNumber; index++) { stageRange[index] = 24; } Av1TransformVector> input128 = default; Av1TransformVector> output128 = default; Av1TransformVector> step128 = default; Av1TransformVector> input256 = default; Av1TransformVector> output256 = default; Av1TransformVector> step256 = default; Av1TransformVector> input512 = default; Av1TransformVector> output512 = default; Av1TransformVector> step512 = default; for (int index = 0; index < length; index++) { input128[index] = Vector128.Create( GetInputValue(index, 0), GetInputValue(index, 1), GetInputValue(index, 2), GetInputValue(index, 3)); input256[index] = Vector256.Create( GetInputValue(index, 0), GetInputValue(index, 1), GetInputValue(index, 2), GetInputValue(index, 3), GetInputValue(index, 4), GetInputValue(index, 5), GetInputValue(index, 6), GetInputValue(index, 7)); input512[index] = Vector512.Create( GetInputValue(index, 0), GetInputValue(index, 1), GetInputValue(index, 2), GetInputValue(index, 3), GetInputValue(index, 4), GetInputValue(index, 5), GetInputValue(index, 6), GetInputValue(index, 7), GetInputValue(index, 8), GetInputValue(index, 9), GetInputValue(index, 10), GetInputValue(index, 11), GetInputValue(index, 12), GetInputValue(index, 13), GetInputValue(index, 14), GetInputValue(index, 15)); } TOperator.Transform(ref input128, ref output128, ref step128, cosBit, stageRange); TOperator.Transform(ref input256, ref output256, ref step256, cosBit, stageRange); TOperator.Transform(ref input512, ref output512, ref step512, cosBit, stageRange); int[] scalarInput = new int[length]; int[] scalarOutput = new int[length]; int[] scalarStep = new int[length]; for (int lane = 0; lane < Vector512.Count; lane++) { for (int index = 0; index < length; index++) { scalarInput[index] = GetInputValue(index, lane); } TOperator.Transform(scalarInput, scalarOutput, scalarStep, cosBit, stageRange); for (int index = 0; index < length; index++) { Assert.Equal(scalarOutput[index], output512[index].GetElement(lane)); if (lane < Vector256.Count) { Assert.Equal(scalarOutput[index], output256[index].GetElement(lane)); } if (lane < Vector128.Count) { Assert.Equal(scalarOutput[index], output128[index].GetElement(lane)); } } } } /// /// Runs every valid forward transform configuration capable of filling a sixteen-lane tile. /// private static void AssertVector512TransformParity() { for (Av1TransformSize transformSize = 0; transformSize < Av1TransformSize.AllSizes; transformSize++) { if (transformSize.GetWidth() < Vector512.Count || transformSize.GetHeight() < Vector512.Count) { continue; } for (Av1TransformType transformType = 0; transformType < Av1TransformType.AllTransformTypes; transformType++) { Av1Transform2dFlipConfiguration allowedConfig = Av1Transform2dFlipConfiguration.CreateForward(transformType, transformSize, 8); if (!allowedConfig.IsAllowed()) { continue; } for (int bitDepth = 8; bitDepth <= 12; bitDepth += 2) { Av1Transform2dFlipConfiguration config = Av1Transform2dFlipConfiguration.CreateForward(transformType, transformSize, bitDepth); DispatchColumn(transformType, transformSize, bitDepth, ref config); } } } } /// /// Creates the complete normative transform matrix shared by the forward and inverse parity tests. /// /// The transform type, size, and bit-depth cases. private static TheoryData CreateValidTransformCases() { TheoryData cases = []; for (Av1TransformSize transformSize = 0; transformSize < Av1TransformSize.AllSizes; transformSize++) { for (Av1TransformType transformType = 0; transformType < Av1TransformType.AllTransformTypes; transformType++) { Av1Transform2dFlipConfiguration config = Av1Transform2dFlipConfiguration.CreateForward(transformType, transformSize, 8); if (!config.IsAllowed()) { continue; } // libaom verifies the low-bit-depth kernel separately from its 10- and 12-bit kernels. Keeping each // depth as a distinct case makes any fixed-point range failure identify the exact configuration. for (int bitDepth = 8; bitDepth <= 12; bitDepth += 2) { cases.Add((int)transformType, (int)transformSize, bitDepth); } } } return cases; } /// /// Closes the static-generic column operator selected by a transform configuration. /// /// The compound transform type. /// The transform-block dimensions. /// The coded sample bit depth. /// The forward transform configuration. private static void DispatchColumn( Av1TransformType transformType, Av1TransformSize transformSize, int bitDepth, ref Av1Transform2dFlipConfiguration config) { switch (config.TransformFunctionTypeColumn) { case Av1TransformFunctionType.Dct4: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct8: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct16: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct32: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct64: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst4: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst8: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst16: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity4: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity8: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity16: DispatchRow(transformType, transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity32: DispatchRow(transformType, transformSize, bitDepth, ref config); break; default: Assert.Fail($"Unexpected column function {config.TransformFunctionTypeColumn} for {transformType} {transformSize}."); break; } } /// /// Closes the static-generic row operator after the column operator has been selected. /// /// The selected column operator. /// The compound transform type. /// The transform-block dimensions. /// The coded sample bit depth. /// The forward transform configuration. private static void DispatchRow( Av1TransformType transformType, Av1TransformSize transformSize, int bitDepth, ref Av1Transform2dFlipConfiguration config) where TColumnOperator : struct, IAv1Transform1dOperator { switch (config.TransformFunctionTypeRow) { case Av1TransformFunctionType.Dct4: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct8: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct16: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct32: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Dct64: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst4: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst8: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Adst16: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity4: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity8: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity16: AssertTransform2dParity(transformSize, bitDepth, ref config); break; case Av1TransformFunctionType.Identity32: AssertTransform2dParity(transformSize, bitDepth, ref config); break; default: Assert.Fail($"Unexpected row function {config.TransformFunctionTypeRow} for {transformType} {transformSize}."); break; } } /// /// Compares scalar and SIMD forward traversals using padded rows and bounded extreme residuals. /// /// The selected column operator. /// The selected row operator. /// The transform-block dimensions. /// The coded sample bit depth. /// The forward transform configuration. private static void AssertTransform2dParity( Av1TransformSize transformSize, int bitDepth, ref Av1Transform2dFlipConfiguration config) where TColumnOperator : struct, IAv1Transform1dOperator where TRowOperator : struct, IAv1Transform1dOperator { int width = transformSize.GetWidth(); int height = transformSize.GetHeight(); int inputStride = width + 5; int maximum = (1 << bitDepth) - 1; short[] input = new short[inputStride * height]; // Padded rows exercise the same edge-block layout used by the encoder. The alternating extrema are the // bounded residual limits used by libaom's SIMD match tests and expose wrapping errors in fixed-point stages. for (int row = 0; row < height; row++) { for (int column = 0; column < width; column++) { int index = (row * width) + column; input[(row * inputStride) + column] = (short)((index & 3) switch { 0 => maximum, 1 => -maximum, 2 => ((index * 73) % ((maximum * 2) + 1)) - maximum, _ => 0, }); } } int coefficientCount = width * height; int workspaceLength = Av1TransformWorkspace.GetRequiredLength(transformSize); int[] scalar = new int[coefficientCount]; int[] vector128 = new int[coefficientCount]; int[] scalarWorkspace = new int[workspaceLength]; int[] vector128Workspace = new int[workspaceLength]; Av1ForwardTransformer.Transform2dScalar(input, scalar, (uint)inputStride, ref config, scalarWorkspace); Av1ForwardTransformer.Transform2dVector128(input, vector128, (uint)inputStride, ref config, vector128Workspace); Assert.Equal(scalar, vector128); // The production dispatcher uses 256-bit lanes only when both axes contain a complete eight-lane tile. if (width >= Vector256.Count && height >= Vector256.Count) { int[] vector256 = new int[coefficientCount]; int[] vector256Workspace = new int[workspaceLength]; Av1ForwardTransformer.Transform2dVector256(input, vector256, (uint)inputStride, ref config, vector256Workspace); Assert.Equal(scalar, vector256); } if (width >= Vector512.Count && height >= Vector512.Count) { int[] vector512 = new int[coefficientCount]; int[] vector512Workspace = new int[workspaceLength]; Av1ForwardTransformer.Transform2dVector512(input, vector512, (uint)inputStride, ref config, vector512Workspace); Assert.Equal(scalar, vector512); } } /// /// Produces deterministic bounded input for one transform position and SIMD lane. /// /// The position within the transform. /// The SIMD lane index. /// The input value. private static int GetInputValue(int index, int lane) => (((index * 73) + (lane * 151)) % 1023) - 511; }