From 0e78190b72775d24e5e24747363c465d7307f7b7 Mon Sep 17 00:00:00 2001 From: Ryan Hirsch rhirsch123 Date: Sun, 26 Apr 2026 09:41:05 +0200 Subject: [PATCH] Speed up find_nnz on neon (Refactored) Passed STC: LLR: 2.97 (-2.94,2.94) <0.00,2.00> Total: 349568 W: 90131 L: 89393 D: 170044 Ptnml(0-2): 725, 37712, 97280, 38234, 833 https://tests.stockfishchess.org/tests/view/69e470510e9667dd5a765046 This is a refactor of https://github.com/official-stockfish/Stockfish/pull/6766 so we don't have to include nnue_architecture.h It also includes the fix in https://github.com/official-stockfish/Stockfish/pull/6769 which was needed to make this work. mstembera made rhirsch123 the author. credits also to anematode. [Edit] Also fixes a bug where _mm512_packus_epi32 was being used instead of the correct _mm512_packs_epi32 for the IceLake version. closes https://github.com/official-stockfish/Stockfish/pull/6774 No functional change --- .../layers/affine_transform_sparse_input.h | 293 ++++++++++-------- 1 file changed, 169 insertions(+), 124 deletions(-) diff --git a/src/nnue/layers/affine_transform_sparse_input.h b/src/nnue/layers/affine_transform_sparse_input.h index 39c5166f3..0b09ae6c4 100644 --- a/src/nnue/layers/affine_transform_sparse_input.h +++ b/src/nnue/layers/affine_transform_sparse_input.h @@ -24,6 +24,7 @@ #include #include #include +#include #include #include "../../bitboard.h" @@ -48,129 +49,6 @@ constexpr int constexpr_lsb(uint64_t bb) { constexpr uint64_t debruijn64 = 0x03F79D71B4CB0A89ULL; return lsb_index64[((bb ^ (bb - 1)) * debruijn64) >> 58]; } - -alignas(CacheLineSize) static constexpr struct OffsetIndices { - - std::uint16_t offset_indices[256][8]; - - constexpr OffsetIndices() : - offset_indices() { - for (int i = 0; i < 256; ++i) - { - std::uint64_t j = i, k = 0; - while (j) - { - offset_indices[i][k++] = constexpr_lsb(j); - j &= j - 1; - } - while (k < 8) - offset_indices[i][k++] = 0; - } - } - -} Lookup; - - #if defined(__GNUC__) || defined(__clang__) - #define RESTRICT __restrict__ - #elif defined(_MSC_VER) - #define RESTRICT __restrict - #else - #define RESTRICT - #endif - -// Find indices of nonzero 32-bit values in a packed byte buffer. -// The input pointer addresses a sequence of 32-bit blocks stored in a -// std::uint8_t array. -template -void find_nnz(const std::uint8_t* RESTRICT input, - std::uint16_t* RESTRICT out, - IndexType& count_out) { - - #if defined(USE_AVX512ICL) - - constexpr IndexType SimdWidthIn = 64; // 512 bits - constexpr IndexType SimdWidthOut = 32; // 512 bits / 16 bits - constexpr IndexType NumChunks = InputDimensions / SimdWidthOut; - const __m512i increment = _mm512_set1_epi16(SimdWidthOut); - __m512i base = _mm512_set_epi16( // Same permute order as _mm512_packus_epi32() - 31, 30, 29, 28, 15, 14, 13, 12, 27, 26, 25, 24, 11, 10, 9, 8, 23, 22, 21, 20, 7, 6, 5, 4, 19, - 18, 17, 16, 3, 2, 1, 0); - - IndexType count = 0; - for (IndexType i = 0; i < NumChunks; ++i) - { - const __m512i inputV0 = _mm512_load_si512(input + i * 2 * SimdWidthIn); - const __m512i inputV1 = _mm512_load_si512(input + i * 2 * SimdWidthIn + SimdWidthIn); - - // Get a bitmask and gather non zero indices - const __m512i inputV01 = _mm512_packus_epi32(inputV0, inputV1); - const __mmask32 nnzMask = _mm512_test_epi16_mask(inputV01, inputV01); - - // Avoid _mm512_mask_compressstoreu_epi16() as it's 256 uOps on Zen4 - __m512i nnz = _mm512_maskz_compress_epi16(nnzMask, base); - _mm512_storeu_si512(out + count, nnz); - - count += popcount(nnzMask); - base = _mm512_add_epi16(base, increment); - } - count_out = count; - - #elif defined(USE_AVX512) - - constexpr IndexType SimdWidth = 16; // 512 bits / 32 bits - constexpr IndexType NumChunks = InputDimensions / SimdWidth; - const __m512i increment = _mm512_set1_epi32(SimdWidth); - __m512i base = _mm512_set_epi32(15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0); - - IndexType count = 0; - for (IndexType i = 0; i < NumChunks; ++i) - { - const __m512i inputV = _mm512_load_si512(input + i * SimdWidth * sizeof(std::uint32_t)); - - // Get a bitmask and gather non zero indices - const __mmask16 nnzMask = _mm512_test_epi32_mask(inputV, inputV); - const __m512i nnzV = _mm512_maskz_compress_epi32(nnzMask, base); - _mm512_mask_cvtepi32_storeu_epi16(out + count, 0xFFFF, nnzV); - count += popcount(nnzMask); - base = _mm512_add_epi32(base, increment); - } - count_out = count; - - #else - - using namespace SIMD; - - constexpr IndexType InputSimdWidth = sizeof(vec_uint_t) / sizeof(std::int32_t); - // Outputs are processed 8 elements at a time, even if the SIMD width is narrower - constexpr IndexType ChunkSize = 8; - constexpr IndexType NumChunks = InputDimensions / ChunkSize; - constexpr IndexType InputsPerChunk = ChunkSize / InputSimdWidth; - - static_assert(InputsPerChunk > 0 && "SIMD width too wide"); - - const auto inputVector = reinterpret_cast(input); - IndexType count = 0; - vec128_t base = vec128_zero; - const vec128_t increment = vec128_set_16(8); - for (IndexType i = 0; i < NumChunks; ++i) - { - // bitmask of nonzero values in this chunk - unsigned nnz = 0; - for (IndexType j = 0; j < InputsPerChunk; ++j) - { - const vec_uint_t inputChunk = inputVector[i * InputsPerChunk + j]; - nnz |= unsigned(vec_nnz(inputChunk)) << (j * InputSimdWidth); - } - const vec128_t offsets = - vec128_load(reinterpret_cast(&Lookup.offset_indices[nnz])); - vec128_storeu(reinterpret_cast(out + count), vec128_add(base, offsets)); - count += popcount(nnz); - base = vec128_add(base, increment); - } - count_out = count; - #endif -} - #endif // Sparse input implementation @@ -199,6 +77,35 @@ class AffineTransformSparseInput { static constexpr IndexType ChunkSize = 1; #endif +#if defined(USE_NEON) + using NNZOutputType = std::conditional_t<(InDims <= 1024), std::uint8_t, std::uint16_t>; +#else + using NNZOutputType = std::uint16_t; +#endif + +#if (USE_SSSE3 | (USE_NEON >= 8)) + alignas(CacheLineSize) static constexpr struct OffsetIndices { + + NNZOutputType offset_indices[256][8]; + + constexpr OffsetIndices() : + offset_indices() { + for (int i = 0; i < 256; ++i) + { + std::uint64_t j = i, k = 0; + while (j) + { + offset_indices[i][k++] = constexpr_lsb(j); + j &= j - 1; + } + while (k < 8) + offset_indices[i][k++] = 0; + } + } + + } Lookup = {}; +#endif + using OutputBuffer = OutputType[PaddedOutputDimensions]; // Hash value embedded in the evaluation file @@ -293,7 +200,7 @@ class AffineTransformSparseInput { #else NumAccums; #endif - std::uint16_t nnz[NumChunks]; + NNZOutputType nnz[NumChunks]; IndexType count; // Find indices of nonzero 32-bit blocks @@ -377,6 +284,144 @@ class AffineTransformSparseInput { } private: +#if (USE_SSSE3 | (USE_NEON >= 8)) + + #if defined(__GNUC__) || defined(__clang__) + #define RESTRICT __restrict__ + #elif defined(_MSC_VER) + #define RESTRICT __restrict + #else + #define RESTRICT + #endif + + // Find indices of nonzero 32-bit blocks in a packed + // std::uint8_t buffer. NumChunks is the number of blocks. + template + static void find_nnz(const std::uint8_t* RESTRICT input, + NNZOutputType* RESTRICT out, + IndexType& count_out) { + + #if defined(USE_AVX512ICL) + + constexpr IndexType SimdWidthIn = 64; // 512 bits + constexpr IndexType SimdWidthOut = 32; // 512 bits / 16 bits + constexpr IndexType SimdChunks = NumChunks / SimdWidthOut; + const __m512i increment = _mm512_set1_epi16(SimdWidthOut); + __m512i base = _mm512_set_epi16( // Same permute order as _mm512_packus_epi32() + 31, 30, 29, 28, 15, 14, 13, 12, 27, 26, 25, 24, 11, 10, 9, 8, 23, 22, 21, 20, 7, 6, 5, 4, + 19, 18, 17, 16, 3, 2, 1, 0); + + IndexType count = 0; + for (IndexType i = 0; i < SimdChunks; ++i) + { + const __m512i inputV0 = _mm512_load_si512(input + i * 2 * SimdWidthIn); + const __m512i inputV1 = _mm512_load_si512(input + i * 2 * SimdWidthIn + SimdWidthIn); + + // Get a bitmask and gather non zero indices + const __m512i inputV01 = _mm512_packs_epi32(inputV0, inputV1); + const __mmask32 nnzMask = _mm512_test_epi16_mask(inputV01, inputV01); + + // Avoid _mm512_mask_compressstoreu_epi16() as it's 256 uOps on Zen4 + __m512i nnz = _mm512_maskz_compress_epi16(nnzMask, base); + _mm512_storeu_si512(out + count, nnz); + + count += popcount(nnzMask); + base = _mm512_add_epi16(base, increment); + } + count_out = count; + + #elif defined(USE_AVX512) + + constexpr IndexType SimdWidth = 16; // 512 bits / 32 bits + constexpr IndexType SimdChunks = NumChunks / SimdWidth; + const __m512i increment = _mm512_set1_epi32(SimdWidth); + __m512i base = _mm512_set_epi32(15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0); + + IndexType count = 0; + for (IndexType i = 0; i < SimdChunks; ++i) + { + const __m512i inputV = _mm512_load_si512(input + i * SimdWidth * sizeof(std::uint32_t)); + + // Get a bitmask and gather non zero indices + const __mmask16 nnzMask = _mm512_test_epi32_mask(inputV, inputV); + const __m512i nnzV = _mm512_maskz_compress_epi32(nnzMask, base); + _mm512_mask_cvtepi32_storeu_epi16(out + count, 0xFFFF, nnzV); + count += popcount(nnzMask); + base = _mm512_add_epi32(base, increment); + } + count_out = count; + + #else + #if defined(USE_NEON) + // NEON path using uint8_t NNZOutputType + if constexpr (std::is_same_v) + { + static_assert(NumChunks <= 256, "NumChunks must be <= 256"); + + static constexpr std::uint16_t nnzMask[8] = {1, 2, 4, 8, 16, 32, 64, 128}; + constexpr IndexType SimdChunks = NumChunks / 8; + const auto inputVector = reinterpret_cast(input); + + IndexType count = 0; + uint64_t base = 0ULL; + const uint64_t increment = 0x0808080808080808ULL; + for (IndexType i = 0; i < SimdChunks; ++i) + { + const uint32x4_t v0 = inputVector[i * 2]; + const uint32x4_t v1 = inputVector[i * 2 + 1]; + const uint16x8_t nnz = + vcombine_u16(vqmovn_u32(vtstq_u32(v0, v0)), vqmovn_u32(vtstq_u32(v1, v1))); + const uint16_t lookup = vaddvq_u16(vandq_u16(nnz, vld1q_u16(nnzMask))); + + uint64_t offsets; + std::memcpy(&offsets, Lookup.offset_indices[lookup], sizeof(offsets)); + const uint64_t indices = offsets + base; + std::memcpy(out + count, &indices, sizeof(indices)); + + count += popcount(lookup); + base += increment; + } + count_out = count; + } + else + #endif + { + using namespace SIMD; + + constexpr IndexType InputSimdWidth = sizeof(vec_uint_t) / sizeof(std::int32_t); + // Outputs are processed 8 elements at a time, even if the SIMD width is narrower + constexpr IndexType SimdChunkSize = 8; + constexpr IndexType SimdChunks = NumChunks / SimdChunkSize; + constexpr IndexType InputsPerChunk = SimdChunkSize / InputSimdWidth; + + static_assert(InputsPerChunk > 0, "SIMD width too wide"); + + const auto inputVector = reinterpret_cast(input); + IndexType count = 0; + vec128_t base = vec128_zero; + const vec128_t increment = vec128_set_16(8); + for (IndexType i = 0; i < SimdChunks; ++i) + { + // bitmask of nonzero values in this chunk + unsigned nnz = 0; + for (IndexType j = 0; j < InputsPerChunk; ++j) + { + const vec_uint_t inputChunk = inputVector[i * InputsPerChunk + j]; + nnz |= unsigned(vec_nnz(inputChunk)) << (j * InputSimdWidth); + } + const vec128_t offsets = + vec128_load(reinterpret_cast(&Lookup.offset_indices[nnz])); + vec128_storeu(reinterpret_cast(out + count), vec128_add(base, offsets)); + count += popcount(nnz); + base = vec128_add(base, increment); + } + count_out = count; + } + #endif + } + #undef RESTRICT +#endif + using BiasType = OutputType; using WeightType = std::int8_t;