Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 33 additions & 44 deletions src/nnue/nnue_feature_transformer.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<BiasType, HalfDimensions>;
using WeightArray = std::array<WeightType, HalfDimensions * PsqDimensions>;
using ThreatAndPpWeightArray = std::array<ThreatWeightType, ThreatAndPpWeightSize>;
using PsqtWeightArray = std::array<PSQTWeightType, PSQTBuckets * PsqDimensions>;
using ThreatAndPpPsqtArray = std::array<PSQTWeightType, ThreatAndPpPsqtSize>;

// Size of forward propagation buffer
static constexpr usize BufferSize = OutputDimensions * sizeof(OutputType);
Expand Down Expand Up @@ -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<ThreatWeightType>(stream, threatWeights(),
ThreatInputDimensions * HalfDimensions);
read_leb_128(stream, threatPsqtWeights(), ThreatFeatureSet::Dimensions * PSQTBuckets);
read_little_endian<ThreatWeightType>(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);
Expand All @@ -193,14 +193,10 @@ class FeatureTransformer {
write_leb_128<BiasType>(stream, copy->biases);


write_little_endian<ThreatWeightType>(stream, copy->threatWeights(),
ThreatInputDimensions * HalfDimensions);
write_leb_128<PSQTWeightType>(stream, copy->threatPsqtWeights(),
ThreatFeatureSet::Dimensions * PSQTBuckets);
write_little_endian<ThreatWeightType>(stream, copy->ppWeights(),
PairInputDimensions * HalfDimensions);
write_leb_128<PSQTWeightType>(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<WeightType>(stream, copy->weights);
write_leb_128<PSQTWeightType>(stream, copy->psqtWeights);
Expand Down Expand Up @@ -432,23 +428,16 @@ class FeatureTransformer {
return psqt;
} // end of function transform()

alignas(CacheLineSize) std::array<BiasType, HalfDimensions> biases;
alignas(
CacheLineSize) std::array<WeightType, HalfDimensions * PSQFeatureSet::Dimensions> 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<ThreatWeightType,
(ThreatFeatureSet::Dimensions + PairFeatureSet::Dimensions)
* HalfDimensions> threatAndPpWeights;
alignas(CacheLineSize)
std::array<PSQTWeightType, PSQTBuckets * PSQFeatureSet::Dimensions> psqtWeights;
// As above
alignas(CacheLineSize) std::array<PSQTWeightType,
(ThreatFeatureSet::Dimensions + PairFeatureSet::Dimensions)
* PSQTBuckets> threatAndPpPsqtWeights;
alignas(CacheLineSize) ThreatAndPpWeightArray threatAndPpWeights;
alignas(CacheLineSize) PsqtWeightArray psqtWeights;
alignas(CacheLineSize) ThreatAndPpPsqtArray threatAndPpPsqtWeights;
};

} // namespace Stockfish::Eval::NNUE
Expand Down
Loading