Split accumulator 3-Way

Squeeze a tiny bit more juice from the original idea in #6336 which this is on top of.

https://tests.stockfishchess.org/tests/view/68dddd85fa806e2e8393c0b9
LLR: 2.95 (-2.94,2.94) <0.00,2.00>
Total: 156320 W: 40925 L: 40447 D: 74948
Ptnml(0-2): 427, 17330, 42172, 17800, 431

4-way doesn't look to be better than this.
https://tests.stockfishchess.org/tests/view/68dde19efa806e2e8393c0c1

closes https://github.com/official-stockfish/Stockfish/pull/6339

No functional change

Co-authored-by: M Stembera <m_stembera@yahoo.com>
This commit is contained in:
Timothy Herchen
2025-10-05 09:33:23 +02:00
committed by Joost VandeVondele
co-authored by M Stembera
parent 7a7c033a86
commit b09339a420
+43 -10
View File
@@ -248,6 +248,7 @@ class AffineTransformSparseInput {
#if defined(USE_AVX512)
using invec_t = __m512i;
using outvec_t = __m512i;
#define vec_add_32 _mm512_add_epi32
#define vec_set_32 _mm512_set1_epi32
#define vec_add_dpbusd_32 SIMD::m512_add_dpbusd_epi32
#elif defined(USE_AVX2)
@@ -274,7 +275,16 @@ class AffineTransformSparseInput {
static constexpr IndexType OutputSimdWidth = sizeof(outvec_t) / sizeof(OutputType);
constexpr IndexType NumChunks = ceil_to_multiple<IndexType>(InputDimensions, 8) / ChunkSize;
constexpr IndexType NumRegs = OutputDimensions / OutputSimdWidth;
constexpr IndexType NumAccums = OutputDimensions / OutputSimdWidth;
// If there's only one accumulator and we're using high-latency dot product instructions,
// split it to create three separate dependency chains and merge at the end
constexpr bool SplitAccums =
#if defined(USE_VNNI)
NumAccums == 1;
#else
false;
#endif
constexpr IndexType NumRegs = SplitAccums ? 3 * NumAccums : NumAccums;
std::uint16_t nnz[NumChunks];
IndexType count;
@@ -285,27 +295,50 @@ class AffineTransformSparseInput {
const outvec_t* biasvec = reinterpret_cast<const outvec_t*>(biases);
outvec_t acc[NumRegs];
for (IndexType k = 0; k < NumRegs; ++k)
for (IndexType k = 0; k < NumAccums; ++k)
acc[k] = biasvec[k];
auto* start = nnz;
auto* end = nnz + count;
const auto* start = nnz;
const auto* end = nnz + count;
// convince GCC to not do weird pointer arithmetic in the following loop
const std::int8_t* weights_cp = weights;
if constexpr (SplitAccums)
{
acc[1] = acc[2] = vec_set_32(0);
while (start < end - 2)
{
const std::ptrdiff_t i0 = *start++;
const std::ptrdiff_t i1 = *start++;
const std::ptrdiff_t i2 = *start++;
const invec_t in0 = vec_set_32(input32[i0]);
const invec_t in1 = vec_set_32(input32[i1]);
const invec_t in2 = vec_set_32(input32[i2]);
const auto col0 =
reinterpret_cast<const invec_t*>(&weights_cp[i0 * OutputDimensions * ChunkSize]);
const auto col1 =
reinterpret_cast<const invec_t*>(&weights_cp[i1 * OutputDimensions * ChunkSize]);
const auto col2 =
reinterpret_cast<const invec_t*>(&weights_cp[i2 * OutputDimensions * ChunkSize]);
vec_add_dpbusd_32(acc[0], in0, *col0);
vec_add_dpbusd_32(acc[1], in1, *col1);
vec_add_dpbusd_32(acc[2], in2, *col2);
}
acc[0] = vec_add_32(vec_add_32(acc[0], acc[1]), acc[2]);
}
while (start < end)
{
const std::ptrdiff_t i = *start;
start++;
const invec_t in = vec_set_32(input32[i]);
const auto col = (const invec_t*) (&weights_cp[i * OutputDimensions * ChunkSize]);
for (IndexType k = 0; k < NumRegs; ++k)
const std::ptrdiff_t i = *start++;
const invec_t in = vec_set_32(input32[i]);
const auto col =
reinterpret_cast<const invec_t*>(&weights_cp[i * OutputDimensions * ChunkSize]);
for (IndexType k = 0; k < NumAccums; ++k)
vec_add_dpbusd_32(acc[k], in, col[k]);
}
outvec_t* outptr = reinterpret_cast<outvec_t*>(output);
for (IndexType k = 0; k < NumRegs; ++k)
for (IndexType k = 0; k < NumAccums; ++k)
outptr[k] = acc[k];
#undef vec_set_32
#undef vec_add_dpbusd_32