Browse Source

Refactor, add MA encoding (decoder-only), finish quantization weights, add bits tests

Files implemented: quant_weights.cc, quant_weights.h, modular/encoding/encoding.cc, modular/encoding/encoding.h, ma_common.h, bits_test.cc
pull/3153/head
winscripter 4 weeks ago
parent
commit
ceb282cc11
  1. 18
      src/ImageSharp/Formats/Jxl/Memory/JxlPlaneBase.cs
  2. 16
      src/ImageSharp/Formats/Jxl/Memory/JxlPlane{T}.cs
  3. 4
      src/ImageSharp/Formats/Jxl/Processing/Butteraugli/Butteraugli.cs
  4. 26
      src/ImageSharp/Formats/Jxl/Processing/JxlScopeGuard.cs
  5. 96
      src/ImageSharp/Formats/Jxl/Processing/JxlSimdUtils.cs
  6. 4
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlMaTreeLookup.cs
  7. 66
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlGroupHeader.cs
  8. 16
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs
  9. 19
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs
  10. 866
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlModularEncoding.cs
  11. 46
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlPropertyDecisionNode.cs
  12. 13
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeLut.cs
  13. 4
      src/ImageSharp/Formats/Jxl/Processing/Modular/JxlModularChannel.cs
  14. 4
      src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlDequantMatrices.cs
  15. 700
      src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlQuantWeights.cs
  16. 8
      src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlQuantizerConstants.cs
  17. 16
      src/ImageSharp/Formats/Jxl/Processing/Splines/JxlSplineUtils.cs
  18. 89
      tests/ImageSharp.Tests/Formats/Jxl/Processing/BitsTests.cs
  19. 2
      tests/ImageSharp.Tests/Formats/Jxl/README.md
  20. 19
      tests/ImageSharp.Tests/Formats/Jxl/TestUtils/ColorEncodingDescriptor.cs

18
src/ImageSharp/Formats/Jxl/Memory/JxlPlaneBase.cs

@ -157,6 +157,24 @@ internal class JxlPlaneBase : IDisposable
return MemoryMarshal.Cast<byte, T>(row);
}
protected Span<T> GetRowMinusBase<T>(int y, int minus)
where T : unmanaged
{
DebugGuard.MustBeLessThan(y, this.YSize, nameof(y));
Span<byte> row = this.Bytes.Span[((y * this.BytesPerRow) - minus)..];
return MemoryMarshal.Cast<byte, T>(row);
}
protected Span<T> GetRowPlusBase<T>(int y, int plus)
where T : unmanaged
{
DebugGuard.MustBeLessThan(y, this.YSize, nameof(y));
Span<byte> row = this.Bytes.Span[((y * this.BytesPerRow) + plus)..];
return MemoryMarshal.Cast<byte, T>(row);
}
/// <summary>
/// Swaps properties &amp; data of this image with the specified image.
/// </summary>

16
src/ImageSharp/Formats/Jxl/Memory/JxlPlane{T}.cs

@ -64,6 +64,22 @@ internal class JxlPlane<T> : JxlPlaneBase
/// <returns>A span which covers memory for the specified row.</returns>
public Span<T> GetRow(int y) => this.GetRowBase<T>(y);
/// <summary>
/// Returns a span for the specified row, but backwards by specified number of elements.
/// </summary>
/// <param name="y">The row index.</param>
/// <param name="minus">Once the offset of the row was derived, this is how much to subtract the offset.</param>
/// <returns>A span which covers memory for the specified row.</returns>
public Span<T> GetRowMinus(int y, int minus) => this.GetRowMinusBase<T>(y, minus);
/// <summary>
/// Returns a span for the specified row, but with offset by specified number of elements.
/// </summary>
/// <param name="y">The row index.</param>
/// <param name="plus">Once the offset of the row was derived, this is the additional offset.</param>
/// <returns>A span which covers memory for the specified row.</returns>
public Span<T> GetRowPlus(int y, int plus) => this.GetRowMinusBase<T>(y, plus);
/// <summary>
/// Returns a span for the specified row within the specified rectangle bounds.
/// </summary>

4
src/ImageSharp/Formats/Jxl/Processing/Butteraugli/Butteraugli.cs

