Skip to content

Commit 9eab47d

Browse files
committed
introduce nnVersionsDictionary
1 parent f668a22 commit 9eab47d

1 file changed

Lines changed: 26 additions & 0 deletions

File tree

‎Common/Tools/PID/pidTPCModule.h‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@
4747
#include <TRandom.h>
4848
#include <TString.h>
4949

50+
#include <array>
5051
#include <chrono>
5152
#include <cstddef>
5253
#include <cstdint>
@@ -55,6 +56,7 @@
5556
#include <memory>
5657
#include <ratio>
5758
#include <string>
59+
#include <string_view>
5860
#include <vector>
5961

6062
#include <math.h>
@@ -454,6 +456,19 @@ class pidTPCModule
454456
constexpr double Ft0cOccupancyNorm = 60000.;
455457
constexpr int NumberOfTpcSectors = 18;
456458

459+
struct NNVersionEntry {
460+
std::string_view versionName{};
461+
int numberOfFeatures{};
462+
int versionNumber{};
463+
};
464+
465+
constexpr std::array<NNVersionEntry, 5> nnVersionsDictionary{
466+
{{"", 6, 1},
467+
{"1", 6, 1},
468+
{"2", 7, 2},
469+
{"3", 8, 3},
470+
{"4", 9, 4}}};
471+
457472
std::vector<float> networkPrediction;
458473

459474
const auto startNetworkTotal = std::chrono::high_resolution_clock::now();
@@ -513,6 +528,17 @@ class pidTPCModule
513528
const uint64_t trackPropSize = inputDimensions * size;
514529
const uint64_t predictionSize = outputDimensions * size;
515530

531+
int nnVersion{0};
532+
for (const auto& nnVersionEntry : nnVersionsDictionary) {
533+
if (networkVersion == nnVersionEntry.versionName && inputDimensions == nnVersionEntry.numberOfFeatures) {
534+
nnVersion = nnVersionEntry.versionNumber;
535+
break;
536+
}
537+
}
538+
if (nnVersion == 0) {
539+
LOG(fatal) << "createNetworkPrediction(): networkVersion '" << networkVersion << "' and number of features " << inputDimensions << " are not compatible according to nnVersionsDictionary";
540+
}
541+
516542
networkPrediction = std::vector<float>(predictionSize * NParticleTypes); // For each mass hypotheses
517543
const float nNclNormalization = response->GetNClNormalization();
518544
float durationNetwork = 0;

0 commit comments

Comments
 (0)