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