@ -60,9 +60,7 @@ internal static class Butteraugli
};
#pragma warning restore
private static readonly DenseMatrix<float> Heatmap;
static Butteraugli() => Heatmap = new(HeatmapData);
private static readonly DenseMatrix<float> Heatmap = new(HeatmapData);
public static ReadOnlySpan<float> Wmul =>
[

26
src/ImageSharp/Formats/Jxl/Processing/JxlScopeGuard.cs

@ -0,0 +1,26 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
namespace SixLabors.ImageSharp.Formats.Jxl.Processing;
/// <summary>
/// Conditionally invokes the specified method when
/// appropriate and when it goes out of scope.
/// </summary>
internal struct JxlScopeGuard(Action action) : IDisposable
{
/// <summary>
/// When true the action will be invoked when Dispose() is called.
/// </summary>
private bool isArmed = true;
public void Disarm() => this.isArmed = false;
public readonly void Dispose()
{
if (this.isArmed)
{
action();
}
}
}

96
src/ImageSharp/Formats/Jxl/Processing/JxlSimdUtils.cs

@ -409,6 +409,102 @@ internal static partial class JxlSimdUtils
return xle0 ? -result : result;
}
// Raises 2 to the power of x
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static Vector128<float> FastPow2f(Vector128<float> x)
{
Vector128<float> floorX = Vector128.Floor(x);
Vector128<int> exponent = Vector128.ConvertToInt32(floorX) + Vector128.Create(127);
Vector128<float> exp = exponent.AsSingle() << 23;
Vector128<float> frac = x - floorX;
Vector128<float> num = frac + Vector128.Create(1.01749063e+01f);
num = Vector128.FusedMultiplyAdd(
num,
frac,
Vector128.Create(4.88687798e+01f));
num = Vector128.FusedMultiplyAdd(
num,
frac,
Vector128.Create(9.85506591e+01f));
num *= exp;
Vector128<float> den =
Vector128.FusedMultiplyAdd(
frac,
Vector128.Create(2.10242958e-01f),
Vector128.Create(-2.22328856e-02f));
den = Vector128.FusedMultiplyAdd(
den,
frac,
Vector128.Create(-1.94414990e+01f));
den = Vector128.FusedMultiplyAdd(
den,
frac,
Vector128.Create(9.85506633e+01f));
return num / den;
}
// Raises 2 to the power of x
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static Vector<float> FastPow2f(Vector<float> x)
{
Vector<float> floorX = Vector.Floor(x);
Vector<int> exponent = Vector.ConvertToInt32(floorX) + Vector.Create(127);
Vector<float> exp = exponent.As<int, float>() << 23;
Vector<float> frac = x - floorX;
Vector<float> num = frac + Vector.Create(1.01749063e+01f);
num = Vector.FusedMultiplyAdd(
num,
frac,
Vector.Create(4.88687798e+01f));
num = Vector.FusedMultiplyAdd(
num,
frac,
Vector.Create(9.85506591e+01f));
num *= exp;
Vector<float> den =
Vector.FusedMultiplyAdd(
frac,
Vector.Create(2.10242958e-01f),
Vector.Create(-2.22328856e-02f));
den = Vector.FusedMultiplyAdd(
den,
frac,
Vector.Create(-1.94414990e+01f));
den = Vector.FusedMultiplyAdd(
den,
frac,
Vector.Create(9.85506633e+01f));
return num / den;
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static Vector128<float> FastPowf(
Vector128<float> @base,
Vector128<float> exponent)
=> FastPow2f(Vector128.Log2(@base) * exponent);
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static Vector<float> FastPowf(
Vector<float> @base,
Vector<float> exponent)
=> FastPow2f(Vector.Log2(@base) * exponent);
/// <summary>
/// Incrementing values to compute the Iota function.
/// </summary>

4
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlMaTreeLookup.cs

@ -3,7 +3,7 @@
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
internal sealed class JxlMaTreeLookup(JxlFlatDecisionNode[] nodes)
internal sealed class JxlMaTreeLookup(List<JxlFlatDecisionNode> nodes)
{
public JxlMaTreeLookupResult Lookup(Span<int> properties)
{
@ -12,7 +12,7 @@ internal sealed class JxlMaTreeLookup(JxlFlatDecisionNode[] nodes)
{
for (int i = 0; i < 2; i++)
{
JxlFlatDecisionNode node = nodes[pos];
JxlFlatDecisionNode node = nodes[(int)pos];
if (node.Property0 < 0)
{
return new(node.ChildID, node.Predictor, node.PredictorOffset, node.Multiplier);

66
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlGroupHeader.cs

@ -0,0 +1,66 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using System.Runtime.InteropServices;
using SixLabors.ImageSharp.Formats.Jxl.Fields;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Transforms;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal sealed class JxlGroupHeader : IJxlFields
{
private bool useGlobalTree;
public JxlGroupHeader() => JxlBundle.Init(this);
public bool UseGlobalTree
{
get => this.useGlobalTree;
set => this.useGlobalTree = value;
}
internal JxlModularHeader WeightedHeader { get; set; } = new();
internal List<JxlTransform> Transforms { get; private set; } = [];
public bool Visit(JxlVisitor visitor)
{
if (!visitor.Boolean(false, ref this.useGlobalTree))
{
return false;
}
if (!visitor.VisitNested(this.WeightedHeader))
{
return false;
}
uint numTransforms = (uint)this.Transforms.Count;
_ = visitor.U32(
JxlFieldExpressions.Value(0),
JxlFieldExpressions.Value(1),
JxlFieldExpressions.BitsOffset(4, 2),
JxlFieldExpressions.BitsOffset(8, 18),
0,
ref numTransforms);
if (visitor.IsReading)
{
this.Transforms = new List<JxlTransform>((int)numTransforms);
}
Span<JxlTransform> sp = CollectionsMarshal.AsSpan(this.Transforms);
for (int i = 0; i < numTransforms; i++)
{
if (!visitor.VisitNested(sp[i]))
{
return false;
}
}
return true;
}
}

16
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs

@ -0,0 +1,16 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal static class JxlMaConstants
{
/// <summary>
/// Total number of MA tree contexts.
/// </summary>
public const int NumTreeContexts = 6;
public const int MaxTreeSize = 1 << 22;
public const int PropertyRangeFast = 512 << 4;
}

19
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs

@ -0,0 +1,19 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal enum JxlMaTreeContext
{
SplitValue,
Property,
Predictor,
Offset,
MultiplierLog,
MultiplierBits
}

866
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlModularEncoding.cs

@ -0,0 +1,866 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using System.Runtime.CompilerServices;
using System.Runtime.InteropServices;
using SixLabors.ImageSharp.Formats.Jxl.Fields;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Decoder;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Transforms;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Primitives;
using FlatTree = System.Collections.Generic.List<SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction.JxlFlatDecisionNode>;
using Tree = System.Collections.Generic.List<SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.JxlPropertyDecisionNode>;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal static class JxlModularEncoding
{
public static bool TreeToLookupTable<T>(FlatTree tree, JxlTreeLut<T> lut)
{
bool hasOffsets = lut.Offsets.Length > 0;
bool hasMultipliers = lut.Multipliers.Length > 0;
Stack<TreeRange> ranges = [];
ranges.Push(new(-JxlMaConstants.PropertyRangeFast - 1, JxlMaConstants.PropertyRangeFast - 1, 0));
while (ranges.Count > 0)
{
TreeRange cur = ranges.Peek();
_ = ranges.Pop();
if (cur.Begin < -JxlMaConstants.PropertyRangeFast - 1 || cur.begin >= JxlMaConstants.PropertyRangeFast - 1 || cur.end > JxlMaConstants.PropertyRangeFast - 1)
{
// Tree is outside the allowed range, exit.
return false;
}
JxlMaNode node = tree[cur.Pos];
// Leaf.
if (node.Property0 == -1)
{
if (node.PredictorOffset is < sbyte.MinValue or > sbyte.MaxValue)
{
return false;
}
if (node.Multiplier is < sbyte.MinValue or > sbyte.MaxValue)
{
return false;
}
if (!hasMultipliers && node.Multiplier != 1)
{
return false;
}
if (!hasOffsets && node.PredictorOffset != 0)
{
return false;
}
for (int i = cur.Begin + 1; i < cur.End + 1; i++)
{
lut.ContextLookup[i + JxlMaConstants.PropertyRangeFast] = node.ChildID;
if (hasMultipliers)
{
lut.Multipliers[i + JxlMaConstants.PropertyRangeFast] = node.Multiplier;
}
if (hasOffsets)
{
lut.Offsets[i + JxlMaConstants.PropertyRangeFast] = node.PredictorOffset;
}
}
continue;
}
if (node.Properties[0] >= NumStaticProperties)
{
ranges.Push(new(node.SplitValues[0], cur.End, node.ChildID));
ranges.Push(new(node.SplitValue0, node.SplitValues[0], node.ChildID + 1));
}
else
{
ranges.Push(new(node.SplitValue0, cur.End, node.ChildID));
}
// <= side
if (node.Properties[1] >= NumStaticProperties)
{
ranges.Push(
new(node.SplitValues[1], node.SplitValue0, node.ChildID + 2));
ranges.Push(
new(cur.Begin, node.SplitValues[1], node.ChildID + 3));
}
else
{
ranges.Push(new(cur.Begin, node.SplitValue0, node.ChildID + 2));
}
}
return true;
}
public static FlatTree FilterTree(Tree globalTree, Span<int> staticProps, ref int numProps, ref bool useWp, ref bool wpOnly, ref bool gradientOnly)
{
numProps = 0;
bool hasWeightedPrediction = false;
bool hasNonWeighted = false;
gradientOnly = true;
void MarkProperty(int p, ref bool gradientOnly)
{
if (p == WpProp)
{
hasWeightedPrediction = true;
}
else if (p >= NumStaticProperties)
{
hasNonWeighted = true;
}
if (p >= NumStaticProperties && p != GradientProp)
{
gradientOnly = false;
}
}
FlatTree output = [];
Queue<int> nodes = [];
nodes.Enqueue(0);
while (nodes.Count > 0)
{
int cur = nodes.Peek();
_ = nodes.Dequeue();
while (globalTree[cur].Property < NumStaticProperties && globalTree[cur].Property != -1)
{
if (staticProps[globalTree[cur].Property] > globalTree[cur].SplitValue)
{
cur = globalTree[cur].LeftChild;
}
else
{
cur = globalTree[cur].RightChild;
}
}
JxlFlatDecisionNode flat = default;
if (globalTree[cur].Property == -1)
{
flat.Property0 = -1;
flat.ChildID = (uint)globalTree[cur].LeftChild;
flat.Predictor = globalTree[cur].Predictor;
flat.PredictorOffset = (int)globalTree[cur].PredictorOffset;
flat.Multiplier = globalTree[cur].Multiplier;
gradientOnly &= flat.Predictor == JxlPredictor.Gradient;
hasWeightedPrediction |= flat.Predictor == JxlPredictor.Weighted;
hasNonWeighted |= flat.Predictor != JxlPredictor.Weighted;
output.Add(flat);
continue;
}
flat.ChildID = (uint)(output.Count + nodes.Count + 1);
flat.Property0 = globalTree[cur].Property;
numProps = Math.Max(flat.Property0 + 1, numProps);
flat.SplitValue0 = globalTree[cur].SplitValue;
for (int i = 0; i < 2; i++)
{
int currentChild = i == 0 ? globalTree[cur].LeftChild : globalTree[cur].RightChild;
while (globalTree[currentChild].Property < kNumStaticProperties && globalTree[currentChild].Property != -1)
{
if (staticProps[globalTree[currentChild].Property] >
globalTree[currentChild].SplitValue)
{
currentChild = globalTree[currentChild].LeftChild;
}
else
{
currentChild = globalTree[currentChild].RightChild;
}
}
if (globalTree[currentChild].Property == -1)
{
flat.Properties[i] = 0;
flat.SplitValues[i] = 0;
nodes.Enqueue(currentChild);
nodes.Enqueue(currentChild);
}
else
{
flat.Properties[i] = globalTree[currentChild].Property;
flat.SplitValues[i] = globalTree[currentChild].SplitValue;
nodes.Enqueue(globalTree[currentChild].LeftChild);
nodes.Enqueue(globalTree[currentChild].RightChild);
numProps = Math.Max(flat.Properties[i] + 1, numProps);
}
}
for (int i = 0; i < 2; i++)
{
short property = flat.Properties[i];
MarkProperty(property, ref gradientOnly);
}
MarkProperty(flat.Property0, ref gradientOnly);
output.Add(flat);
}
if (numProps > JxlMaConstants.NumTreeContexts)
{
numProps = (JxlMath.DivCeil(numProps - JxlMaConstants.NumTreeContexts, kExtraPropsPerChannel) * kExtraPropsPerChannel) + JxlMaConstants.NumTreeContexts;
}
else
{
numProps = JxlMaConstants.NumTreeContexts;
}
useWp = hasWeightedPrediction;
wpOnly = hasWeightedPrediction && !hasNonWeighted;
return output;
}
public static void DecodeModularChannelMAANS(Configuration configuration, bool usesLz77, JxlBitReader bitReader, JxlAnsSymbolReader reader, List<byte> contextMap, Tree globalTree, JxlModularHeader wpHeader, int channelIdx, int groupId, JxlTreeLut<byte> treeLookup, JxlModularImage image, ref uint flRun, ref uint flV)
{
JxlModularChannel channel = image.Channels[channelIdx];
InlineArray2<int> staticProps = default;
staticProps[0] = channelIdx;
staticProps[1] = groupId;
if ((channel.Width & channel.Height) == 0) // equivalent to (channel.Width == 0 || channel.Height == 0)
{
return;
}
bool treeHasWeightedPredictionPropOrPred = false;
bool isWeightedPredictionOnly = false;
bool isGradientOnly = false;
int numProps = 0;
FlatTree tree = FilterTree(globalTree, staticProps, ref numProps, ref treeHasWeightedPredictionPropOrPred, ref isWeightedPredictionOnly, ref isGradientOnly);
Span<JxlFlatDecisionNode> span = CollectionsMarshal.AsSpan(tree);
for (int i = 0; i < span.Length; i++)
{
ref JxlFlatDecisionNode node = ref span[i];
if (node.Property0 == -1)
{
node.ChildID = contextMap[(int)node.ChildID];
}
}
int MakePixel(uint v, int multiplier, int offset)
{
int value = JxlPackSigned.UnpackSigned(v);
return (value * multiplier) + offset;
}
bool GlobalTreeIsAllGradientNoOp()
{
foreach (JxlPropertyDecisionNode n in globalTree)
{
if (n.Property == -1)
{
if (n.Predictor != JxlPredictor.Gradient || n.PredictorOffset != 0 || n.Multiplier != 1)
{
return false;
}
}
else if (n.Property >= NumStaticProperties)
{
return false;
}
}
return true;
}
if (tree.Count == 1)
{
JxlPredictor predictor = tree[0].Predictor;
int offset = tree[0].PredictorOffset;
int multiplier = tree[0].Multiplier;
uint ctx_id = tree[0].ChildID;
if (predictor == JxlPredictor.Zero)
{
if (reader.IsSingleValueAndAdvance(ctx_id, channel.Width * channel.Height, out uint value))
{
int v = MakePixel(value, multiplier, offset);
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
r[..channel.Width].Fill(v);
}
}
else
{
if (multiplier == 1 && offset == 0)
{
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
for (int x = 0; x < channel.Width; x++)
{
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, ctx_id, bitReader);
r[x] = JxlPackSigned.UnpackSigned(v);
}
}
}
else
{
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
for (int x = 0; x < channel.Width; x++)
{
uint v = reader.ReadHybridUintClusteredMaybeInlined(usesLz77, ctx_id, bitReader);
r[x] = MakePixel(v, multiplier, offset);
}
}
}
}
return;
}
else if (usesLz77 && reader.IsHuffRleOnly && GlobalTreeIsAllGradientNoOp())
{
int sv = JxlPackSigned.UnpackSigned(flV);
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
Span<int> rtop = y > 0 ? channel.GetRow(y - 1) : channel.GetRowMinus(y, 1);
Span<int> rtopleft = y > 0 ? channel.GetRowMinus(y - 1, 1) : channel.GetRowMinus(y, 1);
int guess0 = y > 0 ? rtop[0] : 0;
if (flRun == 0)
{
reader.ReadHybridUintClusteredHuffRleOnly(ctx_id, bitReader, ref flV, ref flRun);
sv = JxlPackSigned.UnpackSigned(flV);
}
else
{
flRun--;
}
r[0] = sv + guess0;
for (int x = 1; x < channel.Width; x++)
{
int left = r[x - 1];
int top = rtop[x];
int topleft = rtopleft[x];
int guess = JxlContextPrediction.ClampedGradient(top, left, topleft);
if (flRun == 0)
{
reader.ReadHybridUintClusteredHuffRleOnly(ctx_id, bitReader, ref flV, ref flRun);
sv = JxlPackSigned.UnpackSigned(flV);
}
else
{
flRun--;
}
r[x] = sv + guess;
}
}
return;
}
else if (predictor == JxlPredictor.Gradient && offset == 0 && multiplier == 1)
{
int onerow = channel.Plane.PixelsPerRow;
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
for (int x = 0; x < channel.Width; x++)
{
// Neighbors
int left = x > 0 ? r[x - 1] : y > 0 ? channel.GetRowMinus(y, x - onerow)[0] : 0;
int top = y > 0 ? channel.GetRowMinus(y, x - onerow)[0] : left;
int topleft = x > 0 && y > 0 ? channel.GetRowMinus(y, x - 1 - onerow)[0] : left;
int guess = JxlContextPrediction.ClampedGradient(top, left, topleft);
uint v = reader.ReadHybridUintClusteredMaybeInlined>(
usesLz77,
ctx_id,
bitReader);
r[x] = MakePixel(v, 1, guess);
}
}
return;
}
}
if (isWeightedPredictionOnly)
{
isWeightedPredictionOnly = TreeToLookupTable(tree, treeLookup);
}
if (isGradientOnly)
{
isGradientOnly = TreeToLookupTable(tree, treeLookup);
}
if (isGradientOnly)
{
int onerow = channel.Plane.PixelsPerRow;
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
for (int x = 0; x < channel.Width; x++)
{
// Neighbors
int left = x > 0 ? channel.GetRowPlus(y, x - 1)[0] : y > 0 ? channel.GetRowPlus(y, x - onerow)[0] : 0;
int top = y > 0 ? channel.GetRowPlus(y, x - onerow)[0] : left;
int topleft = x > 0 && y > 0 ? channel.GetRowPlus(y, x - 1 - onerow)[0] : left;
int guess = JxlContextPrediction.ClampedGradient(top, left, topleft);
int pos =
JxlMaConstants.PropertyRangeFast +
Math.Min(
Math.Max(-JxlMaConstants.PropertyRangeFast, top + left - topleft),
JxlMaConstants.PropertyRangeFast - 1);
uint ctx_id = treeLookup.ContextLookup[pos];
uint v = reader.ReadHybridUintClusteredMaybeInlined(usesLz77, ctx_id, br);
r[x] = MakePixel(v, 1, guess);
}
}
}
else if (!usesLz77 && isWeightedPredictionOnly && channel.Width > 8)
{
JxlModularState wpState = new(wpHeader, channel.Width);
Span<int> properties = [0];
for (int y = 0; y < channel.Height; y++)
{
Span<int> r = channel.GetRow(y);
Span<int> rtop = y > 0 ? channel.GetRow(y - 1) : channel.GetRowMinus(y, 1);
Span<int> rtoptop = y > 1 ? channel.GetRow(y - 2) : rtop;
Span<int> rtopleft = y > 0 ? channel.GetRowMinus(y - 1, 1) : channel.GetRowMinus(y, 1);
Span<int> rtopright = y > 0 ? channel.GetRow(y - 1)[1..] : channel.GetRowMinus(y, 1);
int x = 0;
{
int offset = 0;
int left = y > 0 ? rtop[x] : 0;
int toptop = y > 0 ? rtoptop[x] : 0;
int topright = x + 1 < channel.Width && y > 0 ? rtop[x + 1] : left;
int guess = (int)wpState.Predict(true, x, y, channel.Width, left, left, topright, left, toptop, properties, offset);
int pos = JxlMaConstants.PropertyRangeFast + Math.Clamp(properties[0], -JxlMaConstants.PropertyRangeFast, JxlMaConstants.PropertyRangeFast - 1);
uint ctx_id = treeLookup.ContextLookup[pos];
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, ctx_id, bitReader);
r[x] = MakePixel(v, 1, guess);
wpState.UpdatePredictionErrors(r[x], x, y, channel.Width);
}
for (x = 1; x + 1 < channel.Width; x++)
{
int offset = 0;
int guess = (int)wpState.Predict(true, x, y, channel.Width, rtop[x], r[x - 1], rtopright[x], rtopleft[x], rtoptop[x], properties, offset);
int pos = JxlMaConstants.PropertyRangeFast + Math.Clamp(properties[0], -JxlMaConstants.PropertyRangeFast, JxlMaConstants.PropertyRangeFast - 1);
int ctx_id = treeLookup.ContextLookup[pos];
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, ctx_id, bitReader);
r[x] = MakePixel(v, 1, guess);
wpState.UpdatePredictionErrors(r[x], x, y, channel.Width);
}
{
int offset = 0;
int guess = (int)wpState.Predict(true, x, y, channel.Width, rtop[x], r[x - 1], rtop[x], rtopleft[x], rtoptop[x], properties, offset);
int pos = JxlMaConstants.PropertyRangeFast + Math.Clamp(properties[0], -JxlMaConstants.PropertyRangeFast, JxlMaConstants.PropertyRangeFast - 1);
int ctx_id = treeLookup.ContextLookup[pos];
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, ctx_id, bitReader);
r[x] = MakePixel(v, 1, guess);
wpState.UpdatePredictionErrors(r[x], x, y, channel.Width);
}
}
}
else if (!treeHasWeightedPredictionPropOrPred)
{
JxlMaTreeLookup tree_lookup = new(tree);
Span<int> properties = stackalloc int[numProps];
int onerow = channel.Plane.PixelsPerRow;
JxlModularChannel references = new(configuration, properties.Length - NumNonrefProperties, channel.Width, 0, 0);
for (int y = 0; y < channel.Height; y++)
{
Span<int> p = channel.GetRow(y);
PrecomputeReferences(channel, y, image, channelIdx, references);
InitPropsRow(properties, staticProps, y);
if (y > 1 && channel.Width > 8 && references.Width == 0)
{
for (int x = 0; x < 2; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeNoWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references);
uint v = reader.ReadHybridUintClustered(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
}
for (int x = 2; x < channel.Width - 2; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeNoWeightedPredictionNoEdgeCases(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references);
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
}
for (int x = channel.Width - 2; x < channel.Width; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeNoWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references);
uint v = reader.ReadHybridUintClustered(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
}
}
else
{
for (int x = 0; x < channel.Width; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeNoWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references);
uint v = reader.ReadHybridUintClusteredMaybeInlined(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
}
}
}
}
else
{
JxlMaTreeLookup tree_lookup = new(tree);
Span<int> properties = stackalloc int[numProps];
int onerow = channel.Plane.PixelsPerRow;
JxlModularChannel references = new(configuration, properties.Length - NumNonrefProperties, channel.Width, 0, 0);
JxlModularState wpState = new(wpHeader, channel.Width);
for (int y = 0; y < channel.Height; y++)
{
Span<int> p = channel.GetRow(y);
InitPropsRow(properties, staticProps, y);
PrecomputeReferences(channel, y, image, channelIdx, references);
if (!usesLz77 && y > 1 && channel.Width > 8 && references.Width == 0)
{
for (int x = 0; x < 2; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references, wpState);
uint v = reader.ReadHybridUintClustered(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
wpState.UpdatePredictionErrors(p[x], x, y, channel.Width);
}
for (int x = 2; x < channel.Width - 2; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeWeightedPredictionNoEdgeCases(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references, wpState);
uint v = reader.ReadHybridUintClusteredInlined(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
wpState.UpdatePredictionErrors(p[x], x, y, channel.Width);
}
for (int x = channel.Width - 2; x < channel.Width; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references, wpState);
uint v = reader.ReadHybridUintClustered(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
wpState.UpdatePredictionErrors(p[x], x, y, channel.Width);
}
}
else
{
for (int x = 0; x < channel.Width; x++)
{
JxlPredictionResult res = JxlContextPrediction.PredictTreeWeightedPrediction(properties, channel.Width, ref p[x], onerow, x, y, tree_lookup, references, wpState);
uint v = reader.ReadHybridUintClustered(usesLz77, res.Context, bitReader);
p[x] = MakePixel(v, res.Multiplier, res.Guess);
wpState.UpdatePredictionErrors(p[x], x, y, channel.Width);
}
}
}
}
}
public static void DecodeModularChannelMAANS(JxlBitReader bitReader, JxlAnsSymbolReader reader, List<byte> contextMap, Tree globalTree, JxlModularHeader wpHeader, int channelIdx, int groupId, JxlTreeLut<byte> treeLut, JxlModularImage image, ref uint flRun, ref uint flV)
=> DecodeModularChannelMAANS(reader.UsesLz77, bitReader, reader, contextMap, globalTree, wpHeader, channelIdx, groupId, treeLut, image, ref flRun, ref flV);
public static void ValidateChannelDimensions(JxlModularImage image, JxlModularOptions options)
{
int nbChannels = image.Channels.Count;
foreach (bool isDc in (bool[])[true, false])
{
int groupDimensions = options.group_dim * (isDc ? JxlFrameDimensions.BlockDimensions : 1);
int c = image.MetaChannels;
for (; c < nbChannels; c++)
{
JxlModularChannel ch = image.Channels[c];
if (ch.Width > options.group_dim || ch.Height > options.group_dim)
{
break;
}
}
for (; c < nbChannels; c++)
{
JxlModularChannel ch = image.Channels[c];
if (ch.Width == 0 || ch.Height == 0)
{
continue;
}
bool isDcChannel = Math.Min(ch.HorizontalShift, ch.VerticalShift) >= 3;
if (isDcChannel != isDc)
{
continue;
}
int tileDimensions = groupDimensions >> Math.Max(ch.HorizontalShift, ch.VerticalShift);
if (tileDimensions == 0)
{
throw new InvalidOperationException("Inconsistent transforms");
}
}
}
}
public static bool ModularDecode(Configuration configuration, JxlBitReader br, JxlModularImage image, JxlGroupHeader header, int groupId, JxlModularOptions options, Tree globalTree, JxlAnsCode globalCode, List<byte> globalContextMap, bool allowTruncatedGroup)
{
if (image.Channels.Count == 0)
{
return true;
}
if (!JxlBundle.Read(br, header) && !allowTruncatedGroup)
{
throw new InvalidOperationException("Could not read a bundle");
}
image.Transforms = header.Transforms;
foreach (JxlTransform transform in image.Transforms)
{
transform.MetaApply(image);
}
if (image.IsInvalid)
{
throw new InvalidOperationException("Corrupt file. Aborting.");
}
ValidateChannelDimensions(image, options);
int numberOfChannels = image.Channels.Count;
int numChannels = 0;
int distanceMultiplier = 0;
for (int i = 0; i < numberOfChannels; i++)
{
JxlModularChannel channel = image.Channels[i];
if (i >= image.MetaChannels && (channel.Width > options.MaxChannelSize || channel.Height > options.MaxChannelSize))
{
break;
}
if (channel.Width == 0 || channel.Height == 0)
{
continue; // skip empty channels
}
if (channel.Width > distanceMultiplier)
{
distanceMultiplier = channel.Width;
}
numChannels++;
}
if (numChannels == 0)
{
return true;
}
int nextChannel = 0;
using JxlScopeGuard clearGuard = new(() =>
{
for (int c = nextChannel; c < image.Channels.Count; c++)
{
image.Channels[c].Plane.Clear();
}
});
if (allowTruncatedGroup)
{
clearGuard.Disarm();
}
Tree treeStorage = [];
List<byte> contextMapStorage = [];
JxlAnsCode codeStorage = new();
Tree tree = treeStorage;
JxlAnsCode code = codeStorage;
List<byte> contextMap = contextMapStorage;
if (!header.UseGlobalTree)
{
ulong maxTreeSize = 1024;
for (int i = 0; i < numberOfChannels; i++)
{
JxlModularChannel channel = image.Channels[i];
if (i >= image.MetaChannels && (channel.Width > options.MaxChannelSize || channel.Height > options.MaxChannelSize))
{
break;
}
ulong pixels = (ulong)channel.Width * (ulong)channel.Height;
maxTreeSize += pixels;
}
maxTreeSize = Math.Min(1uL << 20, maxTreeSize);
DecodeTree(configuration, br, treeStorage, maxTreeSize));
DecodeHistograms(
configuration,
br,
(treeStorage.Count + 1) / 2,
codeStorage,
contextMapStorage);
}
else
{
if (globalTree.Count == 0)
{
throw new InvalidOperationException("No global tree available but one was requested");
}
tree = globalTree;
code = globalCode;
contextMap = globalContextMap;
}
JxlAnsSymbolReader reader = new(code, br, distanceMultiplier);
JxlTreeLut<byte> treeLut = new(false, false);
uint flRun = 0;
uint flV = 0;
for (; nextChannel < numberOfChannels; nextChannel++)
{
JxlModularChannel channel = image.Channels[nextChannel];
if (nextChannel >= image.MetaChannels &&
(channel.Width > options.MaxChannelSize ||
channel.Height > options.MaxChannelSize))
{
break;
}
if (channel.Width == 0 || channel.Height == 0)
{
continue; // skip empty channels
}
DecodeModularChannelMAANS(br, reader, contextMap, tree, header.WeightedHeader, nextChannel, groupId, treeLut, image, ref flRun, ref flV);
}
clearGuard.Disarm();
if (!reader.CheckAnsFinalState())
{
throw new InvalidOperationException("ANS decode final state failed");
}
return true;
}
public static bool ModularGenericDecompress(Configuration configuration. JxlBitReader br, JxlModularImage image, JxlGroupHeader header, int groupId, JxlModularOptions options, bool undoTransforms, Tree tree, JxlAnsCode code, List<byte> contextMap, bool allowTruncatedGroup)
{
List<Size> reqSizes = new(capacity: image.Channels.Count);
foreach (JxlModularChannel c in image.Channels)
{
reqSizes.Add(new(c.Width, c.Height));
}
header ??= new JxlGroupHeader();
bool decStatus = ModularDecode(configuration, br, image, header, groupId, options, tree, code, contextMap, allowTruncatedGroup);
if (!allowTruncatedGroup)
{
if (!decStatus)
{
return false;
}
}
if (undoTransforms)
{
image.UndoTransforms(header.WeightedHeader);
}
if (image.IsInvalid)
{
throw new InvalidOperationException("Corrupt file. Aborting");
}
if (undoTransforms)
{
if (image.Channels.Count != reqSizes.Count)
{
return false;
}
for (int c = 0; c < reqSizes.Count; c++)
{
if (reqSizes[c].Width != image.Channels[c].Width || reqSizes[c].Height != image.Channels[c].Height)
{
return false;
}
}
}
return true;
}
private readonly record struct TreeRange(int Begin, int End, int Pos);
}

