diff --git a/src/Nethermind/Nethermind.Core.Test/KeccakTests.cs b/src/Nethermind/Nethermind.Core.Test/KeccakTests.cs index ca07d65c5a6f..745146b895e5 100644 --- a/src/Nethermind/Nethermind.Core.Test/KeccakTests.cs +++ b/src/Nethermind/Nethermind.Core.Test/KeccakTests.cs @@ -5,6 +5,7 @@ using System.Buffers; using System.Collections.Generic; using System.Reflection; +using System.Runtime.Intrinsics.X86; using Nethermind.Core.Crypto; using Nethermind.Core.Extensions; using Nethermind.Serialization.Rlp; @@ -168,6 +169,39 @@ public void Computes_known_hash_for_span_and_array() } } + [Test] + public void Avx512_permutation_matches_scalar() + { + if (!Avx512F.IsSupported) + { + Assert.Ignore("AVX-512F intrinsics are not supported on this machine."); + } + + const int stateLength = 25; + ulong[] expected = new ulong[stateLength]; + ulong[] actual = new ulong[stateLength]; + + for (int testCase = 0; testCase < 64; testCase++) + { + for (int lane = 0; lane < stateLength; lane++) + { + actual[lane] = testCase switch + { + 0 => 0, + 1 => ulong.MaxValue, + _ => unchecked((ulong)(testCase * stateLength + lane + 1) * 0x9e3779b97f4a7c15UL) + }; + } + + actual.CopyTo(expected, 0); + + KeccakHash.KeccakF1600Scalar(ref expected[0]); + KeccakHash.KeccakF1600Avx512F(ref actual[0]); + + Assert.That(actual, Is.EqualTo(expected), $"Permutation mismatch for test case {testCase}."); + } + } + [TestCase("0x", "c5d2460186f7233c927e7db2dcc703c0e500b653ca82273b7bfad8045d85a470")] public void Sanity_check(string hexString, string expected) { diff --git a/src/Nethermind/Nethermind.Core/Crypto/Keccak.cs b/src/Nethermind/Nethermind.Core/Crypto/Keccak.cs index 00eb902c2d70..1232a9f11730 100644 --- a/src/Nethermind/Nethermind.Core/Crypto/Keccak.cs +++ b/src/Nethermind/Nethermind.Core/Crypto/Keccak.cs @@ -49,6 +49,7 @@ public static ValueHash256 Compute(string input) } [DebuggerStepThrough] + [SkipLocalsInit] public static ValueHash256 Compute(ReadOnlySpan input) { if (input.Length == 0) @@ -61,6 +62,7 @@ public static ValueHash256 Compute(ReadOnlySpan input) return keccak; } + [SkipLocalsInit] internal static ValueHash256 InternalCompute(byte[] input) { Unsafe.SkipInit(out ValueHash256 keccak); diff --git a/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.cs b/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.cs index 2341a5f492f0..71b321419ef8 100644 --- a/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.cs +++ b/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.cs @@ -15,6 +15,7 @@ public sealed partial class KeccakHash { private const int HASH_SIZE = 32; private const int STATE_SIZE = 200; + private const int STATE_LANES = STATE_SIZE / sizeof(ulong); private const int HASH_DATA_AREA = 136; private byte[] _remainderBuffer = []; @@ -69,6 +70,7 @@ private KeccakHash(int size) public KeccakHash Copy() => new(this); + [SkipLocalsInit] public static void ComputeHash(ReadOnlySpan input, Span output) { if ((uint)(output.Length - 1) >= STATE_SIZE) @@ -83,7 +85,10 @@ public static void ComputeHash(ReadOnlySpan input, Span output) #endif int roundSize = GetRoundSize(output.Length); - Span state = stackalloc ulong[STATE_SIZE / sizeof(ulong)]; + // A struct local rather than stackalloc: localloc would pin this method at Tier0-FullOpts + // (no tiering or dynamic PGO) and add GS-cookie and stack-probe overhead per call. + KeccakState stateBuffer = default; // the sponge state must start all-zero + Span state = stateBuffer; Span stateBytes = MemoryMarshal.AsBytes(state); if (input.Length == Address.Size) @@ -173,6 +178,7 @@ public static uint[] ComputeBytesToUint(ReadOnlySpan input, int size) return output; } + [SkipLocalsInit] public ValueHash256 GenerateValueHash() { Unsafe.SkipInit(out ValueHash256 output); @@ -180,6 +186,7 @@ public ValueHash256 GenerateValueHash() return output; } + [SkipLocalsInit] public void Update(ReadOnlySpan input) { if (_hash is not null) @@ -259,6 +266,7 @@ public void Update(ReadOnlySpan input) } } + [SkipLocalsInit] public void UpdateFinalTo(Span output) { if (_hash is not null) @@ -364,7 +372,8 @@ public void ResetTo(KeccakHash original) private static partial void KeccakF(Span st); - private static int GetRoundSize(int hashSize) => checked(STATE_SIZE - 2 * hashSize); + // Callers bound hashSize to [1, STATE_SIZE], so the arithmetic cannot overflow. + private static int GetRoundSize(int hashSize) => STATE_SIZE - 2 * hashSize; private byte[] GenerateHash() { @@ -469,6 +478,12 @@ private static unsafe void XorVectors(Span state, ReadOnlySpan input private static void ThrowInvalidOutputSize(int length) => throw new ArgumentOutOfRangeException( nameof(length), length, $"Must be between 1 and {STATE_SIZE}."); + [InlineArray(STATE_LANES)] + private struct KeccakState + { + private ulong _lane0; + } + private static class Pool { private const int MaxPooledPerThread = 4; diff --git a/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.std.cs b/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.std.cs index 3ccc9070bcf9..ca34791c15c6 100644 --- a/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.std.cs +++ b/src/Nethermind/Nethermind.Core/Crypto/KeccakHash.std.cs @@ -29,18 +29,32 @@ public sealed partial class KeccakHash 0x8000000000008080UL, 0x0000000080000001UL, 0x8000000080008008UL ]; + // Shared lane-index vectors (Theta's [1,2,3,4,0] == Chi's permute1) + private static readonly Vector512 LaneShift1 = Vector512.Create(1UL, 2UL, 3UL, 4UL, 0UL, 5UL, 6UL, 7UL); + private static readonly Vector512 LaneShift2 = Vector512.Create(2UL, 3UL, 4UL, 0UL, 1UL, 5UL, 6UL, 7UL); + private static readonly Vector512 ThetaRot4 = Vector512.Create(4UL, 0UL, 1UL, 2UL, 3UL, 5UL, 6UL, 7UL); + + // Pre-broadcast round constants (kills the per-round vmovq) + private static readonly Vector512[] RoundConstantVec = + Array.ConvertAll(RoundConstants, rc => Vector512.CreateScalar(rc)); + // update the state with given number of rounds private static partial void KeccakF(Span st) { + Debug.Assert(st.Length == STATE_LANES); + + ref ulong state = ref MemoryMarshal.GetReference(st); if (Avx512F.IsSupported) - KeccakF1600Avx512F(st); + KeccakF1600Avx512F(ref state); else - KeccakF1600(st); + KeccakF1600Scalar(ref state); } - private static void KeccakF1600(Span st) + /// Portable Keccak-f[1600] permutation. + /// Lane 0 of a 25-lane state; all 25 lanes are read and written. + internal static void KeccakF1600Scalar(ref ulong state) { - Debug.Assert(st.Length == 25); + Span st = MemoryMarshal.CreateSpan(ref state, STATE_LANES); ulong aba, abe, abi, abo, abu; ulong aga, age, agi, ago, agu; @@ -250,198 +264,98 @@ private static void KeccakF1600(Span st) st[0] = aba; } + /// AVX-512 Keccak-f[1600] permutation. + /// Lane 0 of a 25-lane state; all 25 lanes are read and written. [SkipLocalsInit] - public static void KeccakF1600Avx512F(Span state) + internal static void KeccakF1600Avx512F(ref ulong state) { - { - // Redundant statement that removes all the in loop bounds checks - _ = state[24]; - } - - // Can straight load and over-read for start elements - Vector512 mask = Vector512.Create(ulong.MaxValue, ulong.MaxValue, ulong.MaxValue, ulong.MaxValue, ulong.MaxValue, 0UL, 0UL, 0UL); - Vector512 c0 = Unsafe.As>(ref MemoryMarshal.GetReference(state)); - // Clear the over-read values from first vectors - c0 = Vector512.BitwiseAnd(mask, c0); - Vector512 c1 = Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 5)); - c1 = Vector512.BitwiseAnd(mask, c1); - Vector512 c2 = Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 10)); - c2 = Vector512.BitwiseAnd(mask, c2); - Vector512 c3 = Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 15)); - c3 = Vector512.BitwiseAnd(mask, c3); - - // Can't over-read for the last elements (8 items in vector 5 to be remaining) - // so read a Vector256 and ulong then combine - Vector256 c4a = Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 20)); - Vector256 c4b = Vector256.Create(state[24], 0UL, 0UL, 0UL); - Vector512 c4 = Vector512.Create(c4a, c4b); - - Vector512 permute1 = Vector512.Create(1UL, 2UL, 3UL, 4UL, 0UL, 5UL, 6UL, 7UL); - Vector512 permute2 = Vector512.Create(2UL, 3UL, 4UL, 0UL, 1UL, 5UL, 6UL, 7UL); - ulong[] roundConstants = RoundConstants; - - // Use constant for loop so Jit expects to loop; unroll once + Debug.Assert(RoundConstantVec.Length == ROUNDS && ROUNDS % 2 == 0); + + ref ulong s = ref state; + + // Lanes 5-7 hold over-read neighbor lanes and need no masking: Theta/Rho/Pi/Chi + // map result lanes 0-4 only from lanes 0-4, and the stores below overwrite 5-7. + Vector512 c0 = Unsafe.As>(ref s); + Vector512 c1 = Unsafe.As>(ref Unsafe.Add(ref s, 5)); + Vector512 c2 = Unsafe.As>(ref Unsafe.Add(ref s, 10)); + Vector512 c3 = Unsafe.As>(ref Unsafe.Add(ref s, 15)); + Vector512 c4 = Vector512.Create( + Unsafe.As>(ref Unsafe.Add(ref s, 20)), + Vector256.CreateScalar(Unsafe.Add(ref s, 24))); + + // The const bound lets the JIT hoist the vector constants out of the loop; + // bounding by RoundConstantVec.Length makes it re-load all of them every iteration. + ref Vector512 roundConstants = ref MemoryMarshal.GetArrayDataReference(RoundConstantVec); for (int round = 0; round < ROUNDS; round += 2) { - // Iteration 1 - { - ulong roundConstant = Unsafe.Add(ref MemoryMarshal.GetArrayDataReference(roundConstants), round); - // Theta step - Vector512 parity = Avx512F.TernaryLogic(Avx512F.TernaryLogic(c0, c1, c2, 0x96), c3, c4, 0x96); - - // Compute Theta - Vector512 bVecRot1Rotated = Avx512F.RotateLeft(Avx512F.PermuteVar8x64(parity, Vector512.Create(1UL, 2UL, 3UL, 4UL, 0UL, 5UL, 6UL, 7UL)), 1); - Vector512 bVecRot4 = Avx512F.PermuteVar8x64(parity, Vector512.Create(4UL, 0UL, 1UL, 2UL, 3UL, 5UL, 6UL, 7UL)); - Vector512 theta = Avx512F.Xor(bVecRot4, bVecRot1Rotated); - - c0 = Avx512F.Xor(c0, theta); - c1 = Avx512F.Xor(c1, theta); - c2 = Avx512F.Xor(c2, theta); - c3 = Avx512F.Xor(c3, theta); - c4 = Avx512F.Xor(c4, theta); - - // Rho step - Vector512 rhoVec0 = Vector512.Create(0UL, 1UL, 62UL, 28UL, 27UL, 0UL, 0UL, 0UL); - c0 = Avx512F.RotateLeftVariable(c0, rhoVec0); - - Vector512 rhoVec1 = Vector512.Create(36UL, 44UL, 6UL, 55UL, 20UL, 0UL, 0UL, 0UL); - c1 = Avx512F.RotateLeftVariable(c1, rhoVec1); - - Vector512 rhoVec2 = Vector512.Create(3UL, 10UL, 43UL, 25UL, 39UL, 0UL, 0UL, 0UL); - c2 = Avx512F.RotateLeftVariable(c2, rhoVec2); - - Vector512 rhoVec3 = Vector512.Create(41UL, 45UL, 15UL, 21UL, 8UL, 0UL, 0UL, 0UL); - c3 = Avx512F.RotateLeftVariable(c3, rhoVec3); - - Vector512 rhoVec4 = Vector512.Create(18UL, 2UL, 61UL, 56UL, 14UL, 0UL, 0UL, 0UL); - c4 = Avx512F.RotateLeftVariable(c4, rhoVec4); - - // Pi step - Vector512 c0Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(0UL, 8 + 1, 2, 3, 4, 5, 6, 7), c1); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 8 + 2, 3, 4, 5, 6, 7), c2); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 8 + 3, 4, 5, 6, 7), c3); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 4, 5, 6, 7), c4); - - Vector512 c1Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(3UL, 8 + 4, 2, 3, 4, 5, 6, 7), c1); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 8 + 0, 3, 4, 5, 6, 7), c2); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 8 + 1, 4, 5, 6, 7), c3); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 2, 5, 6, 7), c4); - - Vector512 c2Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(1UL, 8 + 2, 2, 3, 4, 5, 6, 7), c1); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 8 + 3, 3, 4, 5, 6, 7), c2); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 2, 8 + 4, 4, 5, 6, 7), c3); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 0, 5, 6, 7), c4); - - Vector512 c3Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(4UL, 8 + 0, 2, 3, 4, 5, 6, 7), c1); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 8 + 1, 3, 4, 5, 6, 7), c2); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 2, 8 + 2, 4, 5, 6, 7), c3); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 3, 5, 6, 7), c4); - - Vector512 c4Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(2UL, 8 + 3, 2, 3, 4, 5, 6, 7), c1); - c0 = c0Pi; - c1 = c1Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 8 + 4, 3, 4, 5, 6, 7), c2); - c2 = c2Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 2, 8 + 0, 4, 5, 6, 7), c3); - c3 = c3Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 1, 5, 6, 7), c4); - c4 = c4Pi; - - // Chi step - - c0 = Avx512F.TernaryLogic(c0, Avx512F.PermuteVar8x64(c0, permute1), Avx512F.PermuteVar8x64(c0, permute2), 0xD2); - c1 = Avx512F.TernaryLogic(c1, Avx512F.PermuteVar8x64(c1, permute1), Avx512F.PermuteVar8x64(c1, permute2), 0xD2); - c2 = Avx512F.TernaryLogic(c2, Avx512F.PermuteVar8x64(c2, permute1), Avx512F.PermuteVar8x64(c2, permute2), 0xD2); - c3 = Avx512F.TernaryLogic(c3, Avx512F.PermuteVar8x64(c3, permute1), Avx512F.PermuteVar8x64(c3, permute2), 0xD2); - c4 = Avx512F.TernaryLogic(c4, Avx512F.PermuteVar8x64(c4, permute1), Avx512F.PermuteVar8x64(c4, permute2), 0xD2); - - // Iota step - c0 = Vector512.Xor(c0, Vector512.Create(roundConstant, 0UL, 0UL, 0UL, 0UL, 0UL, 0UL, 0UL)); - } - // Iteration 2 - { - ulong roundConstant = Unsafe.Add(ref MemoryMarshal.GetArrayDataReference(roundConstants), round + 1); - // Theta step - Vector512 parity = Avx512F.TernaryLogic(Avx512F.TernaryLogic(c0, c1, c2, 0x96), c3, c4, 0x96); - - // Compute Theta - Vector512 bVecRot1Rotated = Avx512F.RotateLeft(Avx512F.PermuteVar8x64(parity, Vector512.Create(1UL, 2UL, 3UL, 4UL, 0UL, 5UL, 6UL, 7UL)), 1); - Vector512 bVecRot4 = Avx512F.PermuteVar8x64(parity, Vector512.Create(4UL, 0UL, 1UL, 2UL, 3UL, 5UL, 6UL, 7UL)); - Vector512 theta = Avx512F.Xor(bVecRot4, bVecRot1Rotated); - - c0 = Avx512F.Xor(c0, theta); - c1 = Avx512F.Xor(c1, theta); - c2 = Avx512F.Xor(c2, theta); - c3 = Avx512F.Xor(c3, theta); - c4 = Avx512F.Xor(c4, theta); - - // Rho step - Vector512 rhoVec0 = Vector512.Create(0UL, 1UL, 62UL, 28UL, 27UL, 0UL, 0UL, 0UL); - c0 = Avx512F.RotateLeftVariable(c0, rhoVec0); - - Vector512 rhoVec1 = Vector512.Create(36UL, 44UL, 6UL, 55UL, 20UL, 0UL, 0UL, 0UL); - c1 = Avx512F.RotateLeftVariable(c1, rhoVec1); - - Vector512 rhoVec2 = Vector512.Create(3UL, 10UL, 43UL, 25UL, 39UL, 0UL, 0UL, 0UL); - c2 = Avx512F.RotateLeftVariable(c2, rhoVec2); - - Vector512 rhoVec3 = Vector512.Create(41UL, 45UL, 15UL, 21UL, 8UL, 0UL, 0UL, 0UL); - c3 = Avx512F.RotateLeftVariable(c3, rhoVec3); - - Vector512 rhoVec4 = Vector512.Create(18UL, 2UL, 61UL, 56UL, 14UL, 0UL, 0UL, 0UL); - c4 = Avx512F.RotateLeftVariable(c4, rhoVec4); - - // Pi step - Vector512 c0Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(0UL, 8 + 1, 2, 3, 4, 5, 6, 7), c1); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 8 + 2, 3, 4, 5, 6, 7), c2); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 8 + 3, 4, 5, 6, 7), c3); - c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 4, 5, 6, 7), c4); - - Vector512 c1Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(3UL, 8 + 4, 2, 3, 4, 5, 6, 7), c1); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 8 + 0, 3, 4, 5, 6, 7), c2); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 8 + 1, 4, 5, 6, 7), c3); - c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 2, 5, 6, 7), c4); - - Vector512 c2Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(1UL, 8 + 2, 2, 3, 4, 5, 6, 7), c1); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 8 + 3, 3, 4, 5, 6, 7), c2); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 2, 8 + 4, 4, 5, 6, 7), c3); - c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 0, 5, 6, 7), c4); - - Vector512 c3Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(4UL, 8 + 0, 2, 3, 4, 5, 6, 7), c1); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 8 + 1, 3, 4, 5, 6, 7), c2); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 2, 8 + 2, 4, 5, 6, 7), c3); - c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 3, 5, 6, 7), c4); - - Vector512 c4Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(2UL, 8 + 3, 2, 3, 4, 5, 6, 7), c1); - c0 = c0Pi; - c1 = c1Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 8 + 4, 3, 4, 5, 6, 7), c2); - c2 = c2Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 2, 8 + 0, 4, 5, 6, 7), c3); - c3 = c3Pi; - c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 1, 5, 6, 7), c4); - c4 = c4Pi; - - // Chi step - - c0 = Avx512F.TernaryLogic(c0, Avx512F.PermuteVar8x64(c0, permute1), Avx512F.PermuteVar8x64(c0, permute2), 0xD2); - c1 = Avx512F.TernaryLogic(c1, Avx512F.PermuteVar8x64(c1, permute1), Avx512F.PermuteVar8x64(c1, permute2), 0xD2); - c2 = Avx512F.TernaryLogic(c2, Avx512F.PermuteVar8x64(c2, permute1), Avx512F.PermuteVar8x64(c2, permute2), 0xD2); - c3 = Avx512F.TernaryLogic(c3, Avx512F.PermuteVar8x64(c3, permute1), Avx512F.PermuteVar8x64(c3, permute2), 0xD2); - c4 = Avx512F.TernaryLogic(c4, Avx512F.PermuteVar8x64(c4, permute1), Avx512F.PermuteVar8x64(c4, permute2), 0xD2); - - // Iota step - c0 = Vector512.Xor(c0, Vector512.Create(roundConstant, 0UL, 0UL, 0UL, 0UL, 0UL, 0UL, 0UL)); - } + Round(ref c0, ref c1, ref c2, ref c3, ref c4, Unsafe.Add(ref roundConstants, round)); + Round(ref c0, ref c1, ref c2, ref c3, ref c4, Unsafe.Add(ref roundConstants, round + 1)); } - // Can over-write for first elements - Unsafe.As>(ref MemoryMarshal.GetReference(state)) = c0; - Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 5)) = c1; - Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 10)) = c2; - Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 15)) = c3; - // Can't over-write for last elements so write the upper Vector256 and then ulong - Unsafe.As>(ref Unsafe.Add(ref MemoryMarshal.GetReference(state), 20)) = c4.GetLower(); - state[24] = c4.GetElement(4); + Unsafe.As>(ref s) = c0; + Unsafe.As>(ref Unsafe.Add(ref s, 5)) = c1; + Unsafe.As>(ref Unsafe.Add(ref s, 10)) = c2; + Unsafe.As>(ref Unsafe.Add(ref s, 15)) = c3; + Unsafe.As>(ref Unsafe.Add(ref s, 20)) = c4.GetLower(); + Unsafe.Add(ref s, 24) = c4.GetElement(4); + } + + [SkipLocalsInit] + [MethodImpl(MethodImplOptions.AggressiveInlining)] + private static void Round(ref Vector512 c0, ref Vector512 c1, + ref Vector512 c2, ref Vector512 c3, ref Vector512 c4, Vector512 roundConstant) + { + // Theta + Vector512 parity = Avx512F.TernaryLogic(Avx512F.TernaryLogic(c0, c1, c2, 0x96), c3, c4, 0x96); + Vector512 theta = Avx512F.Xor( + Avx512F.PermuteVar8x64(parity, ThetaRot4), + Avx512F.RotateLeft(Avx512F.PermuteVar8x64(parity, LaneShift1), 1)); + c0 = Avx512F.Xor(c0, theta); c1 = Avx512F.Xor(c1, theta); + c2 = Avx512F.Xor(c2, theta); c3 = Avx512F.Xor(c3, theta); c4 = Avx512F.Xor(c4, theta); + + // Rho + c0 = Avx512F.RotateLeftVariable(c0, Vector512.Create(0UL, 1UL, 62UL, 28UL, 27UL, 0UL, 0UL, 0UL)); + c1 = Avx512F.RotateLeftVariable(c1, Vector512.Create(36UL, 44UL, 6UL, 55UL, 20UL, 0UL, 0UL, 0UL)); + c2 = Avx512F.RotateLeftVariable(c2, Vector512.Create(3UL, 10UL, 43UL, 25UL, 39UL, 0UL, 0UL, 0UL)); + c3 = Avx512F.RotateLeftVariable(c3, Vector512.Create(41UL, 45UL, 15UL, 21UL, 8UL, 0UL, 0UL, 0UL)); + c4 = Avx512F.RotateLeftVariable(c4, Vector512.Create(18UL, 2UL, 61UL, 56UL, 14UL, 0UL, 0UL, 0UL)); + + // Pi + Vector512 c0Pi = Avx512F.PermuteVar8x64x2(c0, Vector512.Create(0UL, 8 + 1, 2, 3, 4, 5, 6, 7), c1); + c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 8 + 2, 3, 4, 5, 6, 7), c2); + c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 8 + 3, 4, 5, 6, 7), c3); + c0Pi = Avx512F.PermuteVar8x64x2(c0Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 4, 5, 6, 7), c4); + + Vector512 c1Pi = Avx512F.PermuteVar8x64x2(c1, Vector512.Create(0UL, 4UL, 8 + 0, 3, 4, 5, 6, 7), c2); + c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 8 + 1, 4, 5, 6, 7), c3); + c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 2, 5, 6, 7), c4); + c1Pi = Avx512F.PermuteVar8x64x2(c1Pi, Vector512.Create(8UL + 3, 1, 2, 3, 4, 5, 6, 7), c0); + + Vector512 c2Pi = Avx512F.PermuteVar8x64x2(c2, Vector512.Create(0UL, 1, 3UL, 8 + 4, 4, 5, 6, 7), c3); + c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 1, 2, 3, 8 + 0, 5, 6, 7), c4); + c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(8UL + 1, 1, 2, 3, 4, 5, 6, 7), c0); + c2Pi = Avx512F.PermuteVar8x64x2(c2Pi, Vector512.Create(0UL, 8 + 2, 2, 3, 4, 5, 6, 7), c1); + + Vector512 c3Pi = Avx512F.PermuteVar8x64x2(c3, Vector512.Create(0UL, 1, 2, 2UL, 8 + 3, 5, 6, 7), c4); + c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(8UL + 4, 1, 2, 3, 4, 5, 6, 7), c0); + c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 8 + 0, 2, 3, 4, 5, 6, 7), c1); + c3Pi = Avx512F.PermuteVar8x64x2(c3Pi, Vector512.Create(0UL, 1, 8 + 1, 3, 4, 5, 6, 7), c2); + + Vector512 c4Pi = Avx512F.PermuteVar8x64x2(c4, Vector512.Create(8 + 2, 1, 2, 3, 1UL, 5, 6, 7), c0); + c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 8 + 3, 2, 3, 4, 5, 6, 7), c1); + c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 8 + 4, 3, 4, 5, 6, 7), c2); + c4Pi = Avx512F.PermuteVar8x64x2(c4Pi, Vector512.Create(0UL, 1, 2, 8 + 0, 4, 5, 6, 7), c3); + + c0 = c0Pi; c1 = c1Pi; c2 = c2Pi; c3 = c3Pi; c4 = c4Pi; + + // Chi + c0 = Avx512F.TernaryLogic(c0, Avx512F.PermuteVar8x64(c0, LaneShift1), Avx512F.PermuteVar8x64(c0, LaneShift2), 0xD2); + c1 = Avx512F.TernaryLogic(c1, Avx512F.PermuteVar8x64(c1, LaneShift1), Avx512F.PermuteVar8x64(c1, LaneShift2), 0xD2); + c2 = Avx512F.TernaryLogic(c2, Avx512F.PermuteVar8x64(c2, LaneShift1), Avx512F.PermuteVar8x64(c2, LaneShift2), 0xD2); + c3 = Avx512F.TernaryLogic(c3, Avx512F.PermuteVar8x64(c3, LaneShift1), Avx512F.PermuteVar8x64(c3, LaneShift2), 0xD2); + c4 = Avx512F.TernaryLogic(c4, Avx512F.PermuteVar8x64(c4, LaneShift1), Avx512F.PermuteVar8x64(c4, LaneShift2), 0xD2); + + // Iota + c0 = Vector512.Xor(c0, roundConstant); } }