From 869340d68e69f0340959be3b1265a19a32e39811 Mon Sep 17 00:00:00 2001 From: James Jackson-South Date: Mon, 7 Sep 2026 15:18:38 +1000 Subject: [PATCH] Add AV1 inter transform rate-distortion estimation Compose DCT and reversible transforms, regular quantization, coefficient rate, local entropy contexts, transform-domain error, and whole-block skip selection using existing workspace storage. Preserve partial-estimate termination and separate skip-header costs from returned statistics. Release .NET 11 VSTest passed 379 focused cases in the current worktree. All 1248 independent complete-estimate comparisons matched optimized libaom exactly across scalar, 128-, 256-, and 512-bit paths. This checkpoint does not complete production motion-controller integration or establish encoder parity. --- .../Av1/Pipeline/Av1TransformBlockEncoder.cs | 280 ++++++++++++++++++ .../Heif/Av1/Av1TransformEstimateTests.cs | 250 ++++++++++++++++ 2 files changed, 530 insertions(+) create mode 100644 tests/ImageSharp.Tests/Formats/Heif/Av1/Av1TransformEstimateTests.cs diff --git a/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1TransformBlockEncoder.cs b/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1TransformBlockEncoder.cs index cf295115e2..24f2ad968e 100644 --- a/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1TransformBlockEncoder.cs +++ b/src/ImageSharp/Formats/Heif/Av1/Pipeline/Av1TransformBlockEncoder.cs @@ -3,6 +3,8 @@ using System.Runtime.CompilerServices; using System.Runtime.InteropServices; +using System.Runtime.Intrinsics; +using SixLabors.ImageSharp.Formats.Heif.Av1.Entropy; using SixLabors.ImageSharp.Formats.Heif.Av1.Pipeline.Quantizers; using SixLabors.ImageSharp.Formats.Heif.Av1.Prediction; using SixLabors.ImageSharp.Formats.Heif.Av1.Prediction.ChromaFromLuma; @@ -1224,6 +1226,284 @@ internal static class Av1TransformBlockEncoder state.TransformType = transformType; } + /// + /// Estimates the luma transform rate and distortion of a prepared inter prediction. + /// + /// The reusable transform storage, overwritten for each transform block. + /// The padded source-minus-prediction block. + /// The number of residual samples between rows. + /// Scratch for one transform's entropy-coding coefficients. + /// The current tile probability state; estimation does not adapt it. + /// The block's top coefficient contexts in four-sample units. + /// The block's left coefficient contexts in four-sample units. + /// The containing prediction block size. + /// The coded extent controlling which padded transform blocks are visited. + /// The transform size selected for the estimate. + /// The effective segment quantizer index. + /// The luma DC quantizer adjustment. + /// The coded sample precision. + /// The quantization sharpness setting. + /// Whether the segment uses reversible transforms and lossless quantization. + /// The current block's rate-distortion multiplier. + /// The rate for signaling the selected transform partition. + /// The rate for signaling a non-skipped prediction block. + /// The rate for signaling a skipped prediction block. + /// The current winning cost used for partial-block termination. + /// The aggregate estimate, excluding the prediction block's skip flag rate. + /// The normalized transform energy before quantization. + /// Whether the aggregate estimate selects transform skip. + /// The decision cost including the skip flag, or the invalid cost for an incomplete estimate. + public static long EstimateInterTransform( + Av1EncoderBlockWorkspace workspace, + ReadOnlySpan residual, + int residualStride, + Span quantizedCoefficients, + Av1SymbolEncoder writer, + ReadOnlySpan aboveContexts, + ReadOnlySpan leftContexts, + Av1BlockSize blockSize, + Size activeSize, + Av1TransformSize transformSize, + int qIndex, + int dcDeltaQ, + Av1BitDepth bitDepth, + int sharpness, + bool lossless, + int rateMultiplier, + int transformSizeRate, + int noSkipRate, + int skipRate, + long bestCost, + out Av1RateDistortionStatistics statistics, + out long sumOfSquares, + out bool skip) + { + int width = transformSize.GetWidth(); + int height = transformSize.GetHeight(); + int widthUnits = width >> Av1Constants.ModeInfoSizeLog2; + int heightUnits = height >> Av1Constants.ModeInfoSizeLog2; + int coefficientCount = transformSize.GetAdjusted().GetSize2d(); + Span transformed = workspace.TransformCoefficients[..coefficientCount]; + Span quantized = quantizedCoefficients[..coefficientCount]; + Span dequantized = workspace.DequantizedCoefficients[..coefficientCount]; + + // A prediction trial changes only its local edge contexts. At most 32 four-sample units lie along + // either edge of a 128-sample block; subsequent transforms see earlier transforms from this trial. + Span above = stackalloc byte[32]; + Span left = stackalloc byte[32]; + aboveContexts.CopyTo(above); + leftContexts.CopyTo(left); + int rate = 0; + long distortion = 0; + sumOfSquares = 0; + skip = true; + long currentCost = Math.Min( + Av1RateDistortion.GetCost(rateMultiplier, noSkipRate + transformSizeRate, 0), + Av1RateDistortion.GetCost(rateMultiplier, skipRate, 0)); + + bool exitEarly = false; + for (int y = 0; y < activeSize.Height; y += height) + { + for (int x = 0; x < activeSize.Width; x += width) + { + // A threshold crossing on the final transform still leaves a complete estimate. Only a + // subsequent unvisited transform invalidates it, so preserve the completed block's statistics. + if (exitEarly) + { + statistics = Av1RateDistortionStatistics.Invalid; + return long.MaxValue; + } + + Span top = above.Slice(x >> Av1Constants.ModeInfoSizeLog2, widthUnits); + Span side = left.Slice(y >> Av1Constants.ModeInfoSizeLog2, heightUnits); + Av1TransformBlockContext context = Av1TileWriter.GetTransformBlockContexts( + Av1ComponentType.Luminance, top, side, blockSize, transformSize); + + ReadOnlySpan transformResidual = residual[((y * residualStride) + x)..]; + ushort endOfBlock; + if (lossless) + { + Av1ForwardTransformer.TransformLossless4x4(transformResidual, transformed, (uint)residualStride); + endOfBlock = Av1ForwardQuantizer.QuantizeLossless(transformed, quantized, dequantized, bitDepth); + } + else + { + Av1ForwardTransformer.Transform2d( + transformResidual, + transformed, + (uint)residualStride, + Av1TransformType.DctDct, + transformSize, + bitDepth.GetBitCount(), + workspace.TransformWorkspace); + + endOfBlock = Av1ForwardQuantizer.QuantizeRegular( + transformed, quantized, dequantized, transformSize, Av1TransformType.DctDct, qIndex, dcDeltaQ, 0, bitDepth, sharpness); + } + + int transformRate = writer.GetCoefficientCost( + transformSize, + Av1TransformType.DctDct, + Av1PredictionMode.DC, + quantized, + Av1ComponentType.Luminance, + context, + endOfBlock, + useReducedTransformSet: false, + Av1FilterIntraMode.AllFilterIntraModes, + usesInterTransformSet: true); + + long transformDistortion = GetTransformError(transformed, dequantized, transformSize, bitDepth, out long transformEnergy); + rate += transformRate; + distortion += transformDistortion; + sumOfSquares += transformEnergy; + skip &= endOfBlock == 0; + + // The running bound chooses the cheaper coded or skipped contribution for each transform. + // The final decision below chooses one skip flag for the whole prediction block. + currentCost += Math.Min( + Av1RateDistortion.GetCost(rateMultiplier, transformRate, transformDistortion), + Av1RateDistortion.GetCost(rateMultiplier, 0, transformEnergy)); + + if (currentCost > bestCost) + { + exitEarly = true; + break; + } + + byte coefficientContext = Av1SymbolContextHelper.GetCoefficientContext( + quantized, transformSize, Av1TransformType.DctDct, endOfBlock); + + top.Fill(coefficientContext); + side.Fill(coefficientContext); + } + } + + long cost; + if (skip) + { + // Empty transforms retain their coefficient-skip rate in the estimate. The block header cost is + // used for this decision but is excluded from the returned rate so the caller can combine planes. + cost = Av1RateDistortion.GetCost(rateMultiplier, skipRate, sumOfSquares); + } + else + { + cost = Av1RateDistortion.GetCost(rateMultiplier, rate + noSkipRate + transformSizeRate, distortion); + rate += transformSizeRate; + if (!lossless) + { + long skipCost = Av1RateDistortion.GetCost(rateMultiplier, skipRate, sumOfSquares); + if (skipCost <= cost) + { + cost = skipCost; + rate = 0; + distortion = sumOfSquares; + skip = true; + } + } + } + + statistics = new Av1RateDistortionStatistics(rateMultiplier, rate, distortion); + return cost; + } + + /// + /// Measures quantization error and unquantized energy in the transform distortion domain. + /// + /// The original transform coefficients. + /// The reconstructed transform coefficients. + /// The transform dimensions controlling coefficient scaling. + /// The coded sample precision. + /// The normalized energy of the original coefficients. + /// The normalized squared quantization error. + public static long GetTransformError( + ReadOnlySpan coefficients, + ReadOnlySpan dequantized, + Av1TransformSize transformSize, + Av1BitDepth bitDepth, + out long sumOfSquares) + { + long error = 0; + long energy = 0; + int i = 0; + ref int coefficientBase = ref MemoryMarshal.GetReference(coefficients); + ref int dequantizedBase = ref MemoryMarshal.GetReference(dequantized); + + // Each Int32 lane holds one coefficient in raster order. Widen before squaring: twelve-bit + // transforms can exceed the signed Int32 square range even though each coefficient and difference fits. + if (Vector512.IsHardwareAccelerated) + { + Vector512 errors = Vector512.Zero; + Vector512 energies = Vector512.Zero; + for (; i <= coefficients.Length - Vector512.Count; i += Vector512.Count) + { + Vector512 values = Vector512.LoadUnsafe(ref coefficientBase, (nuint)i); + Vector512 differences = values - Vector512.LoadUnsafe(ref dequantizedBase, (nuint)i); + (Vector512 lower, Vector512 upper) = Vector512.Widen(values); + (Vector512 lowerDifference, Vector512 upperDifference) = Vector512.Widen(differences); + energies += (lower * lower) + (upper * upper); + errors += (lowerDifference * lowerDifference) + (upperDifference * upperDifference); + } + + energy += Vector512.Sum(energies); + error += Vector512.Sum(errors); + } + + if (Vector256.IsHardwareAccelerated) + { + Vector256 errors = Vector256.Zero; + Vector256 energies = Vector256.Zero; + for (; i <= coefficients.Length - Vector256.Count; i += Vector256.Count) + { + Vector256 values = Vector256.LoadUnsafe(ref coefficientBase, (nuint)i); + Vector256 differences = values - Vector256.LoadUnsafe(ref dequantizedBase, (nuint)i); + (Vector256 lower, Vector256 upper) = Vector256.Widen(values); + (Vector256 lowerDifference, Vector256 upperDifference) = Vector256.Widen(differences); + energies += (lower * lower) + (upper * upper); + errors += (lowerDifference * lowerDifference) + (upperDifference * upperDifference); + } + + energy += Vector256.Sum(energies); + error += Vector256.Sum(errors); + } + + if (Vector128.IsHardwareAccelerated) + { + Vector128 errors = Vector128.Zero; + Vector128 energies = Vector128.Zero; + for (; i <= coefficients.Length - Vector128.Count; i += Vector128.Count) + { + Vector128 values = Vector128.LoadUnsafe(ref coefficientBase, (nuint)i); + Vector128 differences = values - Vector128.LoadUnsafe(ref dequantizedBase, (nuint)i); + (Vector128 lower, Vector128 upper) = Vector128.Widen(values); + (Vector128 lowerDifference, Vector128 upperDifference) = Vector128.Widen(differences); + energies += (lower * lower) + (upper * upper); + errors += (lowerDifference * lowerDifference) + (upperDifference * upperDifference); + } + + energy += Vector128.Sum(energies); + error += Vector128.Sum(errors); + } + + for (; i < coefficients.Length; i++) + { + long value = coefficients[i]; + long difference = value - dequantized[i]; + energy += value * value; + error += difference * difference; + } + + // Normalize high-bit-depth squared values first, rounding once at the accumulated-block boundary. + // Transform scale zero then divides by four, scale one is unchanged, and scale two multiplies by four. + int precisionShift = 2 * (bitDepth.GetBitCount() - 8); + long rounding = (1L << precisionShift) >> 1; + error = (error + rounding) >> precisionShift; + energy = (energy + rounding) >> precisionShift; + int scaleShift = (1 - transformSize.GetScale()) * 2; + sumOfSquares = scaleShift >= 0 ? energy >> scaleShift : energy << -scaleShift; + return scaleShift >= 0 ? error >> scaleShift : error << -scaleShift; + } + public static Span GetPlaneSpan(Buffer2DRegion plane, Point blockOrigin) where TSample : unmanaged { diff --git a/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1TransformEstimateTests.cs b/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1TransformEstimateTests.cs new file mode 100644 index 0000000000..804914a627 --- /dev/null +++ b/tests/ImageSharp.Tests/Formats/Heif/Av1/Av1TransformEstimateTests.cs @@ -0,0 +1,250 @@ +// Copyright (c) Six Labors. +// Licensed under the Six Labors Split License. + +using System.Runtime.InteropServices; +using System.Runtime.Intrinsics; +using SixLabors.ImageSharp.Formats.Heif.Av1; +using SixLabors.ImageSharp.Formats.Heif.Av1.Entropy; +using SixLabors.ImageSharp.Formats.Heif.Av1.Pipeline; +using SixLabors.ImageSharp.Formats.Heif.Av1.Transform; +using SixLabors.ImageSharp.Tests.TestUtilities; + +namespace SixLabors.ImageSharp.Tests.Formats.Heif.Av1; + +/// +/// Verifies motion-winner transform estimates and exports the complete decision boundary for comparison. +/// +[Trait("Format", "Avif")] +public class Av1TransformEstimateTests +{ + /// + /// Exercises complete and partially terminated estimates across sample precision and hardware paths. + /// + [Fact] + public void InterTransformEstimatePreservesDecisionAndContextContracts() + => FeatureTestRunner.RunWithHwIntrinsicsFeature( + ValidateEstimates, + HwIntrinsics.AllowAll | HwIntrinsics.DisableAVX512F | HwIntrinsics.DisableAVX | HwIntrinsics.DisableHWIntrinsic); + + /// + /// Checks storage and skip decisions, retaining independent input samples and the aggregate output. + /// + private static void ValidateEstimates() + { + string vectorWidth = Vector512.IsHardwareAccelerated ? "512" + : Vector256.IsHardwareAccelerated ? "256" + : Vector128.IsHardwareAccelerated ? "128" : "0"; + + string directory = Path.Combine(TestEnvironment.ActualOutputDirectoryFullPath, "Heif", "Av1", "TransformEstimates", vectorWidth); + + Directory.CreateDirectory(directory); + using Av1EncoderBlockWorkspace workspace = new(Configuration.Default); + foreach (Av1BitDepth bitDepth in new[] { Av1BitDepth.EightBit, Av1BitDepth.TenBit, Av1BitDepth.TwelveBit }) + { + foreach (Av1TransformSize transformSize in new[] + { + Av1TransformSize.Size4x4, Av1TransformSize.Size8x8, Av1TransformSize.Size16x16, + Av1TransformSize.Size32x32, Av1TransformSize.Size64x64, Av1TransformSize.Size8x16 + }) + { + foreach (bool split in new[] { false, true }) + { + Av1BlockSize blockSize = transformSize.ToBlockSize(); + if (split) + { + blockSize = transformSize switch + { + Av1TransformSize.Size4x4 => Av1BlockSize.Block8x8, + Av1TransformSize.Size8x8 => Av1BlockSize.Block16x16, + Av1TransformSize.Size16x16 => Av1BlockSize.Block32x32, + Av1TransformSize.Size32x32 => Av1BlockSize.Block64x64, + Av1TransformSize.Size64x64 => Av1BlockSize.Block128x128, + _ => Av1BlockSize.Block16x32 + }; + } + + int width = blockSize.GetWidth(); + int height = blockSize.GetHeight(); + int stride = width + 3; + short[] residual = new short[stride * height]; + int[] quantized = new int[transformSize.GetAdjusted().GetSize2d()]; + byte[] above = new byte[width / 4]; + byte[] left = new byte[height / 4]; + for (int i = 0; i < above.Length; i++) + { + above[i] = (byte)(((i % 3) << Av1Constants.CoefficientContextBitCount) + (i % 5)); + } + + for (int i = 0; i < left.Length; i++) + { + left[i] = (byte)((((i + 1) % 3) << Av1Constants.CoefficientContextBitCount) + (i % 7)); + } + + byte[] originalAbove = (byte[])above.Clone(); + byte[] originalLeft = (byte[])left.Clone(); + foreach (int qIndex in new[] { 0, 90, 255 }) + { + bool lossless = qIndex == 0; + if (lossless && transformSize != Av1TransformSize.Size4x4) + { + continue; + } + + using Av1SymbolEncoder writer = new(Configuration.Default, 65536, qIndex, updateCdf: true); + for (int pattern = 0; pattern < 4; pattern++) + { + int bits = bitDepth.GetBitCount(); + int amplitude = (1 << bits) - 1; + uint random = 73; + for (int y = 0; y < height; y++) + { + for (int x = 0; x < width; x++) + { + random = unchecked((random * 1664525) + 1013904223); + int value = pattern switch + { + 0 => 0, + 1 => (int)(random % (uint)((2 * amplitude) + 1)) - amplitude, + 2 => ((x + y) % 3) - 1, + _ => x == 0 && y == 0 ? amplitude : 0 + }; + + residual[(y * stride) + x] = (short)value; + } + } + + int multiplier = 173; + int partitionRate = split ? 431 : 0; + int noSkipRate = 571; + int skipRate = 619; + int sharpness = pattern; + long cost = Av1TransformBlockEncoder.EstimateInterTransform( + workspace, + residual, + stride, + quantized, + writer, + above, + left, + blockSize, + new Size(width, height), + transformSize, + qIndex, + 0, + bitDepth, + sharpness, + lossless, + multiplier, + partitionRate, + noSkipRate, + skipRate, + long.MaxValue, + out Av1RateDistortionStatistics statistics, + out long energy, + out bool skip); + + Assert.Equal(originalAbove, above); + Assert.Equal(originalLeft, left); + if (pattern == 0) + { + Assert.True(skip); + Assert.Equal(0, statistics.Distortion); + Assert.Equal(0, energy); + Assert.Equal(Av1RateDistortion.GetCost(multiplier, skipRate, 0), cost); + Assert.True(statistics.Rate > 0); + } + else if (lossless) + { + Assert.False(skip); + Assert.Equal(0, statistics.Distortion); + } + + // Running the same trial again must see identical probabilities and incoming neighbors. + long repeated = Av1TransformBlockEncoder.EstimateInterTransform( + workspace, + residual, + stride, + quantized, + writer, + above, + left, + blockSize, + new Size(width, height), + transformSize, + qIndex, + 0, + bitDepth, + sharpness, + lossless, + multiplier, + partitionRate, + noSkipRate, + skipRate, + long.MaxValue, + out Av1RateDistortionStatistics repeatedStatistics, + out long repeatedEnergy, + out bool repeatedSkip); + + Assert.Equal(cost, repeated); + Assert.Equal(statistics, repeatedStatistics); + Assert.Equal(energy, repeatedEnergy); + Assert.Equal(skip, repeatedSkip); + + if (split) + { + long terminated = Av1TransformBlockEncoder.EstimateInterTransform( + workspace, + residual, + stride, + quantized, + writer, + above, + left, + blockSize, + new Size(width, height), + transformSize, + qIndex, + 0, + bitDepth, + sharpness, + lossless, + multiplier, + partitionRate, + noSkipRate, + skipRate, + 0, + out Av1RateDistortionStatistics incomplete, + out _, + out _); + + Assert.Equal(long.MaxValue, terminated); + Assert.Equal(Av1RateDistortionStatistics.Invalid, incomplete); + } + + using BinaryWriter output = new(File.Create(Path.Combine( + directory, $"{bits}-{width}-{height}-{(int)transformSize}-{qIndex}-{pattern}.bin"))); + + foreach (int value in new[] + { + bits, width, height, (int)transformSize, qIndex, sharpness, lossless ? 1 : 0, + multiplier, partitionRate, noSkipRate, skipRate, stride + }) + { + output.Write(value); + } + + output.Write(cost); + output.Write((long)statistics.Rate); + output.Write(statistics.Distortion); + output.Write(energy); + output.Write(skip ? 1L : 0L); + output.Write(above); + output.Write(left); + output.Write(MemoryMarshal.AsBytes(residual.AsSpan())); + } + } + } + } + } + } +}