46
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlPropertyDecisionNode.cs

@ -0,0 +1,46 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal struct JxlPropertyDecisionNode
{
public int SplitValue;
public int Property;
public int LeftChild;
public int RightChild;
public JxlPredictor Predictor;
public long PredictorOffset;
public int Multiplier;
public JxlPropertyDecisionNode(int property, int splitValue, int leftChild, int rightChild, JxlPredictor predictor, long predictorOffset, int multiplier)
{
this.SplitValue = splitValue;
this.Property = property;
this.LeftChild = leftChild;
this.RightChild = rightChild;
this.Predictor = predictor;
this.PredictorOffset = predictorOffset;
this.Multiplier = multiplier;
}
public JxlPropertyDecisionNode()
: this(0, -1, 0, 0, JxlPredictor.Zero, 0, 1)
{
}
public static JxlPropertyDecisionNode Leaf(JxlPredictor predictor, long offset = 0, int multiplier = 1)
=> new(-1, 0, 0, 0, predictor, offset, multiplier);
public static JxlPropertyDecisionNode Split(int p, int splitValue, int leftChild, int rightChild = -1)
{
if (rightChild == -1)
{
rightChild = leftChild + 1;
}
return new(p, splitValue, leftChild, rightChild, JxlPredictor.Zero, 0, 1);
}
}

