From fad73d8cfdaa84dded8f2c0501642be08820c15d Mon Sep 17 00:00:00 2001 From: winscripter <142818255+winscripter@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:37:00 +0400 Subject: [PATCH] Add encoder ANS SIMD bit cost estimation --- .../Entropy/JxlAnsHybridUIntConfiguration.cs | 2 +- .../Jxl/Processing/Encoder/Ans/JxlAnsSimd.cs | 275 ++++++++++++++++++ .../Jxl/Processing/JxlCoefficientOrder.cs | 6 +- 3 files changed, 277 insertions(+), 6 deletions(-) create mode 100644 src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlAnsSimd.cs diff --git a/src/ImageSharp/Formats/Jxl/IO/Entropy/JxlAnsHybridUIntConfiguration.cs b/src/ImageSharp/Formats/Jxl/IO/Entropy/JxlAnsHybridUIntConfiguration.cs index c1d6409076..9ff4306522 100644 --- a/src/ImageSharp/Formats/Jxl/IO/Entropy/JxlAnsHybridUIntConfiguration.cs +++ b/src/ImageSharp/Formats/Jxl/IO/Entropy/JxlAnsHybridUIntConfiguration.cs @@ -31,7 +31,7 @@ internal sealed class JxlAnsHybridUIntConfiguration : IJxlFields public uint LsbMask => (1u << (int)this.LsbInToken) - 1; - public void Encode(uint value, ref uint token, ref uint bitCount, ref uint bits) + public void Encode(uint value, out uint token, out uint bitCount, out uint bits) { if (value < this.SplitToken) { diff --git a/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlAnsSimd.cs b/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlAnsSimd.cs new file mode 100644 index 0000000000..e23bc27009 --- /dev/null +++ b/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlAnsSimd.cs @@ -0,0 +1,275 @@ +// Copyright (c) Six Labors. +// Licensed under the Six Labors Split License. + +using System.Numerics; +using System.Runtime.CompilerServices; +using SixLabors.ImageSharp.Formats.Jxl.IO.Entropy; + +namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Encoder.Ans; + +/// +/// SIMD utilities for ANS entropy encoder +/// +internal static class JxlAnsSimd +{ + private static readonly Vector IotaOffsets = CreateIotaOffsets(); + + private static Vector CreateIotaOffsets() + { + Span values = stackalloc uint[Vector.Count]; + + for (int i = 0; i < values.Length; i++) + { + values[i] = (uint)i; + } + + return new Vector(values); + } + + /// + /// Adds continuously incrementing numbers to the vector. + /// + /// The input vector. + /// vec + [ 1, 2, 3, 4, 5, ... ] + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static Vector Iota(uint vec) => Vector.Create(vec) + IotaOffsets; + + private static uint EstimateTokenCostImpl(uint e, uint m, uint l, ref uint values, int len, ref uint output) + { + Vector split = Vector.Create(1u << (int)e); + Vector expOffset = Vector.Create(127u); + Vector ebOffset = Vector.Create(127u + m + l); + Vector @base = Vector.Create((1u << (int)e) - (e << (int)(m + l))); + Vector mulN = Vector.Create(1u << (int)(m + l)); + Vector maskL = Vector.Create((1u << (int)l) - 1); + Vector maskM = Vector.Create(((1u << (int)m) - 1) << (int)l); + Vector largeThreshold = Vector.Create((1u << 2) - 1); + const uint largeShiftVal = 10; + Vector largeShift = Vector.Create(largeShiftVal); + + Vector extraBits = Vector.Zero; + int lastFull = Vector.Count * (len / Vector.Count); + + for (int i = 0; i < lastFull; i += Vector.Count) + { + Vector val = Vector.LoadUnsafe(ref Unsafe.Add(ref values, i)); + Vector isLarge = Vector.GreaterThan(val, largeThreshold); + Vector valShifted = Vector.ShiftRightLogical(val, (int)largeShiftVal); + Vector notLiteral = Vector.GreaterThanOrEqual(val, split); + Vector valFixed = Vector.ConditionalSelect(isLarge, valShifted, val); + Vector L = val & maskL; + Vector exp = Vector.ShiftRightLogical(valFixed, 23); + Vector expFixed = Vector.ConditionalSelect(isLarge, exp + largeShift, exp); + Vector n = expFixed - expOffset; + Vector eb = expFixed - ebOffset; + Vector M = Vector.ShiftRightLogical(valFixed, (int)(23 - m - l)); + Vector a = @base + (n * mulN); + Vector d = M & maskM; + Vector ebFixed = Vector.ConditionalSelect(notLiteral, eb, Vector.Zero); + Vector c = a | L; + extraBits += ebFixed; + Vector t = c | d; + Vector tFixed = Vector.ConditionalSelect(notLiteral, t, val); + tFixed.StoreUnsafe(ref Unsafe.Add(ref output, i)); + } + + if (lastFull < len) + { + Vector stop = Vector.Create((uint)len); + Vector fence = Iota((uint)lastFull); + Vector take = Vector.LessThan(fence, stop); + Vector val = Vector.LoadUnsafe(ref Unsafe.Add(ref values, lastFull)); + Vector isLarge = Vector.GreaterThan(val, largeThreshold); + Vector valShifted = Vector.ShiftRightLogical(val, (int)largeShiftVal); + Vector notLiteral = Vector.GreaterThanOrEqual(val, split); + Vector valFixed = Vector.ConditionalSelect(isLarge, valShifted, val); + Vector L = val | maskL; + Vector exp = Vector.ShiftRightLogical(valFixed, 23); + Vector exp_fixed = Vector.ConditionalSelect(isLarge, exp + largeShift, exp); + Vector n = exp_fixed - expOffset; + Vector eb = exp_fixed - ebOffset; + Vector M = Vector.ShiftRightLogical(valFixed, 23); + Vector a = @base + (n * mulN); + Vector d = M & maskM; + Vector ebFixed = Vector.ConditionalSelect(notLiteral, eb, Vector.Zero); + Vector ebMasked = Vector.ConditionalSelect(take, ebFixed, Vector.Zero); + Vector c = a | L; + extraBits += ebMasked; + Vector t = c | d; + Vector tFixed = Vector.ConditionalSelect(notLiteral, t, val); + tFixed.StoreUnsafe(ref Unsafe.Add(ref output, lastFull)); + } + + return Vector.Sum(extraBits); + } + + public static uint EstimateTokenCost(ref uint values, int len, JxlAnsHybridUIntConfiguration cfg, ref uint tokens) + { + if (!Vector.IsHardwareAccelerated) + { + // No SIMD support + uint extraBits = 0; + + for (int i = 0; i < len; i++) + { + uint v = Unsafe.Add(ref values, i); + cfg.Encode(v, out uint tok, out uint nbits, out _); // Last parameter is bits + extraBits += nbits; + Unsafe.Add(ref tokens, i) = tok; + } + + return extraBits; + } + else + { + // Have SIMD support + if (cfg.SplitExponent == 0) + { + return EstimateTokenCostImpl(0, 0, 0, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 2) + { + return EstimateTokenCostImpl(2, 0, 1, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 3) + { + if (cfg.MsbInToken == 1) + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(3, 1, 0, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(3, 1, 2, ref values, len, ref tokens); + } + } + else + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(3, 2, 0, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(3, 2, 1, ref values, len, ref tokens); + } + } + } + else if (cfg.SplitExponent == 4) + { + if (cfg.MsbInToken == 1) + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(4, 1, 0, ref values, len, ref tokens); + } + else if (cfg.LsbInToken == 2) + { + return EstimateTokenCostImpl(4, 1, 2, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(4, 1, 3, ref values, len, ref tokens); + } + } + else + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(4, 2, 0, ref values, len, ref tokens); + } + else if (cfg.LsbInToken == 1) + { + return EstimateTokenCostImpl(4, 2, 1, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(4, 2, 2, ref values, len, ref tokens); + } + } + } + else if (cfg.SplitExponent == 5) + { + if (cfg.MsbInToken == 1) + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(5, 1, 0, ref values, len, ref tokens); + } + else if (cfg.LsbInToken == 2) + { + return EstimateTokenCostImpl(5, 1, 2, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(5, 1, 4, ref values, len, ref tokens); + } + } + else + { + if (cfg.LsbInToken == 0) + { + return EstimateTokenCostImpl(5, 2, 0, ref values, len, ref tokens); + } + else if (cfg.LsbInToken == 1) + { + return EstimateTokenCostImpl(5, 2, 1, ref values, len, ref tokens); + } + else if (cfg.LsbInToken == 2) + { + return EstimateTokenCostImpl(5, 2, 2, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(5, 2, 3, ref values, len, ref tokens); + } + } + } + else if (cfg.SplitExponent == 6) + { + if (cfg.MsbInToken == 0) + { + return EstimateTokenCostImpl(6, 0, 0, ref values, len, ref tokens); + } + else if (cfg.MsbInToken == 1) + { + return EstimateTokenCostImpl(6, 1, 5, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(6, 2, 4, ref values, len, ref tokens); + } + } + else if (cfg.SplitExponent is >= 7 and <= 12) + { + if (cfg.SplitExponent == 7) + { + return EstimateTokenCostImpl(7, 0, 0, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 8) + { + return EstimateTokenCostImpl(8, 0, 0, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 9) + { + return EstimateTokenCostImpl(9, 0, 0, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 10) + { + return EstimateTokenCostImpl(10, 0, 0, ref values, len, ref tokens); + } + else if (cfg.SplitExponent == 11) + { + return EstimateTokenCostImpl(11, 0, 0, ref values, len, ref tokens); + } + else + { + return EstimateTokenCostImpl(12, 0, 0, ref values, len, ref tokens); + } + } + + return ~0u; + } + } +} diff --git a/src/ImageSharp/Formats/Jxl/Processing/JxlCoefficientOrder.cs b/src/ImageSharp/Formats/Jxl/Processing/JxlCoefficientOrder.cs index deb5045f2a..4cace39d73 100644 --- a/src/ImageSharp/Formats/Jxl/Processing/JxlCoefficientOrder.cs +++ b/src/ImageSharp/Formats/Jxl/Processing/JxlCoefficientOrder.cs @@ -45,11 +45,7 @@ internal static class JxlCoefficientOrder [MethodImpl(MethodImplOptions.AggressiveInlining)] public static uint CoeffOrderContext(uint value) { - uint token = 0; - uint nbits = 0; - uint bits = 0; - - new JxlAnsHybridUIntConfiguration(0, 0, 0).Encode(value, ref token, ref nbits, ref bits); + new JxlAnsHybridUIntConfiguration(0, 0, 0).Encode(value, out uint token, out uint nbits, out uint bits); return Math.Min(token, PermutationContexts - 1u); }