From b4f2f4f4501e5eed6b3cb997de96a852d8e9df48 Mon Sep 17 00:00:00 2001 From: Disservin Date: Tue, 4 Aug 2026 20:49:20 +0200 Subject: [PATCH] Cleanup FeatureTransformer Types --- src/nnue/nnue_feature_transformer.h | 77 +++++++++++++---------------- 1 file changed, 33 insertions(+), 44 deletions(-) diff --git a/src/nnue/nnue_feature_transformer.h b/src/nnue/nnue_feature_transformer.h index 4f4aa88dfd0..3edd763e2b4 100644 --- a/src/nnue/nnue_feature_transformer.h +++ b/src/nnue/nnue_feature_transformer.h @@ -88,9 +88,22 @@ class FeatureTransformer { // Number of input/output dimensions static constexpr IndexType ThreatInputDimensions = ThreatFeatureSet::Dimensions; static constexpr IndexType PairInputDimensions = PairFeatureSet::Dimensions; - static constexpr IndexType InputDimensions = - PSQFeatureSet::Dimensions + ThreatInputDimensions + PairInputDimensions; - static constexpr IndexType OutputDimensions = HalfDimensions; + static constexpr IndexType PsqDimensions = PSQFeatureSet::Dimensions; + static constexpr IndexType ThreatAndPpDimensions = ThreatInputDimensions + PairInputDimensions; + static constexpr IndexType InputDimensions = PsqDimensions + ThreatAndPpDimensions; + static constexpr IndexType OutputDimensions = HalfDimensions; + static constexpr IndexType ThreatWeightSize = ThreatInputDimensions * HalfDimensions; + static constexpr IndexType ThreatPsqtWeightSize = ThreatInputDimensions * PSQTBuckets; + static constexpr IndexType PairWeightSize = PairInputDimensions * HalfDimensions; + static constexpr IndexType PairPsqtWeightSize = PairInputDimensions * PSQTBuckets; + static constexpr IndexType ThreatAndPpWeightSize = ThreatAndPpDimensions * HalfDimensions; + static constexpr IndexType ThreatAndPpPsqtSize = ThreatAndPpDimensions * PSQTBuckets; + + using BiasesArray = std::array; + using WeightArray = std::array; + using ThreatAndPpWeightArray = std::array; + using PsqtWeightArray = std::array; + using ThreatAndPpPsqtArray = std::array; // Size of forward propagation buffer static constexpr usize BufferSize = OutputDimensions * sizeof(OutputType); @@ -148,33 +161,20 @@ class FeatureTransformer { permute<8>(threatAndPpWeights, InversePackusEpi16Order); } - auto threatWeights() { return threatAndPpWeights.data(); } - auto threatWeights() const { return threatAndPpWeights.data(); } - auto ppWeights() { return &threatAndPpWeights[ThreatFeatureSet::Dimensions * HalfDimensions]; } - auto ppWeights() const { - return &threatAndPpWeights[ThreatFeatureSet::Dimensions * HalfDimensions]; - } - - auto threatPsqtWeights() { return threatAndPpPsqtWeights.data(); } - auto threatPsqtWeights() const { return threatAndPpPsqtWeights.data(); } - auto ppPsqtWeights() { - return &threatAndPpPsqtWeights[ThreatFeatureSet::Dimensions * PSQTBuckets]; - } - auto ppPsqtWeights() const { - return &threatAndPpPsqtWeights[ThreatFeatureSet::Dimensions * PSQTBuckets]; - } + ThreatWeightType* threatWeightData() { return threatAndPpWeights.data(); } + ThreatWeightType* pawnPairWeightData() { return threatWeightData() + ThreatWeightSize; } + PSQTWeightType* threatPsqtData() { return threatAndPpPsqtWeights.data(); } + PSQTWeightType* pawnPairPsqtData() { return threatPsqtData() + ThreatPsqtWeightSize; } // Read network parameters bool read_parameters(std::istream& stream) { read_leb_128(stream, biases); - read_little_endian(stream, threatWeights(), - ThreatInputDimensions * HalfDimensions); - read_leb_128(stream, threatPsqtWeights(), ThreatFeatureSet::Dimensions * PSQTBuckets); - read_little_endian(stream, ppWeights(), - PairInputDimensions * HalfDimensions); - read_leb_128(stream, ppPsqtWeights(), PairFeatureSet::Dimensions * PSQTBuckets); + read_little_endian(stream, threatWeightData(), ThreatWeightSize); + read_leb_128(stream, threatPsqtData(), ThreatPsqtWeightSize); + read_little_endian(stream, pawnPairWeightData(), PairWeightSize); + read_leb_128(stream, pawnPairPsqtData(), PairPsqtWeightSize); read_leb_128(stream, weights); read_leb_128(stream, psqtWeights); @@ -193,14 +193,10 @@ class FeatureTransformer { write_leb_128(stream, copy->biases); - write_little_endian(stream, copy->threatWeights(), - ThreatInputDimensions * HalfDimensions); - write_leb_128(stream, copy->threatPsqtWeights(), - ThreatFeatureSet::Dimensions * PSQTBuckets); - write_little_endian(stream, copy->ppWeights(), - PairInputDimensions * HalfDimensions); - write_leb_128(stream, copy->ppPsqtWeights(), - PairFeatureSet::Dimensions * PSQTBuckets); + write_little_endian(stream, copy->threatWeightData(), ThreatWeightSize); + write_leb_128(stream, copy->threatPsqtData(), ThreatPsqtWeightSize); + write_little_endian(stream, copy->pawnPairWeightData(), PairWeightSize); + write_leb_128(stream, copy->pawnPairPsqtData(), PairPsqtWeightSize); write_leb_128(stream, copy->weights); write_leb_128(stream, copy->psqtWeights); @@ -432,23 +428,16 @@ class FeatureTransformer { return psqt; } // end of function transform() - alignas(CacheLineSize) std::array biases; - alignas( - CacheLineSize) std::array weights; + alignas(CacheLineSize) BiasesArray biases; + alignas(CacheLineSize) WeightArray weights; // Threats and pawn-pair features are concatenated into one array to allow for a single index to address either. // The first pawn-pair feature is at index ThreatFeatureSet::Dimensions. static_assert(PairFeatureSet::IndexBase == ThreatFeatureSet::Dimensions); - alignas(CacheLineSize) std::array threatAndPpWeights; - alignas(CacheLineSize) - std::array psqtWeights; - // As above - alignas(CacheLineSize) std::array threatAndPpPsqtWeights; + alignas(CacheLineSize) ThreatAndPpWeightArray threatAndPpWeights; + alignas(CacheLineSize) PsqtWeightArray psqtWeights; + alignas(CacheLineSize) ThreatAndPpPsqtArray threatAndPpPsqtWeights; }; } // namespace Stockfish::Eval::NNUE