13
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeLut.cs

@ -0,0 +1,13 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal sealed class JxlTreeLut<T>(bool hasOffsets, bool hasMultipliers)
{
public T[] ContextLookup { get; } = new T[2 * JxlMaConstants.PropertyRangeFast];
public byte[] Offsets { get; } = new byte[hasOffsets ? (2 * JxlMaConstants.PropertyRangeFast) : 0];
public byte[] Multipliers { get; } = new byte[hasMultipliers ? (2 * JxlMaConstants.PropertyRangeFast) : 0];
}

4
src/ImageSharp/Formats/Jxl/Processing/Modular/JxlModularChannel.cs

@ -74,4 +74,8 @@ internal sealed class JxlModularChannel
}
public Span<int> GetRow(int y) => this.Plane.GetRow(y);
public Span<int> GetRowMinus(int y, int minus) => this.Plane.GetRowMinus(y, minus);
public Span<int> GetRowPlus(int y, int plus) => this.Plane.GetRowPlus(y, plus);
}

4
src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlDequantMatrices.cs

@ -112,12 +112,12 @@ internal sealed class JxlDequantMatrices
/// <summary>
/// Gets a lookup which represents required widths for each quantizer.
/// </summary>
private static ReadOnlySpan<int> RequiredSizeX => [1, 1, 1, 1, 2, 4, 1, 1, 2, 1, 1, 8, 4, 16, 8, 32, 16];
public static ReadOnlySpan<int> RequiredSizeX => [1, 1, 1, 1, 2, 4, 1, 1, 2, 1, 1, 8, 4, 16, 8, 32, 16];
/// <summary>
/// Gets a lookup which represents required heights for each quantizer.
/// </summary>
private static ReadOnlySpan<int> RequiredSizeY => [1, 1, 1, 1, 2, 4, 2, 4, 4, 1, 1, 8, 8, 16, 16, 32, 32];
public static ReadOnlySpan<int> RequiredSizeY => [1, 1, 1, 1, 2, 4, 2, 4, 4, 1, 1, 8, 8, 16, 16, 32, 32];
/// <summary>
/// Returns the default library with quantizer encodings for all transforms

