diff --git a/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlBt709TransferFunction.cs b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlBt709TransferFunction.cs
index 99b6548b60..dfc9de66ea 100644
--- a/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlBt709TransferFunction.cs
+++ b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlBt709TransferFunction.cs
@@ -1,6 +1,9 @@
// Copyright (c) Six Labors.
// Licensed under the Six Labors Split License.
+using System.Numerics;
+using SixLabors.ImageSharp.Formats.Jxl.Processing;
+
namespace SixLabors.ImageSharp.Formats.Jxl.Cms.TransferFunctions;
///
@@ -8,13 +11,47 @@ namespace SixLabors.ImageSharp.Formats.Jxl.Cms.TransferFunctions;
///
internal static class JxlBt709TransferFunction
{
- public static double EncodedFromDisplay(double d)
+ // Encoded From Display constants
+ private const float Threshold = 0.018f;
+ private const float MulLow = 4.5f;
+ private const float MulHi = 1.099f;
+ private const float PowHi = 0.45f;
+ private const float Sub = -0.099f;
+
+ // Display From Encoded constants
+ private const float InverseThreshold = 0.081f;
+ private const float InverseMulLow = 1f / 4.5f;
+ private const float InverseMulHi = 1f / 1.099f;
+ private const float InversePowHi = 1f / 0.45f;
+ private const float InverseAdd = 0.099f * InverseMulHi;
+
+ public static float EncodedFromDisplay(float d)
{
if (d < Threshold)
{
return MulLow * d;
}
- return (MulHi * Math.Pow(d, PowHi)) + Sub;
+ return (MulHi * MathF.Pow(d, PowHi)) + Sub;
+ }
+
+ public static Vector EncodedFromDisplay(Vector x)
+ {
+ Vector low = Vector.Create(MulLow) * x;
+ Vector high = (Vector.Create(MulHi) * JxlSimdUtils.FastPowf(x, Vector.Create(PowHi))) + Vector.Create(Sub);
+ return Vector.ConditionalSelect(
+ Vector.LessThanOrEqual(x, Vector.Create(Threshold)),
+ low,
+ high);
+ }
+
+ public static Vector DisplayFromEncoded(Vector x)
+ {
+ Vector low = Vector.Create(InverseMulLow) * x;
+ Vector high = JxlSimdUtils.FastPowf((x * Vector.Create(InverseMulHi)) + Vector.Create(InverseAdd), Vector.Create(InversePowHi));
+ return Vector.ConditionalSelect(
+ Vector.LessThan(x, Vector.Create(InverseThreshold)),
+ low,
+ high);
}
}
diff --git a/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlSRgbTransferFunction.cs b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlSRgbTransferFunction.cs
new file mode 100644
index 0000000000..adc817ff78
--- /dev/null
+++ b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlSRgbTransferFunction.cs
@@ -0,0 +1,82 @@
+// Copyright (c) Six Labors.
+// Licensed under the Six Labors Split License.
+
+using System.Numerics;
+
+namespace SixLabors.ImageSharp.Formats.Jxl.Cms.TransferFunctions;
+
+internal static class JxlSRgbTransferFunction
+{
+ private const float ThresholdSRGBToLinear = 0.04045f;
+ private const float ThresholdLinearToSRGB = 0.0031308f;
+ private const float LowDiv = 12.92f;
+ private const float LowDivInverse = 1.0f / LowDiv;
+
+ private static ReadOnlySpan DisplayP =>
+ [
+ 2.200248328e-04f, 2.200248328e-04f, 2.200248328e-04f, 2.200248328e-04f,
+ 1.043637593e-02f, 1.043637593e-02f, 1.043637593e-02f, 1.043637593e-02f,
+ 1.624820318e-01f, 1.624820318e-01f, 1.624820318e-01f, 1.624820318e-01f,
+ 7.961564959e-01f, 7.961564959e-01f, 7.961564959e-01f, 7.961564959e-01f,
+ 8.210152774e-01f, 8.210152774e-01f, 8.210152774e-01f, 8.210152774e-01f
+ ];
+
+ private static ReadOnlySpan DisplayQ =>
+ [
+ 2.631846970e-01f, 2.631846970e-01f, 2.631846970e-01f, 2.631846970e-01f,
+ 1.076976492e+00f, 1.076976492e+00f, 1.076976492e+00f, 1.076976492e+00f,
+ 4.987528350e-01f, 4.987528350e-01f, 4.987528350e-01f, 4.987528350e-01f,
+ -5.512498495e-02f, -5.512498495e-02f, -5.512498495e-02f, -5.512498495e-02f,
+ 6.521209011e-03f, 6.521209011e-03f, 6.521209011e-03f, 6.521209011e-03f
+ ];
+
+ private static ReadOnlySpan EncodedP =>
+ [
+ -5.135152395e-04f, -5.135152395e-04f, -5.135152395e-04f, -5.135152395e-04f,
+ 5.287254571e-03f, 5.287254571e-03f, 5.287254571e-03f, 5.287254571e-03f,
+ 3.903842876e-01f, 3.903842876e-01f, 3.903842876e-01f, 3.903842876e-01f,
+ 1.474205315e+00f, 1.474205315e+00f, 1.474205315e+00f, 1.474205315e+00f,
+ 7.352629620e-01f, 7.352629620e-01f, 7.352629620e-01f, 7.352629620e-01f
+ ];
+
+ private static ReadOnlySpan EncodedQ =>
+ [
+ 1.004519624e-02f, 1.004519624e-02f, 1.004519624e-02f, 1.004519624e-02f,
+ 3.036675394e-01f, 3.036675394e-01f, 3.036675394e-01f, 3.036675394e-01f,
+ 1.340816930e+00f, 1.340816930e+00f, 1.340816930e+00f, 1.340816930e+00f,
+ 9.258482155e-01f, 9.258482155e-01f, 9.258482155e-01f, 9.258482155e-01f,
+ 2.424867759e-02f, 2.424867759e-02f, 2.424867759e-02f, 2.424867759e-02f
+ ];
+
+ public static Vector DisplayFromEncoded(Vector x)
+ {
+ Vector sign = Vector.Create(0x80000000u).As();
+ Vector originalSign = x & sign;
+ x = Vector.AndNot(sign, x);
+
+ Vector linear = x * Vector.Create(LowDivInverse);
+ Vector poly = EvaluateRationalPolynomial(x, DisplayP, DisplayQ);
+ Vector magnitude = Vector.ConditionalSelect(
+ Vector.GreaterThan(x, Vector.Create(ThresholdSRGBToLinear)),
+ poly,
+ linear);
+
+ return Vector.AndNot(sign, magnitude) | originalSign;
+ }
+
+ public static Vector EncodedFromDisplay(Vector x)
+ {
+ Vector sign = Vector.Create(0x80000000u).As();
+ Vector originalSign = x & sign;
+ x = Vector.AndNot(sign, x);
+
+ Vector linear = x * Vector.Create(LowDiv);
+ Vector poly = EvaluateRationalPolynomial(Vector.SquareRoot(x), DisplayP, DisplayQ);
+ Vector magnitude = Vector.ConditionalSelect(
+ Vector.GreaterThan(x, Vector.Create(ThresholdLinearToSRGB)),
+ poly,
+ linear);
+
+ return Vector.AndNot(sign, magnitude) | originalSign;
+ }
+}
diff --git a/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlTransferUtils.cs b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlTransferUtils.cs
new file mode 100644
index 0000000000..8b9760c73d
--- /dev/null
+++ b/src/ImageSharp/Formats/Jxl/Cms/TransferFunctions/JxlTransferUtils.cs
@@ -0,0 +1,71 @@
+// Copyright (c) Six Labors.
+// Licensed under the Six Labors Split License.
+
+using System.Numerics;
+
+namespace SixLabors.ImageSharp.Formats.Jxl.Cms.TransferFunctions;
+
+internal static class JxlTransferUtils
+{
+ public static Vector FastLinearToSrgb(Vector v)
+ {
+ Vector v025_05 = ((v.As() & new Vector(0x3effffff)) | new Vector(0x3e800000)).As();
+
+ Vector d1 = (v025_05 * Vector.Create(0.059914046f)) - Vector.Create(0.108894556f);
+ Vector d2 = (d1 * v025_05) + Vector.Create(0.107963754f);
+ Vector pow = (d2 * v025_05) + Vector.Create(0.018092343f);
+
+ const uint baseBits = 0x40000000;
+
+ ReadOnlySpan powers25to18 =
+ [
+ 0x00, 0x0a, 0x19, 0x26,
+ 0x32, 0x41, 0x4d, 0x5c,
+ 0x68, 0x75, 0x83, 0x8f,
+ 0xa0, 0xaa, 0xb9, 0xc6
+ ];
+
+ ReadOnlySpan powers17to10 =
+ [
+ 0x00, 0xb7, 0x04, 0x0d,
+ 0xcb, 0xe7, 0x41, 0x68,
+ 0x51, 0xd1, 0xeb, 0xf2,
+ 0x00, 0xb7, 0x04, 0x0d
+ ];
+
+ Vector bits = Vector.AsVectorInt32(v);
+ Vector exp = (bits >> 23) - Vector.Create(118);
+
+ exp &= new Vector(0xf);
+
+ Vector p25to18 = GatherByteTable(exp, powers25to18);
+ Vector p17to10 = GatherByteTable(exp, powers17to10);
+
+ Vector mulBits =
+ (p25to18 << 18) |
+ (p17to10 << 10) |
+ new Vector((int)baseBits);
+
+ Vector mul = Vector.AsVectorSingle(mulBits);
+
+ Vector cutoff = new(0.0031308f);
+
+ return Vector.ConditionalSelect(
+ Vector.LessThan(v, cutoff),
+ v * new Vector(12.92f),
+ (pow * mul) - new Vector(0.055f));
+ }
+
+ private static Vector GatherByteTable(Vector indices, ReadOnlySpan table)
+ {
+ int count = Vector.Count;
+ Span result = stackalloc int[count];
+
+ for (int i = 0; i < count; i++)
+ {
+ result[i] = table[indices[i]];
+ }
+
+ return new Vector(result);
+ }
+}
diff --git a/src/ImageSharp/Formats/Jxl/Processing/JxlUnsafe.cs b/src/ImageSharp/Formats/Jxl/Processing/JxlUnsafe.cs
new file mode 100644
index 0000000000..5060a08c66
--- /dev/null
+++ b/src/ImageSharp/Formats/Jxl/Processing/JxlUnsafe.cs
@@ -0,0 +1,17 @@
+// Copyright (c) Six Labors.
+// Licensed under the Six Labors Split License.
+
+using System.Numerics;
+
+namespace SixLabors.ImageSharp.Formats.Jxl.Processing;
+
+internal static unsafe class JxlUnsafe
+{
+ public static bool IsSimdAligned(T* ptr)
+ where T : unmanaged
+ => IsAligned(ptr, Vector.Count * sizeof(T));
+
+ public static bool IsAligned(T* ptr, int alignment)
+ where T : unmanaged
+ => ((long)ptr % alignment) == 0;
+}
diff --git a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs
index 009cb49c6e..115136b1bb 100644
--- a/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs
+++ b/src/ImageSharp/Formats/Jxl/Processing/Modular/Encoding/JxlMaEncoder.cs
@@ -17,7 +17,7 @@ namespace SixLabors.ImageSharp.Formats.Jxl.Processing.Modular.Encoding;
internal static class JxlMaEncoder
{
- internal enum IntersectionType
+ internal enum IntersectionType : byte
{
None,
Partial,
diff --git a/src/ImageSharp/Formats/Jxl/Processing/RenderPipeline/WriteToOutputStage.cs b/src/ImageSharp/Formats/Jxl/Processing/RenderPipeline/WriteToOutputStage.cs
index 2d1b989c9a..62ae1a00fa 100644
--- a/src/ImageSharp/Formats/Jxl/Processing/RenderPipeline/WriteToOutputStage.cs
+++ b/src/ImageSharp/Formats/Jxl/Processing/RenderPipeline/WriteToOutputStage.cs
@@ -4,7 +4,6 @@
using System.Numerics;
using System.Runtime.CompilerServices;
using System.Runtime.InteropServices;
-using System.Runtime.Intrinsics;
using SixLabors.ImageSharp.Formats.Jxl.IO.Metadata;
using SixLabors.ImageSharp.Formats.Jxl.Processing.Image;
@@ -14,6 +13,23 @@ internal sealed class WriteToOutputStage
{
private const int ChunkSize = 1024;
+ private int width;
+ private int height;
+ private IJxlImageOutput main;
+ private int numColors;
+ private bool wantAlpha;
+ private bool hasAlpha;
+ private bool unpremultiplyAlpha;
+ private int alphaC;
+ private bool flipX;
+ private bool flipY;
+ private bool transpose;
+ private readonly List extraChannels = [];
+ private readonly List opaqueAlpha = [];
+ private readonly Configuration configuration;
+ private List> tempIn;
+ private List> tempOut;
+
///
/// Gets the 32x32 blue noise dithering pattern lookup
/// table (from ),
@@ -234,4 +250,85 @@ internal sealed class WriteToOutputStage
JxlExifOrientation.Rotate90 or
JxlExifOrientation.Rotate270 or
JxlExifOrientation.AntiTranspose;
+
+ private unsafe void UnpremultiplyAlpha(int threadId, int len, float** lineBuffers)
+ {
+ // Highly unsafe code! ⚠️
+ float** tempIn = stackalloc float*[4];
+ Vector one = Vector.One;
+
+ for (int c = 0; c < this.main.PixelFormat.Channels; ++c)
+ {
+ // size_t tix = thread_id * main_.num_channels_ + c;
+ // temp_in[c] = temp_in_[tix].address();
+ // memcpy(temp_in[c], line_buffers[c], sizeof(float) * len);
+ int tix = (threadId * this.main.PixelFormat.Channels) + c;
+
+ tempIn[c] = (float*)Unsafe.AsPointer(ref MemoryMarshal.Cast(this.tempIn[tix].Span)[0]);
+
+ MemoryMarshal.CreateSpan(ref Unsafe.AsRef(lineBuffers[c]), len)
+ .CopyTo(MemoryMarshal.CreateSpan(ref Unsafe.AsRef(tempIn[c]), len));
+ }
+
+ Vector smallAlpha = Vector.Create(SmallAlpha);
+
+ for (int ix = 0; ix < len; ix += Vector.Count)
+ {
+ float* ptr = tempIn[this.numColors + ix];
+
+ // Using an aligned and unaligned branch
+ // REVIEW: does the branch outweigh alignment? we will have to benchmark this
+ // when the codec can build
+ if (JxlUnsafe.IsSimdAligned(ptr))
+ {
+ Vector alpha = Vector.LoadAlignedNonTemporal(tempIn[this.numColors] + ix);
+ Vector mul = one / Vector.Max(smallAlpha, alpha);
+
+ for (int c = 0; c < this.numColors; ++c)
+ {
+ float* currPtr = tempIn[c] + ix;
+
+ if (JxlUnsafe.IsSimdAligned(currPtr))
+ {
+ Vector val = Vector.LoadAlignedNonTemporal(currPtr);
+ Vector.StoreAlignedNonTemporal(val * mul, currPtr);
+ }
+ else
+ {
+ Vector val = Vector.Load(currPtr);
+ Vector.Store(val * mul, currPtr);
+ }
+ }
+ }
+ else
+ {
+ Vector alpha = Vector.Load(tempIn[this.numColors] + ix);
+ Vector mul = one / Vector.Max(smallAlpha, alpha);
+
+ for (int c = 0; c < this.numColors; ++c)
+ {
+ float* currPtr = tempIn[c] + ix;
+
+ if (JxlUnsafe.IsSimdAligned(currPtr))
+ {
+ Vector val = Vector.LoadAlignedNonTemporal(currPtr);
+ Vector.StoreAlignedNonTemporal(val * mul, currPtr);
+ }
+ else
+ {
+ Vector val = Vector.Load(currPtr);
+ Vector.Store(val * mul, currPtr);
+ }
+ }
+ }
+ }
+
+ for (int c = 0; c < this.main.PixelFormat.Channels; c++)
+ {
+ fixed (byte* ptr = this.tempIn[c].Span)
+ {
+ lineBuffers[c] = (float*)ptr;
+ }
+ }
+ }
}