Browse Source

Add MA encoder

pull/3153/head
winscripter 4 weeks ago
parent
commit
5653d39664
  1. 6
      src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlToken.cs
  2. 2
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlContextPrediction.cs
  3. 8
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs
  4. 865
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs
  5. 2
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs
  6. 16
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlResidualToken.cs
  7. 884
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeSamples.cs
  8. 15
      src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ListUtils.cs
  9. 15
      src/ImageSharp/Formats/Jxl/Processing/Modular/README.md
  10. 87
      src/ImageSharp/Formats/Jxl/Processing/Primitives/Rng.cs

6
src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlToken.cs

@ -1,13 +1,13 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using SixLabors.ImageSharp.Formats.Jxl.Processing.Splines;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Encoder.Ans;
internal struct JxlToken(JxlSplineEntropyContext c, uint value)
internal struct JxlToken(JxlMaTreeContext c, uint value)
{
public bool IsLz77Length;
public JxlSplineEntropyContext Context = c;
public JxlMaTreeContext Context = c;
public uint Value = value;
}

2
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlContextPrediction.cs

@ -12,7 +12,7 @@ internal static class JxlContextPrediction
{
private const int ExtraPropertiesPerChannel = 4;
private const int NumberOfProperties = 1;
public const int NumberOfProperties = 1;
public static void SetPredictorMode(int i, JxlModularHeader header)
{

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

@ -1,6 +1,8 @@
// 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 static class JxlMaConstants
@ -13,4 +15,10 @@ internal static class JxlMaConstants
public const int MaxTreeSize = 1 << 22;
public const int PropertyRangeFast = 512 << 4;
public const int NumNonrefProperties = 2 + 13 + JxlContextPrediction.NumberOfProperties;
public const int WpProp = NumNonrefProperties - JxlContextPrediction.NumberOfProperties;
public const int GradientProp = 9;
}

865
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs

@ -0,0 +1,865 @@
// 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.InteropServices;
using SixLabors.ImageSharp.Formats.Jxl.IO.Entropy;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Encoder.Ans;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Primitives;
using StaticPropertyRange = System.Runtime.CompilerServices.InlineArray2<System.Runtime.CompilerServices.InlineArray2<int>>;
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 JxlMaEncoder
{
internal enum IntersectionType
{
None,
Partial,
Inside
}
private static int Padded(int x) => JxlMath.RoundUpTo(x, Vector<float>.Count);
public static float EstimateBits(Span<int> counts, int numSymbols)
{
int total = TensorPrimitives.Sum((ReadOnlySpan<int>)counts.Slice(0, numSymbols));
Vector<float> minprob = Vector.Create(1.0f / JxlAnsConstants.AnsTableSize);
Vector<float> inverseTotal = Vector.Create(1.0f / total);
Vector<float> bitsLanes = Vector<float>.Zero;
for (int i = 0; i < numSymbols; i += Vector<float>.Count)
{
Vector<int> countsIv = new(counts[i..]);
Vector<float> countsFv = Vector.ConvertToSingle(countsIv);
Vector<float> probs = countsFv * inverseTotal;
Vector<float> mprobs = probs * minprob;
Vector<float> nbps = Vector.ConditionalSelect(
Vector.Equals(countsIv, Vector.Create(total)),
Vector<float>.Zero,
Vector.Log2(mprobs));
bitsLanes -= countsFv * nbps;
}
return Vector.Sum(bitsLanes);
}
public static void MakeSplitNode(int pos, int property, int splitValue, JxlPredictor leftPredictor, int leftOffset, JxlPredictor rightPredictor, int rightOffset, Tree tree)
{
ref JxlPropertyDecisionNode treePos = ref CollectionsMarshal.AsSpan(tree)[pos];
treePos.LeftChild = tree.Count;
treePos.RightChild = tree.Count + 1;
treePos.SplitValue = splitValue;
treePos.Property = property;
JxlPropertyDecisionNode newRightNode = new()
{
Property = -1,
Predictor = rightPredictor,
PredictorOffset = rightOffset,
Multiplier = 1
};
JxlPropertyDecisionNode newLeftNode = new()
{
Property = -1,
Predictor = leftPredictor,
PredictorOffset = leftOffset,
Multiplier = 1
};
tree.Add(newRightNode);
tree.Add(newLeftNode);
}
public static IntersectionType BoxIntersects(StaticPropertyRange needle, StaticPropertyRange haystack, ref int partialAxis, ref int partialValue)
{
bool partial = false;
for (int i = 0; i < JxlPredictorFacts.StaticProperties; i++)
{
if (haystack[i][0] >= needle[i][1])
{
return IntersectionType.None;
}
if (haystack[i][1] <= needle[i][0])
{
return IntersectionType.None;
}
if (haystack[i][0] <= needle[i][0] && haystack[i][1] >= needle[i][1])
{
continue;
}
partial = true;
partialAxis = i;
int innerOffset = haystack[i][0] > needle[i][0] && haystack[i][0] < needle[i][1] ? 0 : 1; // yes, if true then 0 not 1
partialValue = haystack[i][innerOffset] - 1;
}
return partial ? IntersectionType.Partial : IntersectionType.Inside;
}
public static void SplitTreeSamples(bool s, JxlTreeSamples samples, int begin, int pos, int end, int prop, int value)
{
int beginPos = begin;
int endPos = pos;
do
{
while (beginPos < pos && samples.GetProperty(s, prop, beginPos) <= value)
{
beginPos++;
}
while (endPos < end && samples.GetProperty(s, prop, endPos) > value)
{
endPos++;
}
if (beginPos < pos && endPos < end)
{
samples.Swap(beginPos, endPos);
}
beginPos++;
endPos++;
}
while (beginPos < pos && endPos < end);
}
// Simple overload so we don't have to type out CollectionsMarshal.AsSpan all the time.
public static void CollectExtraBitsIncrease(bool s, JxlTreeSamples treeSamples, List<JxlResidualToken> residualTokens, List<int> countIncrease, List<int> extraBitsIncrease, int begin, int end, int propertyIndex, int maxSymbols)
=> CollectExtraBitsIncrease(
s,
treeSamples,
residualTokens,
CollectionsMarshal.AsSpan(countIncrease),
CollectionsMarshal.AsSpan(extraBitsIncrease),
begin,
end,
propertyIndex,
maxSymbols);
public static void CollectExtraBitsIncrease(bool s, JxlTreeSamples treeSamples, List<JxlResidualToken> residualTokens, Span<int> countIncrease, Span<int> extraBitsIncrease, int begin, int end, int propertyIndex, int maxSymbols)
{
for (int i2 = begin; i2 < end; i2++)
{
JxlResidualToken rt = residualTokens[i2];
int cnt = treeSamples.GetCount(i2);
int p = treeSamples.GetProperty(s, propertyIndex, i2);
int sym = rt.Token;
int ebi = rt.NumberOfBits * cnt;
countIncrease[(p * maxSymbols) + sym] += cnt;
extraBitsIncrease[p] += ebi;
}
}
public static unsafe void FindBestSplit(JxlTreeSamples treeSamples, float threshold, List<JxlModularMultiplierInfo> mulInfo, StaticPropertyRange initialStaticPropertyRange, float fastDecodeMultiplier, Tree tree)
{
Stack<NodeInfo> nodes = [];
nodes.Push(new(0, 0, treeSamples.NumberOfDistinctSamples, initialStaticPropertyRange));
int numPred = treeSamples.NumberOfPredictors;
int numProp = treeSamples.NumberOfProperties;
Span<int> totalExtraBits = stackalloc int[numPred];
while (nodes.Count > 0)
{
NodeInfo last = nodes.Peek();
int pos = last.Pos;
int begin = last.Begin;
int end = last.End;
StaticPropertyRange staticPropertyRange = last.StaticPropertyRange;
_ = nodes.Pop();
if (begin == end)
{
continue;
}
SplitInfo bestSplitStaticConstant = default;
SplitInfo bestSplitStatic = default;
SplitInfo bestSplitNonStatic = default;
SplitInfo bestSplitNoWeightedPrediction = default;
if (begin > end)
{
throw new InvalidOperationException("Begin should be <= end");
}
if (end > treeSamples.NumberOfDistinctSamples)
{
throw new InvalidOperationException("End index out of range");
}
int maxSymbols = 0;
for (int pred = 0; pred < numPred; pred++)
{
for (int i = begin; i < end; i++)
{
int token = treeSamples.GetToken(pred, i);
maxSymbols = maxSymbols > token + 1 ? maxSymbols : token + 1;
}
}
maxSymbols = Padded(maxSymbols);
int[] counts = ArrayPool<int>.Shared.Rent(maxSymbols * numPred);
for (int pred = 0; pred < numPred; pred++)
{
int extraBits = 0;
List<JxlResidualToken> rtokens = treeSamples.GetResidualTokensForPrediction(pred);
for (int i = begin; i < end; i++)
{
JxlResidualToken rt = rtokens[i];
int count = treeSamples.GetCount(i);
int eb = rt.NumberOfBits * count;
counts[(pred * maxSymbols) + rt.Token] += count;
extraBits += eb;
}
totalExtraBits[pred] = extraBits;
}
float baseBits = 0;
{
int pred = treeSamples.FindPredictorIndex(tree[pos].Predictor);
baseBits = EstimateBits(counts.AsSpan(pred * maxSymbols), maxSymbols) + totalExtraBits[pred];
}
ref SplitInfo best = ref bestSplitNonStatic;
SplitInfo forcedSplit = default;
foreach (JxlModularMultiplierInfo mmi in mulInfo)
{
int axis = 0;
int val = 0;
IntersectionType t = BoxIntersects(staticPropertyRange, mmi.Range, ref axis, ref val);
if (t == IntersectionType.None)
{
continue;
}
if (t == IntersectionType.Inside)
{
CollectionsMarshal.AsSpan(tree)[pos].Multiplier = (int)mmi.Multiplier;
break;
}
if (t == IntersectionType.Partial)
{
forcedSplit.Value = treeSamples.QuantizeStaticProperty(axis, val);
forcedSplit.Property = axis;
forcedSplit.LeftCost = forcedSplit.RightCost = (baseBits / 2) - threshold;
forcedSplit.LeftPredictor = forcedSplit.RightPredictor = CollectionsMarshal.AsSpan(tree)[pos].Predictor;
best = ref forcedSplit;
best.Position = begin;
if (best.Property != treeSamples.PropertyFromIndex(best.Property))
{
throw new InvalidOperationException("Invalid property");
}
if (best.Property < treeSamples.NumberOfStaticProperties)
{
for (int x = begin; x < end; x++)
{
if (treeSamples.GetProperty(true, best.Property, x) <= best.Value)
{
best.Position++;
}
}
}
else
{
int prop = best.Property - treeSamples.NumberOfStaticProperties;
for (int x = begin; x < end; x++)
{
if (treeSamples.GetProperty(false, prop, x) <= best.Value)
{
best.Position++;
}
}
}
break;
}
}
if (!Unsafe.AreSame(ref best, ref forcedSplit))
{
List<int> countsIncrease = [];
List<int> extraBitsIncrease = [];
List<CostInfo> leftCosts = [];
List<CostInfo> rightCosts = [];
int[] aboveCounts = ArrayPool<int>.Shared.Rent(maxSymbols);
int[] belowCounts = ArrayPool<int>.Shared.Rent(maxSymbols);
float changePredictionPenalty = 800.0f / (100.0f + threshold);
for (int prop = 0; prop < numProp && baseBits > threshold; prop++)
{
leftCosts.Clear();
rightCosts.Clear();
int propertySize = treeSamples.CountPropertyValues(prop);
if (extraBitsIncrease.Count < propertySize)
{
countsIncrease.Grow(0, propertySize * maxSymbols);
extraBitsIncrease.Grow(0, propertySize);
}
int[] propertyValueUsedCount = ArrayPool<int>.Shared.Rent(propertySize);
propertyValueUsedCount.AsSpan().Clear();
int firstUsed = propertySize;
int lastUsed = 0;
if (prop < treeSamples.NumberOfStaticProperties)
{
for (int i = begin; i < end; i++)
{
int p = treeSamples.GetProperty(true, prop, i);
propertyValueUsedCount[p]++;
lastUsed = Math.Max(lastUsed, p);
firstUsed = Math.Max(firstUsed, p);
}
}
else
{
int prop_idx = prop - treeSamples.NumberOfStaticProperties;
for (int i = begin; i < end; i++)
{
int p = treeSamples.GetProperty(false, prop_idx, i);
propertyValueUsedCount[p]++;
lastUsed = Math.Max(lastUsed, p);
firstUsed = Math.Max(firstUsed, p);
}
}
leftCosts.Grow(default, lastUsed - firstUsed);
rightCosts.Grow(default, lastUsed - firstUsed);
for (int pred = 0; pred < numPred; pred++)
{
List<JxlResidualToken> rtokens = treeSamples.GetResidualTokensForPrediction(pred);
if (prop < treeSamples.NumberOfStaticProperties)
{
CollectExtraBitsIncrease(true, treeSamples, rtokens, countsIncrease, extraBitsIncrease, begin, end, prop, maxSymbols);
}
else
{
CollectExtraBitsIncrease(false, treeSamples, rtokens, countsIncrease, extraBitsIncrease, begin, end, prop - treeSamples.NumberOfStaticProperties, maxSymbols);
}
counts.AsSpan().Slice(pred * maxSymbols, maxSymbols).CopyTo(aboveCounts);
belowCounts.AsSpan().Slice(0, maxSymbols).Clear();
int extraBitsBelow = 0;
for (int i = firstUsed; i < lastUsed; i++)
{
if (propertyValueUsedCount[i] == 0)
{
continue;
}
extraBitsBelow += extraBitsIncrease[i];
extraBitsIncrease[i] = 0;
for (int sym = 0; sym < maxSymbols; sym++)
{
aboveCounts[sym] -= countsIncrease[(i * maxSymbols) + sym];
belowCounts[sym] += countsIncrease[(i * maxSymbols) + sym];
countsIncrease[(i * maxSymbols) + sym] = 0;
}
float rightCost = EstimateBits(aboveCounts.AsSpan(), maxSymbols) + totalExtraBits[pred] - extraBitsBelow;
float leftCost = EstimateBits(belowCounts.AsSpan(), maxSymbols) + extraBitsBelow;
if (extraBitsBelow > totalExtraBits[pred])
{
throw new InvalidOperationException("Too many extra bits");
}
float penalty = 0;
if (treeSamples.PredictorFromIndex(pred) != tree[pos].Predictor &&
tree[pos].Predictor != JxlPredictor.Weighted)
{
penalty = changePredictionPenalty;
}
if (treeSamples.PredictorFromIndex(pred) == JxlPredictor.Weighted)
{
penalty += 1e-8f;
}
if (treeSamples.PredictorFromIndex(pred) == JxlPredictor.Zero)
{
penalty -= 1e-8f;
}
if (rightCost + penalty < rightCosts[i - firstUsed].TotalCost)
{
CostInfo cost = rightCosts[i - firstUsed];
cost.Cost = rightCost;
cost.ExtraCost = penalty;
cost.Predictor = treeSamples.PredictorFromIndex(pred);
rightCosts[i - firstUsed] = cost;
}
if (leftCost + penalty < leftCosts[i - firstUsed].TotalCost)
{
CostInfo cost = leftCosts[i - firstUsed];
cost.Cost = leftCost;
cost.ExtraCost = penalty;
cost.Predictor = treeSamples.PredictorFromIndex(pred);
leftCosts[i - firstUsed] = cost;
}
}
}
int split = begin;
for (int i = firstUsed; i < lastUsed; i++)
{
if (propertyValueUsedCount[i] == 0)
{
continue;
}
split += propertyValueUsedCount[i];
float rightCost = rightCosts[i - firstUsed].Cost;
float leftCost = leftCosts[i - firstUsed].Cost;
bool usesWeightedPrediction = treeSamples.PropertyFromIndex(prop) == JxlMaConstants.WpProp ||
leftCosts[i - firstUsed].Predictor == JxlPredictor.Weighted ||
rightCosts[i - firstUsed].Predictor == JxlPredictor.Weighted;
bool zeroEntropySide = rightCost == 0 || leftCost == 0;
// Using pointer specifically for this variable instead of
// ref, as we can't assign a ref variable conditionally.
SplitInfo* referenceToBest =
treeSamples.PropertyFromIndex(prop) < JxlPredictorFacts.StaticProperties
? (zeroEntropySide ? &bestSplitStaticConstant : &bestSplitStatic)
: (usesWeightedPrediction ? &bestSplitNonStatic : &bestSplitNoWeightedPrediction);
if (leftCost + rightCost < referenceToBest->Cost)
{
referenceToBest->Property = prop;
referenceToBest->Value = i;
referenceToBest->Position = split;
referenceToBest->LeftCost = leftCost;
referenceToBest->LeftPredictor = leftCosts[i - firstUsed].Predictor;
referenceToBest->RightCost = rightCost;
referenceToBest->RightPredictor = rightCosts[i - firstUsed].Predictor;
}
}
extraBitsIncrease[lastUsed] = 0;
for (int sym = 0; sym < maxSymbols; sym++)
{
countsIncrease[(lastUsed * maxSymbols) + sym] = 0;
}
}
if (bestSplitNoWeightedPrediction.Cost + threshold < baseBits &&
bestSplitNoWeightedPrediction.Cost <= fastDecodeMultiplier * best.Cost)
{
best = ref bestSplitNoWeightedPrediction;
}
if (bestSplitStatic.Cost + threshold < baseBits &&
bestSplitStatic.Cost <= fastDecodeMultiplier * best.Cost)
{
best = ref bestSplitStatic;
}
if (bestSplitStaticConstant.Cost + threshold < baseBits)
{
best = ref bestSplitStaticConstant;
}
ArrayPool<int>.Shared.Return(aboveCounts);
ArrayPool<int>.Shared.Return(belowCounts);
}
if (best.Cost + threshold < baseBits)
{
int p = treeSamples.PropertyFromIndex(best.Property);
int dequant = treeSamples.UnquantizeProperty(best.Property, best.Value);
MakeSplitNode(pos, p, dequant, best.LeftPredictor, 0, best.RightPredictor, 0, tree);
if (best.Property < treeSamples.NumberOfStaticProperties)
{
SplitTreeSamples(true, treeSamples, begin, best.Position, end, best.Property, best.Value);
}
else
{
SplitTreeSamples(false, treeSamples, begin, best.Position, end, best.Property - treeSamples.NumberOfStaticProperties, best.Value);
}
StaticPropertyRange newStaticPropertyRange = staticPropertyRange;
if (p < JxlPredictorFacts.StaticProperties)
{
if (dequant + 1 > newStaticPropertyRange[p][1])
{
throw new InvalidOperationException("Dequantized coefficient is out of range");
}
newStaticPropertyRange[p][1] = dequant + 1;
if (newStaticPropertyRange[p][0] >= newStaticPropertyRange[p][1])
{
throw new InvalidOperationException("Static property is out of range");
}
}
nodes.Push(new(tree[pos].RightChild, begin, best.Position, newStaticPropertyRange));
newStaticPropertyRange = staticPropertyRange;
if (p < JxlPredictorFacts.StaticProperties)
{
if (newStaticPropertyRange[p][0] > dequant + 1)
{
throw new InvalidOperationException("Static property must be <= dequantized coefficient");
}
newStaticPropertyRange[p][0] = dequant + 1;
if (newStaticPropertyRange[p][0] >= newStaticPropertyRange[p][1])
{
throw new InvalidOperationException("Static property is out of range");
}
}
nodes.Push(new NodeInfo(tree[pos].LeftChild, best.Position, end, newStaticPropertyRange));
}
ArrayPool<int>.Shared.Return(counts);
}
}
public static void ComputeBestTree(JxlTreeSamples treeSamples, float threshold, List<JxlModularMultiplierInfo> mulInfo, StaticPropertyRange staticPropertyRange, float fastDecodeMultiplier, Tree tree)
{
JxlPropertyDecisionNode node = new()
{
Property = -1,
Predictor = treeSamples.PredictorFromIndex(0),
PredictorOffset = 0,
Multiplier = 1
};
tree.Add(node);
if (treeSamples.NumberOfProperties >= 64)
{
throw new InvalidOperationException($"Too many properties: {treeSamples.NumberOfPredictors}");
}
FindBestSplit(treeSamples, threshold, mulInfo, staticPropertyRange, fastDecodeMultiplier, tree);
}
public static List<int> QuantizeHistogram(ReadOnlySpan<int> histogram, int numChunks)
{
if (histogram.Length == 0 || numChunks == 0)
{
return [];
}
int sum = TensorPrimitives.Sum(histogram);
if (sum == 0)
{
return [];
}
List<int> thresholds = [];
long cumulativeSum = 0;
long threshold = 1;
for (int i = 0; i < histogram.Length; i++)
{
cumulativeSum += histogram[i];
if (cumulativeSum * numChunks >= threshold * sum)
{
thresholds.Add(i);
while (cumulativeSum * numChunks >= threshold * sum)
{
threshold++;
}
}
}
if (thresholds.Count > numChunks)
{
throw new InvalidOperationException("Too many thresholds");
}
thresholds.RemoveAt(thresholds.Count - 1);
return thresholds;
}
public static List<int> QuantizeSamples(ReadOnlySpan<int> samples, int numChunks)
{
const int range = 512;
if (samples.Length == 0)
{
return [];
}
int min = Math.Clamp(TensorPrimitives.Min(samples), -range, range);
Span<int> counts = stackalloc int[2048].Slice(0, (2 * range) + 1);
for (int i = 0; i < samples.Length; i++)
{
int s = samples[i];
int sampleOffset = Math.Clamp(s, -range, range) - min;
counts[sampleOffset]++;
}
List<int> thresholds = QuantizeHistogram(counts, numChunks);
// For fast processing via TensorPrimitives
Span<int> thresholdsSpan = CollectionsMarshal.AsSpan(thresholds);
TensorPrimitives.Add(thresholdsSpan, min, thresholdsSpan);
return thresholds;
}
public static void QuantMap(Span<int> from, Span<int> to, int numPegs, int bias)
{
int mapped = 0;
for (int i = 0; i < numPegs; i++)
{
while (mapped < from.Length && i - bias > from[mapped])
{
mapped++;
}
to[i] = mapped;
}
}
public static void CollectPixelSamples(Configuration configuration, JxlModularImage image, JxlModularOptions options, int groupId, Span<int> groupPixelCount, Span<int> channelPixelCount, List<int> pixelSamples, List<int> diffSamples)
{
if (options.NumberOfRepeats == 0)
{
throw new InvalidOperationException("No repeats");
}
// Power of 2 size is unknown.
Span<int> alignedGroupPixelCount = stackalloc int[groupPixelCount.Length <= groupId ? groupId + 1 : groupPixelCount.Length];
groupPixelCount.CopyTo(alignedGroupPixelCount);
Span<int> alignedChannelPixelCount = stackalloc int[channelPixelCount.Length < image.Channels.Count ? image.Channels.Count : channelPixelCount.Length];
channelPixelCount.CopyTo(alignedGroupPixelCount);
Rng rng = new((ulong)groupId);
float fraction = MathF.Min(options.NumberOfRepeats * 0.1f, 0.99f);
Rng.GeometricDistribution dist = new(fraction);
int totalPixels = 0;
List<int> channelIds = [];
int i;
for (i = 0; i < image.Channels.Count; i++)
{
JxlModularChannel channel = image.Channels[i];
if (i >= image.MetaChannels && (channel.Width > options.MaxChannelSize || channel.Height > options.MaxChannelSize))
{
break;
}
if (channel.Width <= 1 || channel.Height == 0)
{
continue;
}
channelIds.Add(i);
groupPixelCount[groupId] += channel.Width * channel.Height;
channelPixelCount[i] += channel.Width * channel.Height;
totalPixels += channel.Width * channel.Height;
}
if (channelIds.Count == 0)
{
throw new InvalidOperationException("No channel IDs");
}
pixelSamples.Grow(0, (int)(pixelSamples.Count + (fraction * totalPixels)));
diffSamples.Grow(0, (int)(diffSamples.Count + (fraction * totalPixels)));
i = 0;
int y = 0;
int x = 0;
void Advance(uint amount)
{
x += (int)amount;
while (x >= image.Channels[channelIds[i]].Width)
{
x -= image.Channels[channelIds[i]].Width;
y++;
if (y == image.Channels[channelIds[i]].Height)
{
i++;
y = 0;
if (i >= channelIds.Count)
{
return;
}
}
}
}
Advance(rng.Geometric(in dist));
for (; i < channelIds.Count; Advance(rng.Geometric(in dist) + 1))
{
Span<int> row = image.Channels[channelIds[i]].GetRow(y);
pixelSamples.Add(row[x]);
int xp = x == 0 ? 1 : x - 1;
diffSamples.Add(row[x] - row[xp]);
}
}
public static void TokenizeTree(Tree tree, List<JxlToken> tokens, Tree decoderTree)
{
if (tree.Count > JxlMaConstants.MaxTreeSize)
{
throw new InvalidOperationException("Too many tree nodes");
}
Queue<int> q = [];
q.Enqueue(0);
int leafId = 0;
decoderTree.Clear();
while (q.Count > 0)
{
int cur = q.Peek();
_ = q.Dequeue();
if (tree[cur].Property < -1)
{
throw new InvalidOperationException("Property value is too small");
}
tokens.Add(new JxlToken(JxlMaTreeContext.Property, (uint)tree[cur].Property + 1));
if (tree[cur].Property == -1)
{
tokens.Add(new JxlToken(JxlMaTreeContext.Predictor, (uint)tree[cur].Predictor));
tokens.Add(new JxlToken(JxlMaTreeContext.Offset, JxlPackSigned.PackUnsigned((int)tree[cur].PredictorOffset)));
uint mulLog = JxlMath.Num0BitsBelowLS1Bit_Nonzero(tree[cur].Multiplier);
int mulBits = (tree[cur].Multiplier >> (int)mulLog) - 1;
tokens.Add(new JxlToken(JxlMaTreeContext.MultiplierLog, mulLog));
tokens.Add(new JxlToken(JxlMaTreeContext.MultiplierBits, (uint)mulBits));
if (tree[cur].Predictor >= JxlPredictor.Best)
{
throw new InvalidOperationException("Invalid predictor");
}
decoderTree.Add(new JxlPropertyDecisionNode(-1, 0, leafId, 0, tree[cur].Predictor, tree[cur].PredictorOffset, tree[cur].Multiplier));
leafId++;
continue;
}
decoderTree.Add(new JxlPropertyDecisionNode(tree[cur].Property, tree[cur].SplitValue, decoderTree.Count + q.Count + 1, decoderTree.Count + q.Count + 2, JxlPredictor.Zero, 0, 1));
q.Enqueue(tree[cur].LeftChild);
q.Enqueue(tree[cur].RightChild);
tokens.Add(new JxlToken(JxlMaTreeContext.SplitValue, JxlPackSigned.PackUnsigned(tree[cur].SplitValue)));
}
}
private readonly record struct NodeInfo(int Pos, int Begin, int End, StaticPropertyRange StaticPropertyRange);
private struct SplitInfo()
{
public int Property { get; set; }
public int Value { get; set; }
public int Position { get; set; }
public float LeftCost { get; set; } = float.MaxValue;
public float RightCost { get; set; } = float.MaxValue;
public JxlPredictor LeftPredictor { get; set; } = JxlPredictor.Zero;
public JxlPredictor RightPredictor { get; set; } = JxlPredictor.Zero;
public readonly float Cost => this.LeftCost + this.RightCost;
}
private struct CostInfo()
{
public float Cost { get; set; } = float.MaxValue;
public float ExtraCost { get; set; }
public JxlPredictor Predictor { get; set; }
public readonly float TotalCost => this.Cost + this.ExtraCost;
}
}

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

@ -3,7 +3,7 @@
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal enum JxlMaTreeContext
internal enum JxlMaTreeContext : byte
{
SplitValue,

16
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlResidualToken.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 struct JxlResidualToken
{
public int Token;
public int NumberOfBits;
public JxlResidualToken(int token, int numberOfBits)
{
this.Token = token;
this.NumberOfBits = numberOfBits;
}
}

884
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeSamples.cs

@ -0,0 +1,884 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using System.Numerics.Tensors;
using System.Runtime.CompilerServices;
using System.Runtime.InteropServices;
using SixLabors.ImageSharp.Common.Helpers;
using SixLabors.ImageSharp.Formats.Jxl.IO.Entropy;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding.ContextPrediction;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Primitives;
using StaticPropertyRange = System.Runtime.CompilerServices.InlineArray2<System.Runtime.CompilerServices.InlineArray2<int>>;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
/// <summary>
/// Contains all the necessary data needed to build a tree.
/// </summary>
internal sealed class JxlTreeSamples
{
/// <summary>
/// This deduplication table entry marks an unused entry.
/// </summary>
private const int DeduplicationEntryUnused = -1;
private const int PropertyRange = 511;
/// <summary>
/// Residual information: token and number of extra bits per predictor.
/// </summary>
private List<List<JxlResidualToken>> residuals = [];
/// <summary>
/// Number of occurrences of each sample.
/// </summary>
private readonly List<int> sampleCounts = [];
/// <summary>
/// Quantized static property values.
/// </summary>
private InlineArray2<List<int>> staticProperties;
/// <summary>
/// Property values, quantized to at most 256 distinct values.
/// </summary>
private readonly List<List<byte>> properties = [];
/// <summary>
/// Decompactification info for <see cref="properties"/>.
/// </summary>
private readonly List<List<int>> compactProperties = [];
/// <summary>
/// List of properties to use.
/// </summary>
private List<int> propertiesToUse = [];
/// <summary>
/// List of predictors to use.
/// </summary>
private List<JxlPredictor> predictors = [];
/// <summary>
/// Mapping property value -> quantized property value.
/// </summary>
private InlineArray2<List<int>> staticPropertyMapping;
private readonly List<List<int>> propertyMapping = [];
/// <summary>
/// Table for deduplication.
/// </summary>
private List<int> deduplicationTable = [];
/// <summary>
/// Gets a value indicating whether there are any residual samples.
/// </summary>
public bool HasSamples => this.residuals.Count > 0 && this.residuals[0].Count > 0;
/// <summary>
/// Gets the total number of distinct samples.
/// </summary>
public int NumberOfDistinctSamples => this.sampleCounts.Count;
/// <summary>
/// Gets the total number of samples.
/// </summary>
public int NumberOfSamples { get; private set; }
/// <summary>
/// Gets the number of quantized static property values.
/// </summary>
public int NumberOfStaticProperties { get; private set; }
/// <summary>
/// Gets the total number of predictors.
/// </summary>
public int NumberOfPredictors => this.predictors.Count;
/// <summary>
/// Gets the total number of properties.
/// </summary>
public int NumberOfProperties => this.propertiesToUse.Count;
/// <summary>
/// Returns a List of residual tokens for the specified kind of prediction.
/// </summary>
/// <param name="prediction">The type of prediction.</param>
/// <returns>Residual tokens corresponding to the prediction specified by the parameter <paramref name="prediction"/>.</returns>
public List<JxlResidualToken> GetResidualTokensForPrediction(int prediction) => this.residuals[prediction];
/// <summary>
/// Returns a reference to the residual token.
/// </summary>
/// <param name="pred">Kind of prediction.</param>
/// <param name="i">Index of the residual token within that prediction.</param>
/// <returns>The residual token for prediction <paramref name="pred"/> indexed <paramref name="i"/>.</returns>
public ref JxlResidualToken GetResidualToken(int pred, int i) => ref CollectionsMarshal.AsSpan(this.residuals[pred])[i];
/// <summary>
/// Returns a token of the residual for prediction <paramref name="prediction"/> index <paramref name="index"/>.
/// </summary>
/// <param name="prediction">The kind of prediction.</param>
/// <param name="index">Index of the residual token within that prediction.</param>
/// <returns>For residual token whose prediction is <paramref name="prediction"/> and index is <paramref name="index"/>, returns its token coefficient.</returns>
public int GetToken(int prediction, int index) => this.residuals[prediction][index].Token;
/// <summary>
/// Returns the number of occurrences for sample <paramref name="i"/>.
/// </summary>
/// <param name="i">The index of the sample.</param>
/// <returns>Number of times <paramref name="i"/> appears.</returns>
public int GetCount(int i) => this.sampleCounts[i];
/// <summary>
/// Finds the index of the predictor <paramref name="predictor"/>.
/// </summary>
/// <param name="predictor">The predictor to find the index for.</param>
/// <returns>Index of the predictor in the predictors storage.</returns>
/// <exception cref="InvalidOperationException">Thrown when the predictor can't be found.</exception>
public int FindPredictorIndex(JxlPredictor predictor)
{
ReadOnlySpan<JxlPredictor> span = CollectionsMarshal.AsSpan(this.predictors);
int index = span.IndexOf(predictor);
if (index < 0)
{
// Should not happen.
throw new InvalidOperationException("Cannot find the index of the predictor");
}
return index;
}
/// <summary>
/// Finds the index of the property <paramref name="property"/>.
/// </summary>
/// <param name="property">The property to find the index for.</param>
/// <returns>Index of the property in the properties storage.</returns>
/// <exception cref="InvalidOperationException">Thrown when the property isn't valid.</exception>
public int FindPropertyIndex(int property)
{
ReadOnlySpan<int> span = CollectionsMarshal.AsSpan(this.propertiesToUse);
int index = span.IndexOf(property);
if (index != this.propertiesToUse[^1])
{
// Should not happen.
throw new InvalidOperationException("Invalid property");
}
return index;
}
/// <summary>
/// Returns the number of properties for a <paramref name="propertyIndex"/>.
/// </summary>
/// <param name="propertyIndex">Index of the property.</param>
/// <returns>Number of property values for property with index <paramref name="propertyIndex"/>.</returns>
public int CountPropertyValues(int propertyIndex) => this.compactProperties[propertyIndex].Count + 1;
/// <summary>
/// Returns the value of a property.
/// </summary>
/// <param name="useStaticProperty">Prefer a static property?</param>
/// <param name="propertyIndex">The index of the properties table.</param>
/// <param name="i">The index of the property.</param>
/// <returns>
/// Property for index <paramref name="i" /> within table <paramref name="propertyIndex"/>.
/// </returns>
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public int GetProperty(bool useStaticProperty, int propertyIndex, int i)
{
if (useStaticProperty)
{
return this.staticProperties[propertyIndex][i];
}
else
{
return this.properties[propertyIndex][i];
}
}
/// <summary>
/// Returns the dequantized property.
/// </summary>
/// <param name="propertyIndex">Index of the target property.</param>
/// <param name="quant">Property quantizer.</param>
/// <returns>The dequantized property.</returns>
/// <exception cref="InvalidOperationException">Thrown when the quant is out of range.</exception>
public int UnquantizeProperty(int propertyIndex, int quant)
{
List<int> compactProperties = this.compactProperties[propertyIndex];
if (quant >= compactProperties.Count)
{
throw new InvalidOperationException("Quant is out of range");
}
return compactProperties[quant];
}
/// <summary>
/// Returns a predictor for index <paramref name="index"/>.
/// </summary>
/// <param name="index">The index of the predictor.</param>
/// <returns>Predictor at index <paramref name="index"/>.</returns>
public JxlPredictor PredictorFromIndex(int index)
{
DebugGuard.MustBeLessThan(index, this.predictors.Count, nameof(index));
return this.predictors[index];
}
/// <summary>
/// Returns a property for index <paramref name="index"/>.
/// </summary>
/// <param name="index">The index of the property.</param>
/// <returns>Property at index <paramref name="index"/>.</returns>
public int PropertyFromIndex(int index)
{
DebugGuard.MustBeLessThan(index, this.propertiesToUse.Count, nameof(index));
return this.propertiesToUse[index];
}
/// <summary>
/// Invoked after processing samples completed.
/// </summary>
public void AllSamplesDone() => this.deduplicationTable = [];
/// <summary>
/// Returns the quantized property value of <paramref name="v"/>.
/// </summary>
/// <param name="property">Index of the property.</param>
/// <param name="v">Value to quantize.</param>
/// <returns>Quantized value.</returns>
public int QuantizeProperty(int property, int v)
{
DebugGuard.MustBeGreaterThanOrEqualTo(property, this.NumberOfStaticProperties, nameof(property));
v = Math.Clamp(v, -PropertyRange, PropertyRange) + PropertyRange;
return this.propertyMapping[property][v];
}
/// <summary>
/// Returns the quantized static property value of <paramref name="v"/>.
/// </summary>
/// <param name="property">Index of the static property.</param>
/// <param name="v">Value to quantize.</param>
/// <returns>Quantized value.</returns>
public int QuantizeStaticProperty(int property, int v)
{
DebugGuard.MustBeGreaterThanOrEqualTo(property, this.NumberOfStaticProperties, nameof(property));
v = Math.Clamp(v, -PropertyRange, PropertyRange) + PropertyRange;
return this.staticPropertyMapping[property][v];
}
public void SetPredictor(JxlPredictor predictor, JxlTreeMode wpTreeMode)
{
if (wpTreeMode == JxlTreeMode.WpOnly)
{
this.predictors = [JxlPredictor.Weighted];
// equivalent of residuals.resize(1)
if (this.residuals.Count > 1)
{
this.residuals = [this.residuals[0]];
}
else if (this.residuals.Count < 1)
{
this.residuals = [[]];
}
return;
}
if (wpTreeMode == JxlTreeMode.NoWp && predictor == JxlPredictor.Weighted)
{
throw new InvalidOperationException("Invalid predictor settings");
}
if (predictor == JxlPredictor.Variable)
{
for (int i = 0; i < JxlPredictorFacts.ModularPredictors; i++)
{
this.predictors.Add((JxlPredictor)i);
}
Span<JxlPredictor> predictorsSpan = CollectionsMarshal.AsSpan(this.predictors);
RuntimeUtility.Swap(ref predictorsSpan[0], ref predictorsSpan[(int)JxlPredictor.Weighted]);
RuntimeUtility.Swap(ref predictorsSpan[1], ref predictorsSpan[(int)JxlPredictor.Gradient]);
}
else if (predictor == JxlPredictor.Best)
{
this.predictors = [JxlPredictor.Weighted, JxlPredictor.Gradient];
}
else
{
this.predictors = [predictor];
}
if (wpTreeMode == JxlTreeMode.NoWp)
{
// delete all weighted predictors
_ = this.predictors.RemoveAll(p => p == JxlPredictor.Weighted);
}
this.residuals.Grow([], this.predictors.Count);
}
public void SetProperties(List<int> properties, JxlTreeMode wpTreeMode)
{
this.propertiesToUse = properties;
if (wpTreeMode == JxlTreeMode.WpOnly)
{
this.propertiesToUse = [JxlMaConstants.WpProp];
}
else if (wpTreeMode == JxlTreeMode.GradientOnly)
{
this.propertiesToUse = [JxlMaConstants.GradientProp];
}
else if (wpTreeMode == JxlTreeMode.NoWp)
{
// delete all weighted properties, tree mode
// says "no weighted predictors"
_ = this.propertiesToUse.RemoveAll(x => x == JxlMaConstants.WpProp);
}
if (this.propertiesToUse.Count == 0)
{
// could happen in default tree mode or when properties parameter
// is empty
throw new InvalidOperationException("Invalid property set configuration");
}
this.NumberOfStaticProperties = 0;
for (int i = 0; i < this.propertiesToUse.Count; ++i)
{
int prop = this.propertiesToUse[i];
if (prop < JxlPredictorFacts.StaticProperties)
{
if (i != prop)
{
throw new InvalidOperationException("Index is not equal to the property");
}
this.NumberOfStaticProperties++;
}
}
this.properties.Resize(this.propertiesToUse.Count - this.NumberOfStaticProperties);
}
public void InitializeTable(int logSize)
{
int size = 1 << logSize;
if (this.deduplicationTable.Count == size)
{
return;
}
this.deduplicationTable.Resize(size, DeduplicationEntryUnused);
for (int i = 0; i < this.NumberOfDistinctSamples; i++)
{
if (this.sampleCounts[i] != ushort.MaxValue)
{
this.AddToTable(i);
}
}
}
public void AddToTableAndMerge(int a)
{
int pos1 = Hash1(a);
int pos2 = Hash2(a);
if (this.deduplicationTable[pos1] != DeduplicationEntryUnused && this.IsSameSample(a, this.deduplicationTable[pos1]))
{
if (this.sampleCounts[a] != 1)
{
throw new InvalidOperationException("Sample count must be 1");
}
this.sampleCounts[this.deduplicationTable[pos1]]++;
if (this.sampleCounts[this.deduplicationTable[pos1]] == ushort.MaxValue)
{
this.deduplicationTable[pos1] = DeduplicationEntryUnused;
}
return;
}
if (this.deduplicationTable[pos2] != DeduplicationEntryUnused && this.IsSameSample(a, this.deduplicationTable[pos2]))
{
if (this.sampleCounts[a] != 1)
{
throw new InvalidOperationException("Sample count must be 1");
}
this.sampleCounts[this.deduplicationTable[pos2]]++;
if (this.sampleCounts[this.deduplicationTable[pos2]] == ushort.MaxValue)
{
this.deduplicationTable[pos2] = DeduplicationEntryUnused;
}
return;
}
this.AddToTable(a);
}
public void AddToTable(int a)
{
int pos1 = Hash1(a);
int pos2 = Hash2(a);
if (this.deduplicationTable[pos1] == DeduplicationEntryUnused)
{
this.deduplicationTable[pos1] = a;
}
else if (this.deduplicationTable[pos2] == DeduplicationEntryUnused)
{
this.deduplicationTable[pos2] = a;
}
}
public void PrepareForSamples(int extraNumSamples)
{
foreach (List<JxlResidualToken> residual in this.residuals)
{
residual.Grow(default, residual.Count + extraNumSamples);
}
for (int i = 0; i < this.NumberOfStaticProperties; i++)
{
this.staticProperties[i].Grow(0, this.staticProperties[i].Count + extraNumSamples);
}
foreach (List<byte> prop in this.properties)
{
prop.Grow((byte)0, prop.Count + extraNumSamples);
}
int totalNumSamples = extraNumSamples + this.sampleCounts.Count;
int nextSize = JxlMath.CeilLog2Nonzero(totalNumSamples * 3 / 2);
this.InitializeTable(nextSize);
}
public int Hash1(int a)
{
const ulong constant = 0x1e35a7bd;
ulong h = constant;
foreach (List<JxlResidualToken> r in this.residuals)
{
h = (h * constant) + (ulong)r[a].Token;
h = (h * constant) + (ulong)r[a].NumberOfBits;
}
for (int i = 0; i < this.NumberOfStaticProperties; i++)
{
h = (h * constant) + (ulong)this.staticProperties[i][a];
}
foreach (List<byte> property in this.properties)
{
h = (h * constant) + property[a];
}
return (int)((h >> 16) & (ulong)(this.deduplicationTable.Count - 1));
}
public int Hash2(int a)
{
const ulong constant = 0x1e35a7bd1e35a7bd;
ulong h = constant;
for (int i = 0; i < this.NumberOfStaticProperties; i++)
{
h = (h * constant) ^ (ulong)this.staticProperties[i][a];
}
foreach (List<byte> property in this.properties)
{
h = (h * constant) ^ property[a];
}
foreach (List<JxlResidualToken> r in this.residuals)
{
h = (h * constant) ^ (ulong)r[a].Token;
h = (h * constant) ^ (ulong)r[a].NumberOfBits;
}
return (int)((h >> 16) & (ulong)(this.deduplicationTable.Count - 1));
}
public bool IsSameSample(int a, int b)
{
foreach (List<JxlResidualToken> r in this.residuals)
{
if (r[a].Token != r[b].Token)
{
return false;
}
if (r[a].NumberOfBits != r[b].NumberOfBits)
{
return false;
}
}
for (int i = 0; i < this.NumberOfStaticProperties; ++i)
{
if (this.staticProperties[i][a] != this.staticProperties[i][b])
{
return false;
}
}
foreach (List<byte> p in this.properties)
{
if (p[a] != p[b])
{
return false;
}
}
return false;
}
public void AddSample(int pixel, Span<int> properties, Span<int> predictions)
{
for (int i = 0; i < this.predictors.Count; i++)
{
int v = pixel - predictions[(int)this.predictors[i]];
new JxlAnsHybridUIntConfiguration(4, 1, 2).Encode(JxlPackSigned.PackUnsigned(v), out uint tok, out uint nbits, out uint bits);
if (tok >= 256)
{
throw new InvalidOperationException("Token is too large");
}
if (nbits >= 256)
{
throw new InvalidOperationException("Number of bits is too large");
}
JxlResidualToken token = new((int)tok, (int)nbits);
this.residuals[i].Add(token);
}
for (int i = 0; i < this.NumberOfStaticProperties; ++i)
{
this.staticProperties[i].Add(this.QuantizeStaticProperty(i, properties[i]));
}
for (int i = this.NumberOfStaticProperties; i < this.propertiesToUse.Count; i++)
{
this.properties[i - this.NumberOfStaticProperties].Add(unchecked((byte)this.QuantizeProperty(i, properties[this.propertiesToUse[i]])));
}
this.sampleCounts.Add(1);
this.NumberOfSamples++;
this.AddToTableAndMerge(this.sampleCounts.Count - 1);
foreach (List<JxlResidualToken> residual in this.residuals)
{
// remove last item from List<T>
residual.RemoveAt(residual.Count - 1);
}
for (int i = 0; i < this.NumberOfStaticProperties; i++)
{
// ditto
this.staticProperties[i].RemoveAt(this.staticProperties[i].Count - 1);
}
foreach (List<byte> property in this.properties)
{
// ditto
property.RemoveAt(property.Count - 1);
}
// ditto
this.sampleCounts.RemoveAt(this.sampleCounts.Count - 1);
}
public void Swap(int a, int b)
{
if (a == b)
{
return;
}
foreach (List<JxlResidualToken> r in this.residuals)
{
// Get a Span for this List so we can get a ref to its items
// for use in RuntimeUtility.Swap (tuple-swap is slightly slower)
Span<JxlResidualToken> sp = CollectionsMarshal.AsSpan(r);
RuntimeUtility.Swap(ref sp[a], ref sp[b]);
}
for (int i = 0; i < this.NumberOfStaticProperties; i++)
{
// Ditto
Span<int> sp = CollectionsMarshal.AsSpan(this.staticProperties[i]);
RuntimeUtility.Swap(ref sp[a], ref sp[b]);
}
foreach (List<byte> p in this.properties)
{
// Ditto
Span<byte> sp = CollectionsMarshal.AsSpan(p);
RuntimeUtility.Swap(ref sp[a], ref sp[b]);
}
// Ditto
Span<int> sampleCounts = CollectionsMarshal.AsSpan(this.sampleCounts);
RuntimeUtility.Swap(ref sampleCounts[a], ref sampleCounts[b]);
}
public void PreQuantizeProperties(
Configuration configuration,
StaticPropertyRange range,
List<JxlModularMultiplierInfo> multiplierInfo,
List<int> groupPixelCount,
List<int> channelPixelCount,
List<int> pixelSamples,
List<int> diffSamples,
int maxPropertyValues)
{
List<int> groupMultiplierThresholds = [];
List<int> channelMultiplierThresholds = [];
foreach (JxlModularMultiplierInfo v in multiplierInfo)
{
if (v.Range[0][0] != range[0][0])
{
channelMultiplierThresholds.Add(v.Range[0][0] - 1);
}
if (v.Range[0][1] != range[0][1])
{
channelMultiplierThresholds.Add(v.Range[0][1] - 1);
}
if (v.Range[1][0] != range[1][0])
{
groupMultiplierThresholds.Add(v.Range[1][0] - 1);
}
if (v.Range[1][1] != range[1][1])
{
groupMultiplierThresholds.Add(v.Range[1][1] - 1);
}
}
channelMultiplierThresholds.Sort();
channelMultiplierThresholds.Resize(0, channelMultiplierThresholds.Distinct().Count());
groupMultiplierThresholds.Sort();
groupMultiplierThresholds.Resize(0, groupMultiplierThresholds.Distinct().Count());
this.compactProperties.Resize([], this.propertiesToUse.Count);
List<int> QuantizeChannel()
{
if (channelMultiplierThresholds.Count > 0)
{
return channelMultiplierThresholds;
}
return JxlMaEncoder.QuantizeHistogram(
CollectionsMarshal.AsSpan(groupPixelCount),
maxPropertyValues);
}
List<int> QuantizeGroupId()
{
if (groupMultiplierThresholds.Count > 0)
{
return groupMultiplierThresholds;
}
return JxlMaEncoder.QuantizeHistogram(
CollectionsMarshal.AsSpan(groupPixelCount),
maxPropertyValues);
}
List<int> QuantizeCoordinate()
{
List<int> quantized = new(maxPropertyValues - 1);
for (int i = 0; i + 1 < maxPropertyValues; i++)
{
quantized[i] = ((i + 1) * 256 / maxPropertyValues) - 1;
}
return quantized;
}
List<int> absPixelThresholds = [];
List<int> pixelThresholds = [];
List<int> QuantizePixelProperty()
{
if (pixelThresholds.Count == 0)
{
pixelThresholds = JxlMaEncoder.QuantizeSamples(
CollectionsMarshal.AsSpan(pixelSamples),
maxPropertyValues);
}
return pixelThresholds;
}
List<int> QuantizeAbsolutePixelProperty()
{
if (absPixelThresholds.Count == 0)
{
_ = QuantizePixelProperty(); // compute the non-abs thresholds
Span<int> pixelSamplesSpan = CollectionsMarshal.AsSpan(pixelSamples);
TensorPrimitives.Abs(pixelSamplesSpan, pixelSamplesSpan);
absPixelThresholds = JxlMaEncoder.QuantizeSamples(pixelSamplesSpan, maxPropertyValues);
}
return absPixelThresholds;
}
List<int> absoluteDiffThresholds = [];
List<int> diffThresholds = [];
List<int> QuantizeDiffProperty()
{
if (diffThresholds.Count == 0)
{
diffThresholds = JxlMaEncoder.QuantizeSamples(
CollectionsMarshal.AsSpan(diffSamples),
maxPropertyValues);
}
return diffThresholds;
}
List<int> QuantizeAbsoluteDiffProperty()
{
if (absoluteDiffThresholds.Count == 0)
{
_ = QuantizeDiffProperty();
Span<int> diffSamplesSpan = CollectionsMarshal.AsSpan(diffSamples);
TensorPrimitives.Abs(diffSamplesSpan, diffSamplesSpan);
absoluteDiffThresholds = JxlMaEncoder.QuantizeSamples(diffSamplesSpan, maxPropertyValues);
}
return absoluteDiffThresholds;
}
List<int> QuantizeWeightedPrediction()
{
// TODO: static ReadOnlySpan<int> ... => [...]?
if (maxPropertyValues < 32)
{
return [-127, -63, -31, -15, -7, -3, -1, 0, 1,
3, 7, 15, 31, 63, 127];
}
else if (maxPropertyValues < 64)
{
return [-255, -191, -127, -95, -63, -47, -31, -23,
-15, -11, -7, -5, -3, -1, 0, 1,
3, 5, 7, 11, 15, 23, 31, 47,
63, 95, 127, 191, 255];
}
else
{
return [-255, -223, -191, -159, -127, -111, -95, -79, -63, -55, -47,
-39, -31, -27, -23, -19, -15, -13, -11, -9, -7, -6,
-5, -4, -3, -2, -1, 0, 1, 2, 3, 4, 5,
6, 7, 9, 11, 13, 15, 19, 23, 27, 31, 39,
47, 55, 63, 79, 95, 111, 127, 159, 191, 223, 255];
}
}
this.propertyMapping.Resize(0, this.propertiesToUse.Count - this.NumberOfStaticProperties);
for (int i = 0; i < this.propertiesToUse.Count; i++)
{
if (this.propertiesToUse[i] == 0)
{
this.compactProperties[i] = QuantizeChannel();
}
else if (this.propertiesToUse[i] == 1)
{
this.compactProperties[i] = QuantizeGroupId();
}
else if (this.propertiesToUse[i] is 2 or 3)
{
this.compactProperties[i] = QuantizeCoordinate();
}
else if (this.propertiesToUse[i] == 6 || this.propertiesToUse[i] == 7 ||
this.propertiesToUse[i] == 8 ||
(this.propertiesToUse[i] >= JxlMaConstants.NumNonrefProperties && (this.propertiesToUse[i] - JxlMaConstants.NumNonrefProperties) % 4 == 1))
{
this.compactProperties[i] = QuantizePixelProperty();
}
else if (this.propertiesToUse[i] == 4 || this.propertiesToUse[i] == 5 ||
(this.propertiesToUse[i] >= JxlMaConstants.NumNonrefProperties && (this.propertiesToUse[i] - JxlMaConstants.NumNonrefProperties) % 4 == 0))
{
this.compactProperties[i] = QuantizeAbsolutePixelProperty();
}
else if (this.propertiesToUse[i] >= JxlMaConstants.NumNonrefProperties && (this.propertiesToUse[i] - JxlMaConstants.NumNonrefProperties) % 4 == 2)
{
this.compactProperties[i] = QuantizeAbsoluteDiffProperty();
}
else if (this.propertiesToUse[i] == JxlMaConstants.WpProp)
{
this.compactProperties[i] = QuantizeWeightedPrediction();
}
else
{
this.compactProperties[i] = QuantizeDiffProperty();
}
if (i < this.NumberOfStaticProperties)
{
JxlMaEncoder.QuantMap(
CollectionsMarshal.AsSpan(this.compactProperties[i]),
CollectionsMarshal.AsSpan(this.staticPropertyMapping[i]),
(PropertyRange * 2) + 1,
PropertyRange);
}
else
{
JxlMaEncoder.QuantMap(
CollectionsMarshal.AsSpan(this.compactProperties[i]),
CollectionsMarshal.AsSpan(this.propertyMapping[i - this.NumberOfStaticProperties]),
(PropertyRange * 2) + 1,
PropertyRange);
}
}
}
}

15
src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ListUtils.cs

@ -0,0 +1,15 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal static class ListUtils
{
public static void Grow<T>(this List<T> list, T item, int desiredSize)
{
while (list.Count < desiredSize)
{
list.Add(item);
}
}
}

15
src/ImageSharp/Formats/Jxl/Processing/Modular/README.md

@ -0,0 +1,15 @@
# Implementation of Modular coding mode for JPEG XL
JPEG XL files typically consist of two distinct coding modes: VarDCT
and Modular. VarDCT is lossy and transform-based, while Modular uses
the predictive coding approach and can achieve lossless or near-lossless
compression.
This folder has the implementation of the Modular encoder/decoder.
### Folder structure
In the Transforms folder, there are implementations for the Reversible
Color Transform (RCT), Squeeze Transform and Palette, both inverse
and forward.
In the Encoding folder there are predictors using neighbors, averages,
weighted prediction and tree-based prediction.

87
src/ImageSharp/Formats/Jxl/Processing/Primitives/Rng.cs

@ -0,0 +1,87 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
using System.Runtime.CompilerServices;
using SixLabors.ImageSharp.Common.Helpers;
namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Primitives;
/// <summary>
/// Deterministic random number generator used for compatibility
/// with JPEG XL.
/// </summary>
internal struct Rng
{
private ulong s0;
private ulong s1;
public Rng(ulong seed)
{
this.s0 = 0x94D049BB133111EBUL;
this.s1 = 0xBF58476D1CE4E5B9UL + seed;
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public ulong Next()
{
ulong s1 = this.s0;
ulong s0 = this.s1;
ulong bits = s1 + s0;
this.s0 = s0;
s1 ^= s1 << 23;
s1 ^= s0 ^ (s1 >> 18) ^ (s0 >> 5);
this.s1 = s1;
return bits;
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public long UniformI(long begin, long end) => (long)(this.Next() % (ulong)(end - begin)) + begin;
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public ulong UniformU(ulong begin, ulong end) => (this.Next() % (end - begin)) + begin;
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public float UniformF(float begin, float end)
{
uint u = (uint)(this.Next() >> (64 - 23)) | 0x3F800000u;
float f = BitConverter.UInt32BitsToSingle(u);
return ((end - begin) * (f - 1.0f)) + begin;
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public bool Bernoulli(float p) => this.UniformF(0, 1) < p;
internal readonly struct GeometricDistribution
{
public readonly float Value { get; }
public GeometricDistribution(float value) => this.Value = value;
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public static GeometricDistribution Make(float p) => new(1.0f / MathF.Log(1.0f - p));
}
[MethodImpl(MethodImplOptions.AggressiveInlining)]
public uint Geometric(in GeometricDistribution dist)
{
float f = this.UniformF(0, 1);
float invLog1mp = dist.Value;
float log = MathF.Log(1.0f - f) * invLog1mp;
return (uint)log;
}
public void Shuffle<T>(Span<T> span)
{
for (nuint i = 0; i + 1 < (nuint)span.Length; i++)
{
nuint a = (nuint)this.UniformU(i, (nuint)span.Length);
RuntimeUtility.Swap(ref span[(int)a], ref span[(int)i]);
}
}
}
Loading…
Cancel
Save