700
src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlQuantWeights.cs

@ -1,7 +1,16 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using System.Buffers;
using System.Numerics;
using System.Numerics.Tensors;
using System.Runtime.CompilerServices;
using System.Runtime.Intrinsics;
using SixLabors.ImageSharp.Formats.Jxl.Fields;
using SixLabors.ImageSharp.Formats.Jxl.Processing.AcStrategy;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Dct;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Decoder;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Primitives;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Quantization;
@ -15,6 +24,8 @@ internal static class JxlQuantWeights
public const int Log2NumQuantModes = 3;
private const float AlmostZero = 1e-8f;
/// <summary>
/// DCT quantizer encoding. (6 distance bands)
/// </summary>
@ -526,4 +537,693 @@ internal static class JxlQuantWeights
]
],
8));
private static ReadOnlySpan<float> AfvFrequencies =>
[
0xBAD,
0xBAD,
0.8517778890324296f,
5.37778436506804f,
0xBAD,
0xBAD,
4.734747904497923f,
5.449245381693219f,
1.6598270267479331f,
4,
7.275749096817861f,
10.423227632456525f,
2.662932286148962f,
7.630657783650829f,
8.962388608184032f,
12.97166202570235f,
];
private static Vector128<float> Gather(ReadOnlySpan<float> data, Vector128<int> indices)
{
Vector128<float> result = Vector128<float>.Zero;
for (int i = 0; i < 4; i++)
{
result = result.WithElement(i, data[indices[i]]);
}
return result;
}
public static void GetQuantWeightsDCT2(float[][] dct2Weights, Span<float> weights)
{
for (int c = 0; c < 3; c++)
{
int start = c * 64;
weights[start] = 0xBAD;
weights[start + 1] = weights[start + 8] = dct2Weights[c][0];
weights[start + 9] = dct2Weights[c][1];
for (int y = 0; y < 2; y++)
{
for (int x = 0; x < 2; x++)
{
weights[start + (y * 8) + x + 2] = dct2Weights[c][2];
weights[start + ((y + 2) * 8) + x] = dct2Weights[c][2];
}
}
for (int y = 0; y < 2; y++)
{
for (int x = 0; x < 2; x++)
{
weights[start + ((y + 2) * 8) + x + 2] = dct2Weights[c][3];
}
}
for (int y = 0; y < 4; y++)
{
for (int x = 0; x < 4; x++)
{
weights[start + (y * 8) + x + 4] = dct2Weights[c][4];
weights[start + ((y + 4) * 8) + x] = dct2Weights[c][4];
}
}
for (int y = 0; y < 4; y++)
{
for (int x = 0; x < 4; x++)
{
weights[start + ((y + 4) * 8) + x + 4] = dct2Weights[c][5];
}
}
}
}
public static void GetQuantWeightsIdentity(float[][] idWeights, Span<float> weights)
{
for (int c = 0; c < 3; c++)
{
int c64 = 64 * c;
for (int i = 0; i < 64; i++)
{
weights[c64 + i] = idWeights[c][0];
}
weights[c64 + 1] = idWeights[c][1];
weights[c64 + 8] = idWeights[c][1];
weights[c64 + 9] = idWeights[c][2];
}
}
public static float Interpolate(float pos, float max, Span<float> array, int len)
{
float scaledPos = pos * (len - 1) / max;
int idx = (int)scaledPos;
float a = array[idx];
float b = array[idx + 1];
return a * MathF.Pow(b / a, scaledPos - idx);
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static float Mult(float v) => v > 0f
? 1.0f + v
: 1.0f / (1.0f - v);
public static Vector128<float> InterpolateVec(Vector128<float> scaledPos, ReadOnlySpan<float> array)
{
Vector128<int> idx = Vector128.ConvertToInt32(scaledPos);
Vector128<float> frac = scaledPos - Vector128.ConvertToSingle(idx);
Vector128<float> a = Gather(array, idx);
Vector128<float> b = Gather(array[1..], idx);
return a * JxlSimdUtils.FastPowf(b / a, frac);
}
public static bool GetQuantWeights(
int rows,
int cols,
float[][] distanceBands,
int numBands,
Span<float> output)
{
Span<float> bands = stackalloc float[JxlQuantizerConstants.MaxDistanceBands];
for (int c = 0; c < 3; c++)
{
bands[0] = distanceBands[c][0];
if (bands[0] < AlmostZero)
{
return false;
}
for (int i = 1; i < numBands; i++)
{
bands[i] = bands[i - 1] * distanceBands[c][i];
if (bands[i] < AlmostZero)
{
return false;
}
}
float scale = (numBands - 1) / (JxlDctScales.Sqrt2 + 1e-6f);
float rcpCol = scale / (cols - 1);
float rcpRow = scale / (rows - 1);
for (int y = 0; y < rows; y++)
{
float dy = y * rcpRow;
float dy2 = dy * dy;
for (int x = 0; x < cols; x += 4)
{
Vector128<float> dx =
Vector128.Create((float)x, x + 1, x + 2, x + 3)
* Vector128.Create(rcpCol);
Vector128<float> scaledDistance =
Vector128.Sqrt(
(dx * dx) + Vector128.Create(dy2));
Vector128<float> weight =
numBands == 1
? Vector128.Create(bands[0])
: InterpolateVec(scaledDistance, bands);
weight.CopyTo(output[((c * cols * rows) + (y * cols) + x)..]);
}
}
}
return true;
}
public static bool ComputeQuantTable(JxlQuantizerEncoding encoding, Span<float> table, Span<float> inverseTable, int tableNum, JxlQuantTable kind, ref int pos)
{
const int n = JxlFrameDimensions.BlockDimensions;
int quantSizeTable = (int)kind;
int wrows = 8 * JxlDequantMatrices.RequiredSizeX[quantSizeTable];
int wcols = 8 * JxlDequantMatrices.RequiredSizeY[quantSizeTable];
int num = wrows * wcols;
Span<float> weights = stackalloc float[3 * num];
switch (encoding.Mode)
{
case JxlQuantMode.Library:
{
// Library and copy quant encoding should get replaced by the actual
// parameters by the caller.
return false;
}
case JxlQuantMode.Id:
{
if (num != JxlFrameDimensions.DctBlockSize)
{
return false;
}
GetQuantWeightsIdentity(encoding.IdWeights!, weights);
break;
}
case JxlQuantMode.Dct2:
{
if (num != JxlFrameDimensions.DctBlockSize)
{
return false;
}
GetQuantWeightsDCT2(encoding.Dct2Weights!, weights);
break;
}
case JxlQuantMode.Dct4:
{
if (num != JxlFrameDimensions.DctBlockSize)
{
return false;
}
Span<float> weights4x4 = stackalloc float[3 * 4 * 4];
// Always use 4x4 GetQuantWeights for DCT4 quantization tables.
if (!GetQuantWeights(4, 4, encoding.DctParameters!.DistanceBands, encoding.DctParameters.NumDistanceBands, weights4x4))
{
return false;
}
for (int c = 0; c < 3; c++)
{
for (int y = 0; y < JxlFrameDimensions.BlockDimensions; y++)
{
for (int x = 0; x < JxlFrameDimensions.BlockDimensions; x++)
{
weights[(c * num) + (y * JxlFrameDimensions.BlockDimensions) + x] = weights4x4[(c * 16) + ((y / 2) * 4) + (x / 2)];
}
}
weights[(c * num) + 1] /= encoding.Dct4Multipliers![c][0];
weights[(c * num) + n] /= encoding.Dct4Multipliers[c][0];
weights[(c * num) + n + 1] /= encoding.Dct4Multipliers[c][1];
}
break;
}
case JxlQuantMode.Dct4x8:
{
if (num != JxlFrameDimensions.DctBlockSize)
{
return false;
}
Span<float> weights4x8 = stackalloc float[3 * 4 * 8];
if (!GetQuantWeights(4, 8, encoding.DctParameters!.DistanceBands, encoding.DctParameters.NumDistanceBands, weights4x8))
{
return false;
}
for (int c = 0; c < 3; c++)
{
for (int y = 0; y < JxlFrameDimensions.BlockDimensions; y++)
{
for (int x = 0; x < JxlFrameDimensions.BlockDimensions; x++)
{
weights[(c * num) + (y * JxlFrameDimensions.BlockDimensions) + x] = weights4x8[(c * 32) + ((y / 2) * 8) + x];
}
}
weights[(c * num) + n] /= encoding.Dct4x8Multipliers![c];
}
break;
}
case JxlQuantMode.Dct:
{
if (!GetQuantWeights(wrows, wcols, encoding.DctParameters!.DistanceBands, encoding.DctParameters.NumDistanceBands, weights))
{
return false;
}
break;
}
case JxlQuantMode.Raw:
{
if (encoding.QuantizationTable is null || encoding.QuantizationTable.Length != 3 * num)
{
throw new InvalidOperationException("Invalid raw quantizer table encoding");
}
Span<int> qtable = encoding.QuantizationTable.AsSpan();
for (int i = 0; i < 3 * num; i++)
{
weights[i] = 1f / (encoding.QuantizationTableDenominator * qtable[i]);
}
break;
}
case JxlQuantMode.Afv:
{
Span<float> weights4x8 = stackalloc float[3 * 4 * 8];
if (!GetQuantWeights(4, 8, encoding.DctParameters!.DistanceBands, encoding.DctParameters.NumDistanceBands, weights4x8))
{
return false;
}
Span<float> weights4x4 = stackalloc float[3 * 4 * 4];
if (!GetQuantWeights(4, 4, encoding.DctParametersAfv4x4!.DistanceBands, encoding.DctParametersAfv4x4.NumDistanceBands, weights4x4))
{
return false;
}
const float lo = 0.8517778890324296f;
const float hi = 12.97166202570235f - lo + 1e-6f;
Span<float> bands = [0, 0, 0, 0];
for (int c = 0; c < 3; c++)
{
bands[0] = encoding.AfvWeights![c][5];
if (bands[0] < AlmostZero)
{
throw new InvalidOperationException("Invalid AFV bands");
}
for (int i = 1; i < 4; i++)
{
bands[i] = bands[i - 1] * Mult(encoding.AfvWeights[c][i + 5]);
if (bands[i] < AlmostZero)
{
throw new InvalidOperationException("Invalid AFV bands");
}
}
int start = c * 64;
void SetWeight(int x, int y, float value, Span<float> weights) => weights[start + (y * 8) + x] = value;
weights[start] = 1;
SetWeight(0, 1, encoding.AfvWeights[c][0], weights);
SetWeight(1, 0, encoding.AfvWeights[c][1], weights);
SetWeight(0, 2, encoding.AfvWeights[c][2], weights);
SetWeight(2, 0, encoding.AfvWeights[c][3], weights);
SetWeight(2, 2, encoding.AfvWeights[c][4], weights);
// All other AFV weights.
for (int y = 0; y < 4; y++)
{
for (int x = 0; x < 4; x++)
{
if (x < 2 && y < 2)
{
continue;
}
float interpolatedVal = Interpolate(AfvFrequencies[(y * 4) + x] - lo, hi, bands, 4);
SetWeight(2 * x, 2 * y, interpolatedVal, weights);
}
}
// Put 4x8 weights in odd rows, except (1, 0).
for (int y = 0; y < JxlFrameDimensions.BlockDimensions / 2; y++)
{
for (int x = 0; x < JxlFrameDimensions.BlockDimensions; x++)
{
if (x == 0 && y == 0)
{
continue;
}
weights[(c * num) + (((2 * y) + 1) * JxlFrameDimensions.BlockDimensions) + x] = weights4x8[(c * 32) + (y * 8) + x];
}
}
// Put 4x4 weights in even rows / odd columns, except (0, 1).
for (int y = 0; y < JxlFrameDimensions.BlockDimensions / 2; y++)
{
for (int x = 0; x < JxlFrameDimensions.BlockDimensions / 2; x++)
{
if (x == 0 && y == 0)
{
continue;
}
weights[(c * num) + ((2 * y) * JxlFrameDimensions.BlockDimensions) + (2 * x) + 1] = weights4x4[(c * 16) + (y * 4) + x];
}
}
}
break;
}
}
int prevPos = pos;
// Don't zero-init
Span<float> invVal = stackalloc float[64];
Span<float> val = stackalloc float[64];
for (int i = 0; i < num * 3; i += 64)
{
weights.Slice(i, 64).CopyTo(invVal);
// TODO: there's an unlikely check right here in
// reference:
// if (JXL_UNLIKELY(!AllFalse(d, Ge(inv_val, Set(d, 1.0f / kAlmostZero))) ||
// !AllFalse(d, Lt(inv_val, Set(d, kAlmostZero)))))
// {
// throw new InvalidOperationException("Invalid quantization table");
// }
// should we trade performance for an unlikely check?
val.Fill(1.0f);
TensorPrimitives.Divide(val, invVal, val);
val.CopyTo(table.Slice(pos + i, 64));
invVal.CopyTo(inverseTable.Slice(pos + i, 64));
}
pos += 3 * num;
int xs = JxlDequantMatrices.RequiredSizeX[quantSizeTable];
int ys = JxlDequantMatrices.RequiredSizeY[quantSizeTable];
JxlForwardCoefficientOrder.CoefficientLayout(ref ys, ref xs);
for (int c = 0; c < 3; c++)
{
for (int y = 0; y < ys; y++)
{
for (int x = 0; x < xs; x++)
{
inverseTable[prevPos + (c * ys * xs * JxlFrameDimensions.DctBlockSize) + (y * JxlFrameDimensions.BlockDimensions * xs) + x] = 0;
}
}
}
return true;
}
public static bool DecodeDctParameters(JxlBitReader reader, JxlDctQuantWeightParameters parameters)
{
parameters.NumDistanceBands = (int)reader.ReadBits32(JxlQuantizerConstants.Log2MaxDistanceBands) + 1;
for (int c = 0; c < 3; c++)
{
for (int i = 0; i < parameters.NumDistanceBands; i++)
{
if (!JxlF16Coder.Read(reader, ref parameters.DistanceBands[c][i]))
{
return false;
}
}
if (parameters.DistanceBands[c][0] < AlmostZero)
{
throw new InvalidOperationException("Distance band seed is too small");
}
parameters.DistanceBands[c][0] *= 64f;
}
return true;
}
public static bool Decode(Configuration configuration, JxlBitReader br, JxlQuantizerEncoding encoding, int requiredSizeX, int requiredSizeY, int idx, JxlModularFrameDecoder modularFrameDecoder)
{
int requiredSize = requiredSizeX * requiredSizeY;
requiredSizeX *= JxlFrameDimensions.BlockDimensions;
requiredSizeY *= JxlFrameDimensions.BlockDimensions;
int mode = (int)br.ReadBits32(JxlQuantizerConstants.Log2NumQuantModes);
switch ((JxlQuantMode)mode)
{
case JxlQuantMode.Library:
{
encoding.Predefined = (byte)br.ReadBits32(CeilLog2NumPredefinedTables);
if (encoding.Predefined >= NumPredefinedTables)
{
throw new InvalidOperationException("Invalid predefined table");
}
break;
}
case JxlQuantMode.Id:
{
if (requiredSize != 1)
{
throw new InvalidOperationException("Invalid mode");
}
for (int c = 0; c < 3; c++)
{
for (int i = 0; i < 3; i++)
{
if (!JxlF16Coder.Read(br, ref encoding.IdWeights![c][i]))
{
return false;
}
if (Math.Abs(encoding.IdWeights[c][i]) < AlmostZero)
{
throw new InvalidOperationException("ID Quantizer is too small");
}
encoding.IdWeights[c][i] *= 64;
}
}
break;
}
case JxlQuantMode.Dct2:
{
if (requiredSize != 1)
{
throw new InvalidOperationException("Invalid mode");
}
for (int c = 0; c < 3; c++)
{
for (int i = 0; i < 6; i++)
{
if (!JxlF16Coder.Read(br, ref encoding.Dct2Weights![c][i]))
{
return false;
}
if (Math.Abs(encoding.Dct2Weights[c][i]) < AlmostZero)
{
throw new InvalidOperationException("Quantizer is too small");
}
encoding.Dct2Weights[c][i] *= 64;
}
}
break;
}
case JxlQuantMode.Dct4x8:
{
if (requiredSize != 1)
{
throw new InvalidOperationException("Invalid mode");
}
for (int c = 0; c < 3; c++)
{
if (!JxlF16Coder.Read(br, ref encoding.Dct4x8Multipliers![c]))
{
return false;
}
if (Math.Abs(encoding.Dct4x8Multipliers[c]) < AlmostZero)
{
throw new InvalidOperationException("DCT4X8 multiplier is too small");
}
}
if (!DecodeDctParameters(br, encoding.DctParameters!))
{
return false;
}
break;
}
case JxlQuantMode.Dct4:
{
if (requiredSize != 1)
{
throw new InvalidOperationException("Invalid mode");
}
for (int c = 0; c < 3; c++)
{
for (int i = 0; i < 2; i++)
{
if (!JxlF16Coder.Read(br, ref encoding.Dct4Multipliers![c][i]))
{
return false;
}
if (Math.Abs(encoding.Dct4Multipliers[c][i]) < AlmostZero)
{
throw new InvalidOperationException("DCT4 multiplier is too small");
}
}
}
if (!DecodeDctParameters(br, encoding.DctParameters!))
{
return false;
}
break;
}
case JxlQuantMode.Afv:
{
if (requiredSize != 1)
{
throw new InvalidOperationException("Invalid mode");
}
for (int c = 0; c < 3; c++)
{
for (int i = 0; i < 9; i++)
{
if (!JxlF16Coder.Read(br, ref encoding.AfvWeights![c][i]))
{
return false;
}
}
for (int i = 0; i < 6; i++)
{
encoding.AfvWeights![c][i] *= 64;
}
}
if (!DecodeDctParameters(br, encoding.DctParameters!))
{
return false;
}
if (!DecodeDctParameters(br, encoding.DctParametersAfv4x4!))
{
return false;
}
break;
}
case JxlQuantMode.Dct:
{
if (!DecodeDctParameters(br, encoding.DctParameters!))
{
return false;
}
break;
}
case JxlQuantMode.Raw:
{
// Set mode early, to avoid mem-leak.
encoding.Mode = JxlQuantMode.Raw;
if (!JxlModularFrameDecoder.DecodeQuantTable(configuration, requiredSizeX, requiredSizeY, br, encoding, idx, modularFrameDecoder))
{
return false;
}
break;
}
default:
throw new InvalidOperationException("Invalid quant table encoding");
}
encoding.Mode = (JxlQuantMode)mode;
return true;
}
}

