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()));
+ }
+ }
+ }
+ }
+ }
+ }
+}