diff --git a/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlToken.cs b/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlToken.cs index 2b07fb14d9..feb293a98d 100644 --- a/src/ImageSharp/Formats/Jxl/Processing/Encoder/Ans/JxlToken.cs +++ b/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; } diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlContextPrediction.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlContextPrediction.cs index f856ef0e39..caeb0e68c6 100644 --- a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ContextPrediction/JxlContextPrediction.cs +++ b/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) { diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs index db4f1be9cb..6e375e3da7 100644 --- a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaConstants.cs +++ b/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; } diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs new file mode 100644 index 0000000000..009cb49c6e --- /dev/null +++ b/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>; +using Tree = System.Collections.Generic.List; + +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.Count); + + public static float EstimateBits(Span counts, int numSymbols) + { + int total = TensorPrimitives.Sum((ReadOnlySpan)counts.Slice(0, numSymbols)); + + Vector minprob = Vector.Create(1.0f / JxlAnsConstants.AnsTableSize); + Vector inverseTotal = Vector.Create(1.0f / total); + Vector bitsLanes = Vector.Zero; + + for (int i = 0; i < numSymbols; i += Vector.Count) + { + Vector countsIv = new(counts[i..]); + Vector countsFv = Vector.ConvertToSingle(countsIv); + Vector probs = countsFv * inverseTotal; + Vector mprobs = probs * minprob; + + Vector nbps = Vector.ConditionalSelect( + Vector.Equals(countsIv, Vector.Create(total)), + Vector.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 residualTokens, List countIncrease, List 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 residualTokens, Span countIncrease, Span 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 mulInfo, StaticPropertyRange initialStaticPropertyRange, float fastDecodeMultiplier, Tree tree) + { + Stack nodes = []; + nodes.Push(new(0, 0, treeSamples.NumberOfDistinctSamples, initialStaticPropertyRange)); + + int numPred = treeSamples.NumberOfPredictors; + int numProp = treeSamples.NumberOfProperties; + + Span 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.Shared.Rent(maxSymbols * numPred); + + for (int pred = 0; pred < numPred; pred++) + { + int extraBits = 0; + List 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 countsIncrease = []; + List extraBitsIncrease = []; + + List leftCosts = []; + List rightCosts = []; + + int[] aboveCounts = ArrayPool.Shared.Rent(maxSymbols); + int[] belowCounts = ArrayPool.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.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 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.Shared.Return(aboveCounts); + ArrayPool.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.Shared.Return(counts); + } + } + + public static void ComputeBestTree(JxlTreeSamples treeSamples, float threshold, List 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 QuantizeHistogram(ReadOnlySpan histogram, int numChunks) + { + if (histogram.Length == 0 || numChunks == 0) + { + return []; + } + + int sum = TensorPrimitives.Sum(histogram); + + if (sum == 0) + { + return []; + } + + List 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 QuantizeSamples(ReadOnlySpan samples, int numChunks) + { + const int range = 512; + + if (samples.Length == 0) + { + return []; + } + + int min = Math.Clamp(TensorPrimitives.Min(samples), -range, range); + + Span 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 thresholds = QuantizeHistogram(counts, numChunks); + + // For fast processing via TensorPrimitives + Span thresholdsSpan = CollectionsMarshal.AsSpan(thresholds); + TensorPrimitives.Add(thresholdsSpan, min, thresholdsSpan); + + return thresholds; + } + + public static void QuantMap(Span from, Span 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 groupPixelCount, Span channelPixelCount, List pixelSamples, List diffSamples) + { + if (options.NumberOfRepeats == 0) + { + throw new InvalidOperationException("No repeats"); + } + + // Power of 2 size is unknown. + Span alignedGroupPixelCount = stackalloc int[groupPixelCount.Length <= groupId ? groupId + 1 : groupPixelCount.Length]; + groupPixelCount.CopyTo(alignedGroupPixelCount); + + Span 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 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 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 tokens, Tree decoderTree) + { + if (tree.Count > JxlMaConstants.MaxTreeSize) + { + throw new InvalidOperationException("Too many tree nodes"); + } + + Queue 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; + } +} diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs index 2c06e80ba7..8c04283516 100644 --- a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaTreeContext.cs +++ b/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, diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlResidualToken.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlResidualToken.cs new file mode 100644 index 0000000000..8a8b387200 --- /dev/null +++ b/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; + } +} diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeSamples.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlTreeSamples.cs new file mode 100644 index 0000000000..1d25122266 --- /dev/null +++ b/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>; + +namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding; + +/// +/// Contains all the necessary data needed to build a tree. +/// +internal sealed class JxlTreeSamples +{ + /// + /// This deduplication table entry marks an unused entry. + /// + private const int DeduplicationEntryUnused = -1; + + private const int PropertyRange = 511; + + /// + /// Residual information: token and number of extra bits per predictor. + /// + private List> residuals = []; + + /// + /// Number of occurrences of each sample. + /// + private readonly List sampleCounts = []; + + /// + /// Quantized static property values. + /// + private InlineArray2> staticProperties; + + /// + /// Property values, quantized to at most 256 distinct values. + /// + private readonly List> properties = []; + + /// + /// Decompactification info for . + /// + private readonly List> compactProperties = []; + + /// + /// List of properties to use. + /// + private List propertiesToUse = []; + + /// + /// List of predictors to use. + /// + private List predictors = []; + + /// + /// Mapping property value -> quantized property value. + /// + private InlineArray2> staticPropertyMapping; + + private readonly List> propertyMapping = []; + + /// + /// Table for deduplication. + /// + private List deduplicationTable = []; + + /// + /// Gets a value indicating whether there are any residual samples. + /// + public bool HasSamples => this.residuals.Count > 0 && this.residuals[0].Count > 0; + + /// + /// Gets the total number of distinct samples. + /// + public int NumberOfDistinctSamples => this.sampleCounts.Count; + + /// + /// Gets the total number of samples. + /// + public int NumberOfSamples { get; private set; } + + /// + /// Gets the number of quantized static property values. + /// + public int NumberOfStaticProperties { get; private set; } + + /// + /// Gets the total number of predictors. + /// + public int NumberOfPredictors => this.predictors.Count; + + /// + /// Gets the total number of properties. + /// + public int NumberOfProperties => this.propertiesToUse.Count; + + /// + /// Returns a List of residual tokens for the specified kind of prediction. + /// + /// The type of prediction. + /// Residual tokens corresponding to the prediction specified by the parameter . + public List GetResidualTokensForPrediction(int prediction) => this.residuals[prediction]; + + /// + /// Returns a reference to the residual token. + /// + /// Kind of prediction. + /// Index of the residual token within that prediction. + /// The residual token for prediction indexed . + public ref JxlResidualToken GetResidualToken(int pred, int i) => ref CollectionsMarshal.AsSpan(this.residuals[pred])[i]; + + /// + /// Returns a token of the residual for prediction index . + /// + /// The kind of prediction. + /// Index of the residual token within that prediction. + /// For residual token whose prediction is and index is , returns its token coefficient. + public int GetToken(int prediction, int index) => this.residuals[prediction][index].Token; + + /// + /// Returns the number of occurrences for sample . + /// + /// The index of the sample. + /// Number of times appears. + public int GetCount(int i) => this.sampleCounts[i]; + + /// + /// Finds the index of the predictor . + /// + /// The predictor to find the index for. + /// Index of the predictor in the predictors storage. + /// Thrown when the predictor can't be found. + public int FindPredictorIndex(JxlPredictor predictor) + { + ReadOnlySpan 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; + } + + /// + /// Finds the index of the property . + /// + /// The property to find the index for. + /// Index of the property in the properties storage. + /// Thrown when the property isn't valid. + public int FindPropertyIndex(int property) + { + ReadOnlySpan 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; + } + + /// + /// Returns the number of properties for a . + /// + /// Index of the property. + /// Number of property values for property with index . + public int CountPropertyValues(int propertyIndex) => this.compactProperties[propertyIndex].Count + 1; + + /// + /// Returns the value of a property. + /// + /// Prefer a static property? + /// The index of the properties table. + /// The index of the property. + /// + /// Property for index within table . + /// + [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]; + } + } + + /// + /// Returns the dequantized property. + /// + /// Index of the target property. + /// Property quantizer. + /// The dequantized property. + /// Thrown when the quant is out of range. + public int UnquantizeProperty(int propertyIndex, int quant) + { + List compactProperties = this.compactProperties[propertyIndex]; + + if (quant >= compactProperties.Count) + { + throw new InvalidOperationException("Quant is out of range"); + } + + return compactProperties[quant]; + } + + /// + /// Returns a predictor for index . + /// + /// The index of the predictor. + /// Predictor at index . + public JxlPredictor PredictorFromIndex(int index) + { + DebugGuard.MustBeLessThan(index, this.predictors.Count, nameof(index)); + + return this.predictors[index]; + } + + /// + /// Returns a property for index . + /// + /// The index of the property. + /// Property at index . + public int PropertyFromIndex(int index) + { + DebugGuard.MustBeLessThan(index, this.propertiesToUse.Count, nameof(index)); + + return this.propertiesToUse[index]; + } + + /// + /// Invoked after processing samples completed. + /// + public void AllSamplesDone() => this.deduplicationTable = []; + + /// + /// Returns the quantized property value of . + /// + /// Index of the property. + /// Value to quantize. + /// Quantized value. + 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]; + } + + /// + /// Returns the quantized static property value of . + /// + /// Index of the static property. + /// Value to quantize. + /// Quantized value. + 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 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 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 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 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 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 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 property in this.properties) + { + h = (h * constant) ^ property[a]; + } + + foreach (List 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 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 p in this.properties) + { + if (p[a] != p[b]) + { + return false; + } + } + + return false; + } + + public void AddSample(int pixel, Span properties, Span 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 residual in this.residuals) + { + // remove last item from List + 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 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 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 sp = CollectionsMarshal.AsSpan(r); + + RuntimeUtility.Swap(ref sp[a], ref sp[b]); + } + + for (int i = 0; i < this.NumberOfStaticProperties; i++) + { + // Ditto + Span sp = CollectionsMarshal.AsSpan(this.staticProperties[i]); + + RuntimeUtility.Swap(ref sp[a], ref sp[b]); + } + + foreach (List p in this.properties) + { + // Ditto + Span sp = CollectionsMarshal.AsSpan(p); + + RuntimeUtility.Swap(ref sp[a], ref sp[b]); + } + + // Ditto + Span sampleCounts = CollectionsMarshal.AsSpan(this.sampleCounts); + + RuntimeUtility.Swap(ref sampleCounts[a], ref sampleCounts[b]); + } + + public void PreQuantizeProperties( + Configuration configuration, + StaticPropertyRange range, + List multiplierInfo, + List groupPixelCount, + List channelPixelCount, + List pixelSamples, + List diffSamples, + int maxPropertyValues) + { + List groupMultiplierThresholds = []; + List 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 QuantizeChannel() + { + if (channelMultiplierThresholds.Count > 0) + { + return channelMultiplierThresholds; + } + + return JxlMaEncoder.QuantizeHistogram( + CollectionsMarshal.AsSpan(groupPixelCount), + maxPropertyValues); + } + + List QuantizeGroupId() + { + if (groupMultiplierThresholds.Count > 0) + { + return groupMultiplierThresholds; + } + + return JxlMaEncoder.QuantizeHistogram( + CollectionsMarshal.AsSpan(groupPixelCount), + maxPropertyValues); + } + + List QuantizeCoordinate() + { + List quantized = new(maxPropertyValues - 1); + + for (int i = 0; i + 1 < maxPropertyValues; i++) + { + quantized[i] = ((i + 1) * 256 / maxPropertyValues) - 1; + } + + return quantized; + } + + List absPixelThresholds = []; + List pixelThresholds = []; + + List QuantizePixelProperty() + { + if (pixelThresholds.Count == 0) + { + pixelThresholds = JxlMaEncoder.QuantizeSamples( + CollectionsMarshal.AsSpan(pixelSamples), + maxPropertyValues); + } + + return pixelThresholds; + } + + List QuantizeAbsolutePixelProperty() + { + if (absPixelThresholds.Count == 0) + { + _ = QuantizePixelProperty(); // compute the non-abs thresholds + + Span pixelSamplesSpan = CollectionsMarshal.AsSpan(pixelSamples); + TensorPrimitives.Abs(pixelSamplesSpan, pixelSamplesSpan); + + absPixelThresholds = JxlMaEncoder.QuantizeSamples(pixelSamplesSpan, maxPropertyValues); + } + + return absPixelThresholds; + } + + List absoluteDiffThresholds = []; + List diffThresholds = []; + + List QuantizeDiffProperty() + { + if (diffThresholds.Count == 0) + { + diffThresholds = JxlMaEncoder.QuantizeSamples( + CollectionsMarshal.AsSpan(diffSamples), + maxPropertyValues); + } + + return diffThresholds; + } + + List QuantizeAbsoluteDiffProperty() + { + if (absoluteDiffThresholds.Count == 0) + { + _ = QuantizeDiffProperty(); + + Span diffSamplesSpan = CollectionsMarshal.AsSpan(diffSamples); + TensorPrimitives.Abs(diffSamplesSpan, diffSamplesSpan); + + absoluteDiffThresholds = JxlMaEncoder.QuantizeSamples(diffSamplesSpan, maxPropertyValues); + } + + return absoluteDiffThresholds; + } + + List QuantizeWeightedPrediction() + { + // TODO: static ReadOnlySpan ... => [...]? + 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); + } + } + } +} diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ListUtils.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/ListUtils.cs new file mode 100644 index 0000000000..c203fc66a0 --- /dev/null +++ b/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(this List list, T item, int desiredSize) + { + while (list.Count < desiredSize) + { + list.Add(item); + } + } +} diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/README.md b/src/ImageSharp/Formats/Jxl/Processing/Modular/README.md new file mode 100644 index 0000000000..3339605a40 --- /dev/null +++ b/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. diff --git a/src/ImageSharp/Formats/Jxl/Processing/Primitives/Rng.cs b/src/ImageSharp/Formats/Jxl/Processing/Primitives/Rng.cs new file mode 100644 index 0000000000..e11f243f8b --- /dev/null +++ b/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; + +/// +/// Deterministic random number generator used for compatibility +/// with JPEG XL. +/// +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(Span 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]); + } + } +}