8
src/ImageSharp/Formats/Jxl/Processing/Quantization/JxlQuantizerConstants.cs

@ -8,6 +8,14 @@ namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Quantization;
/// </summary>
internal static class JxlQuantizerConstants
{
public const int Log2MaxDistanceBands = 4;
public const int MaxDistanceBands = 1 + (1 << Log2MaxDistanceBands);
public const int CeilLog2NumPredefinedTables = 0;
public const int Log2NumQuantModes = 3;
/// <summary>
/// Total number of quantization tables.
/// </summary>

16
src/ImageSharp/Formats/Jxl/Processing/Splines/JxlSplineUtils.cs

@ -13,15 +13,17 @@ internal static class JxlSplineUtils
{
public const float DesiredRenderingDistance = 1f;
private const float PiBy32 = MathF.PI / 32;
private static ReadOnlySpan<float> ContinuousIDCTMultipliers =>
[
MathF.PI / 32 * 0, MathF.PI / 32 * 1, MathF.PI / 32 * 2, MathF.PI / 32 * 3, MathF.PI / 32 * 4,
MathF.PI / 32 * 5, MathF.PI / 32 * 6, MathF.PI / 32 * 7, MathF.PI / 32 * 8, MathF.PI / 32 * 9,
MathF.PI / 32 * 10, MathF.PI / 32 * 11, MathF.PI / 32 * 12, MathF.PI / 32 * 13, MathF.PI / 32 * 14,
MathF.PI / 32 * 15, MathF.PI / 32 * 16, MathF.PI / 32 * 17, MathF.PI / 32 * 18, MathF.PI / 32 * 19,
MathF.PI / 32 * 20, MathF.PI / 32 * 21, MathF.PI / 32 * 22, MathF.PI / 32 * 23, MathF.PI / 32 * 24,
MathF.PI / 32 * 25, MathF.PI / 32 * 26, MathF.PI / 32 * 27, MathF.PI / 32 * 28, MathF.PI / 32 * 29,
MathF.PI / 32 * 30, MathF.PI / 32 * 31,
PiBy32 * 0, PiBy32 * 1, PiBy32 * 2, PiBy32 * 3, PiBy32 * 4,
PiBy32 * 5, PiBy32 * 6, PiBy32 * 7, PiBy32 * 8, PiBy32 * 9,
PiBy32 * 10, PiBy32 * 11, PiBy32 * 12, PiBy32 * 13, PiBy32 * 14,
PiBy32 * 15, PiBy32 * 16, PiBy32 * 17, PiBy32 * 18, PiBy32 * 19,
PiBy32 * 20, PiBy32 * 21, PiBy32 * 22, PiBy32 * 23, PiBy32 * 24,
PiBy32 * 25, PiBy32 * 26, PiBy32 * 27, PiBy32 * 28, PiBy32 * 29,
PiBy32 * 30, PiBy32 * 31,
];
public static float ContinuousInverseDCT(in Dct32 dct, float t)

89
tests/ImageSharp.Tests/Formats/Jxl/Processing/BitsTests.cs

@ -0,0 +1,89 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using SixLabors.ImageSharp.Formats.Jxl.Processing;
namespace SixLabors.ImageSharp.Tests.Formats.Jxl.Processing;
public class BitsTests
{
[Fact]
public void TestNumZeroBits()
{
Assert.Equal(32u, JxlMath.Num0BitsAboveMS1Bit(0u));
Assert.Equal(64u, JxlMath.Num0BitsAboveMS1Bit(0uL));
Assert.Equal(32u, JxlMath.Num0BitsBelowLS1Bit(0u));
Assert.Equal(64u, JxlMath.Num0BitsBelowLS1Bit(0uL));
Assert.Equal(31u, JxlMath.Num0BitsAboveMS1Bit(1u));
Assert.Equal(30u, JxlMath.Num0BitsAboveMS1Bit(2u));
Assert.Equal(63u, JxlMath.Num0BitsAboveMS1Bit(1uL));
Assert.Equal(62u, JxlMath.Num0BitsAboveMS1Bit(2uL));
Assert.Equal(0u, JxlMath.Num0BitsBelowLS1Bit(1u));
Assert.Equal(0u, JxlMath.Num0BitsBelowLS1Bit(1uL));
Assert.Equal(1u, JxlMath.Num0BitsBelowLS1Bit(2u));
Assert.Equal(1u, JxlMath.Num0BitsBelowLS1Bit(2uL));
Assert.Equal(0u, JxlMath.Num0BitsAboveMS1Bit(0x80000000u));
Assert.Equal(0u, JxlMath.Num0BitsAboveMS1Bit(0x8000000000000000uL));
Assert.Equal(31u, JxlMath.Num0BitsBelowLS1Bit(0x80000000u));
Assert.Equal(63u, JxlMath.Num0BitsBelowLS1Bit(0x8000000000000000uL));
}
[Fact]
public void TestFloorLog2()
{
Span<int> expected = [0, 1, 1, 2, 2, 2, 2];
for (int i = 1; i <= 7; ++i)
{
Assert.Equal(expected[i - 1], JxlMath.FloorLog2Nonzero(i));
Assert.Equal((ulong)expected[i - 1], JxlMath.FloorLog2Nonzero((ulong)i));
}
Assert.Equal(11u, JxlMath.FloorLog2Nonzero(0x00000fffu)); // 4095
Assert.Equal(12u, JxlMath.FloorLog2Nonzero(0x00001000u)); // 4096
Assert.Equal(12u, JxlMath.FloorLog2Nonzero(0x00001001u)); // 4097
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0x80000000u));
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0x80000001u));
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0xFFFFFFFFu));
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0x80000000uL));
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0x80000001uL));
Assert.Equal(31u, JxlMath.FloorLog2Nonzero(0xFFFFFFFFuL));
Assert.Equal(63u, JxlMath.FloorLog2Nonzero(0x8000000000000000uL));
Assert.Equal(63u, JxlMath.FloorLog2Nonzero(0x8000000000000001uL));
Assert.Equal(63u, JxlMath.FloorLog2Nonzero(0xFFFFFFFFFFFFFFFFuL));
}
[Fact]
public void TestCeilLog2()
{
Span<int> expected = [0, 1, 2, 2, 3, 3, 3];
for (int i = 1; i <= 7; ++i)
{
Assert.Equal(expected[i - 1], JxlMath.CeilLog2Nonzero(i));
Assert.Equal((ulong)expected[i - 1], JxlMath.CeilLog2Nonzero((ulong)i));
}
Assert.Equal(12u, JxlMath.CeilLog2Nonzero(0x00000fffu)); // 4095
Assert.Equal(12u, JxlMath.CeilLog2Nonzero(0x00001000u)); // 4096
Assert.Equal(13u, JxlMath.CeilLog2Nonzero(0x00001001u)); // 4097
Assert.Equal(31u, JxlMath.CeilLog2Nonzero(0x80000000u));
Assert.Equal(32u, JxlMath.CeilLog2Nonzero(0x80000001u));
Assert.Equal(32u, JxlMath.CeilLog2Nonzero(0xFFFFFFFFu));
Assert.Equal(31u, JxlMath.CeilLog2Nonzero(0x80000000uL));
Assert.Equal(32u, JxlMath.CeilLog2Nonzero(0x80000001uL));
Assert.Equal(32u, JxlMath.CeilLog2Nonzero(0xFFFFFFFFuL));
Assert.Equal(63u, JxlMath.CeilLog2Nonzero(0x8000000000000000uL));
Assert.Equal(64u, JxlMath.CeilLog2Nonzero(0x8000000000000001uL));
Assert.Equal(64u, JxlMath.CeilLog2Nonzero(0xFFFFFFFFFFFFFFFFuL));
}
}

2
tests/ImageSharp.Tests/Formats/Jxl/README.md

@ -1,5 +1,5 @@
# JPEG XL Tests
### Preprocessor directives
Enable a ALLOW_JPEGXL_SLOW_TESTS preprocessor directive to allow
Enable a `ALLOW_JPEGXL_SLOW_TESTS` preprocessor directive to allow
slower tests that test the library more extensively.

19
tests/ImageSharp.Tests/Formats/Jxl/TestUtils/ColorEncodingDescriptor.cs

@ -0,0 +1,19 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using SixLabors.ImageSharp.Formats.Jxl.Cms;
namespace SixLabors.ImageSharp.Tests.Formats.Jxl.TestUtils;
internal struct ColorEncodingDescriptor
{
public JxlColorSpace ColorSpace { get; set; }
public JxlWhitePoint WhitePoint { get; set; }
public JxlPrimaries Primaries { get; set; }
public JxlTransferFunction TransferFunction { get; set; }
public JxlRenderingIntent RenderingIntent { get; set; }
}
Loading…
Cancel
Save